ARotting's picture
Publish First-order MAML, pooled, and random sine initializations
47dfc70 verified
Raw
History Blame Contribute Delete
7.78 kB
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()