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