| from __future__ import annotations |
|
|
| import json |
| from collections import OrderedDict |
| from pathlib import Path |
|
|
| import numpy as np |
| import pandas as pd |
| import torch |
| import trackio |
| from model import SineRegressor, parameter_count |
| from safetensors.torch import save_file |
| from tasks import sample_points, sample_task |
| from torch.func import functional_call |
| from torch.nn import functional as F |
|
|
| PROJECT_DIR = Path(__file__).resolve().parent |
| ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "meta-sine-foundry" |
| DATA_DIR = PROJECT_DIR / "data" |
| INNER_LR = 0.01 |
|
|
|
|
| def adapted_parameters( |
| model: SineRegressor, |
| parameters: OrderedDict, |
| support_x: torch.Tensor, |
| support_y: torch.Tensor, |
| *, |
| create_graph: bool, |
| ) -> OrderedDict: |
| prediction = functional_call(model, parameters, (support_x,)) |
| loss = F.mse_loss(prediction, support_y) |
| gradients = torch.autograd.grad( |
| loss, |
| tuple(parameters.values()), |
| create_graph=create_graph, |
| ) |
| if not create_graph: |
| gradients = tuple(gradient.detach() for gradient in gradients) |
| return OrderedDict( |
| (name, parameter - INNER_LR * gradient) |
| for (name, parameter), gradient in zip( |
| parameters.items(), gradients, strict=True |
| ) |
| ) |
|
|
|
|
| def meta_train(iterations: int = 2_000) -> tuple[SineRegressor, list[dict]]: |
| model = SineRegressor() |
| optimizer = torch.optim.Adam(model.parameters(), lr=1e-3) |
| rng = np.random.default_rng(2043) |
| history = [] |
| for iteration in range(1, iterations + 1): |
| query_losses = [] |
| base_parameters = OrderedDict(model.named_parameters()) |
| for _ in range(12): |
| amplitude, phase = sample_task(rng) |
| support_x, support_y = sample_points( |
| amplitude, phase, 10, rng |
| ) |
| query_x, query_y = sample_points(amplitude, phase, 20, rng) |
| adapted = base_parameters |
| for _ in range(5): |
| adapted = adapted_parameters( |
| model, |
| adapted, |
| support_x, |
| support_y, |
| create_graph=False, |
| ) |
| query_losses.append( |
| F.mse_loss( |
| functional_call(model, adapted, (query_x,)), |
| query_y, |
| ) |
| ) |
| loss = torch.stack(query_losses).mean() |
| optimizer.zero_grad() |
| loss.backward() |
| optimizer.step() |
| if iteration % 100 == 0: |
| record = { |
| "training_iteration": iteration, |
| "meta_query_mse": float(loss.detach()), |
| } |
| history.append(record) |
| trackio.log(record) |
| return model, history |
|
|
|
|
| def pooled_train(iterations: int = 2_000) -> SineRegressor: |
| model = SineRegressor() |
| optimizer = torch.optim.Adam(model.parameters(), lr=1e-3) |
| rng = np.random.default_rng(3043) |
| for _ in range(iterations): |
| predictions = [] |
| targets = [] |
| for _ in range(12): |
| amplitude, phase = sample_task(rng) |
| x, y = sample_points(amplitude, phase, 30, rng) |
| predictions.append(model(x)) |
| targets.append(y) |
| loss = F.mse_loss(torch.cat(predictions), torch.cat(targets)) |
| optimizer.zero_grad() |
| loss.backward() |
| optimizer.step() |
| return model |
|
|
|
|
| def task_error( |
| model: SineRegressor, |
| amplitude: float, |
| phase: float, |
| seed: int, |
| adaptation_steps: int, |
| ) -> float: |
| rng = np.random.default_rng(seed) |
| support_x, support_y = sample_points(amplitude, phase, 5, rng) |
| query_x, query_y = sample_points(amplitude, phase, 100, rng) |
| parameters = OrderedDict( |
| (name, parameter.detach().clone().requires_grad_(True)) |
| for name, parameter in model.named_parameters() |
| ) |
| for _ in range(adaptation_steps): |
| parameters = adapted_parameters( |
| model, |
| parameters, |
| support_x, |
| support_y, |
| create_graph=False, |
| ) |
| with torch.no_grad(): |
| prediction = functional_call(model, parameters, (query_x,)) |
| return float(F.mse_loss(prediction, query_y)) |
|
|
|
|
| def evaluate_models(models: dict[str, SineRegressor]) -> tuple[dict, list[dict]]: |
| rng = np.random.default_rng(4043) |
| tasks = [sample_task(rng) for _ in range(200)] |
| rows = [] |
| for task_index, (amplitude, phase) in enumerate(tasks): |
| for name, model in models.items(): |
| for steps in [0, 1, 5]: |
| rows.append( |
| { |
| "task": task_index, |
| "amplitude": amplitude, |
| "phase": phase, |
| "model": name, |
| "adaptation_steps": steps, |
| "query_mse": task_error( |
| model, |
| amplitude, |
| phase, |
| seed=50_000 + task_index, |
| adaptation_steps=steps, |
| ), |
| } |
| ) |
| summary = {} |
| for name in models: |
| summary[name] = {} |
| for steps in [0, 1, 5]: |
| values = [ |
| row["query_mse"] |
| for row in rows |
| if row["model"] == name and row["adaptation_steps"] == steps |
| ] |
| summary[name][f"after_{steps}_steps"] = { |
| "mean_query_mse": float(np.mean(values)), |
| "median_query_mse": float(np.median(values)), |
| } |
| return summary, rows |
|
|
|
|
| def main() -> None: |
| torch.manual_seed(2043) |
| torch.set_num_threads(1) |
| trackio.init( |
| project="meta-sine-foundry", |
| name="first-order-maml-v1", |
| config={ |
| "meta_iterations": 2_000, |
| "tasks_per_iteration": 12, |
| "meta_inner_steps": 5, |
| "support_points_train": 10, |
| "support_points_test": 5, |
| "inner_learning_rate": INNER_LR, |
| }, |
| ) |
| meta_model, history = meta_train() |
| pooled_model = pooled_train() |
| torch.manual_seed(5043) |
| random_model = SineRegressor() |
| models = { |
| "first_order_maml": meta_model, |
| "pooled_pretraining": pooled_model, |
| "random_initialization": random_model, |
| } |
| summary, rows = evaluate_models(models) |
| report = { |
| "benchmark": "Five-shot sinusoid adaptation", |
| "parameters_per_model": parameter_count(meta_model), |
| "heldout_tasks": 200, |
| "support_points": 5, |
| "inner_learning_rate": INNER_LR, |
| "summary": summary, |
| "training_history": history, |
| } |
| ARTIFACT_DIR.mkdir(parents=True, exist_ok=True) |
| DATA_DIR.mkdir(parents=True, exist_ok=True) |
| for name, model in models.items(): |
| save_file(model.state_dict(), ARTIFACT_DIR / f"{name}.safetensors") |
| (ARTIFACT_DIR / "evaluation.json").write_text( |
| json.dumps(report, indent=2), encoding="utf-8" |
| ) |
| pd.DataFrame(rows).to_parquet(DATA_DIR / "heldout_tasks.parquet", index=False) |
| trackio.log( |
| { |
| "maml_one_step_mse": summary["first_order_maml"]["after_1_steps"][ |
| "mean_query_mse" |
| ], |
| "maml_five_step_mse": summary["first_order_maml"]["after_5_steps"][ |
| "mean_query_mse" |
| ], |
| "pooled_five_step_mse": summary["pooled_pretraining"][ |
| "after_5_steps" |
| ]["mean_query_mse"], |
| "random_five_step_mse": summary["random_initialization"][ |
| "after_5_steps" |
| ]["mean_query_mse"], |
| } |
| ) |
| trackio.finish() |
| print(json.dumps(report, indent=2)) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|