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