ARotting's picture
Publish First-order MAML, pooled, and random sine initializations
47dfc70 verified
Raw
History Blame Contribute Delete
3.49 kB
from __future__ import annotations
from collections import OrderedDict
from pathlib import Path
import gradio as gr
import numpy as np
import plotly.graph_objects as go
import torch
from model import SineRegressor
from safetensors.torch import load_file
from tasks import sample_points
from train import adapted_parameters
ARTIFACT_DIR = Path(__file__).resolve().parent / "artifacts" / "meta-sine-foundry"
MODELS = {}
for label, filename in [
("First-order MAML", "first_order_maml"),
("Pooled pretraining", "pooled_pretraining"),
("Random initialization", "random_initialization"),
]:
model = SineRegressor()
model.load_state_dict(load_file(ARTIFACT_DIR / f"{filename}.safetensors"))
model.eval()
MODELS[label] = model
def adapt_curve(
amplitude: float, phase: float, steps: int, seed: int
) -> tuple[go.Figure, dict]:
rng = np.random.default_rng(int(seed))
support_x, support_y = sample_points(amplitude, phase, 5, rng)
grid = torch.linspace(-5, 5, 300).unsqueeze(1)
truth = amplitude * torch.sin(grid + phase)
figure = go.Figure()
figure.add_trace(
go.Scatter(
x=grid[:, 0],
y=truth[:, 0],
mode="lines",
name="True task",
line={"color": "#ffffff", "width": 3},
)
)
errors = {}
for label, model in MODELS.items():
parameters = OrderedDict(
(name, parameter.detach().clone().requires_grad_(True))
for name, parameter in model.named_parameters()
)
for _ in range(int(steps)):
parameters = adapted_parameters(
model,
parameters,
support_x,
support_y,
create_graph=False,
)
with torch.no_grad():
prediction = torch.func.functional_call(model, parameters, (grid,))
error = float(torch.mean((prediction - truth) ** 2))
errors[label] = round(error, 5)
figure.add_trace(
go.Scatter(
x=grid[:, 0],
y=prediction[:, 0],
mode="lines",
name=label,
)
)
figure.add_trace(
go.Scatter(
x=support_x[:, 0],
y=support_y[:, 0],
mode="markers",
name="Five support points",
marker={"size": 10, "color": "#f59e0b"},
)
)
figure.update_layout(
title=f"Five-shot adaptation after {int(steps)} gradient steps",
xaxis_title="x",
yaxis_title="y",
template="plotly_dark",
)
return figure, errors
with gr.Blocks(title="Meta-Sine Foundry") as demo:
gr.Markdown(
"# Meta-Sine Foundry\n"
"Create an unseen sinusoid, reveal five points, and watch three "
"initializations adapt with the same gradient rule."
)
with gr.Row():
amplitude = gr.Slider(0.1, 5.0, 3.0, step=0.1, label="Amplitude")
phase = gr.Slider(0, float(np.pi), 1.0, step=0.1, label="Phase")
steps = gr.Slider(0, 10, 5, step=1, label="Adaptation steps")
seed = gr.Number(2043, precision=0, label="Support seed")
run = gr.Button("Adapt to task", variant="primary")
curves = gr.Plot()
errors = gr.JSON()
run.click(adapt_curve, [amplitude, phase, steps, seed], [curves, errors])
demo.load(adapt_curve, [amplitude, phase, steps, seed], [curves, errors])
if __name__ == "__main__":
demo.launch()