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()