File size: 3,417 Bytes
0fc8fd1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
from __future__ import annotations

import json
from pathlib import Path

import gradio as gr
import numpy as np
import plotly.graph_objects as go
import torch
from model import ConditionalNeuralProcess
from safetensors.torch import load_file

ARTIFACT_DIR = Path(__file__).resolve().parent / "artifacts" / "neural-process-pocket"
MODEL = ConditionalNeuralProcess()
MODEL.load_state_dict(load_file(ARTIFACT_DIR / "model.safetensors"))
MODEL.eval()
REPORT = json.loads((ARTIFACT_DIR / "evaluation.json").read_text(encoding="utf-8"))


@torch.inference_mode()
def infer_function(
    amplitude: float,
    phase: float,
    frequency: float,
    context_points: int,
    seed: int,
) -> tuple[go.Figure, dict]:
    rng = np.random.default_rng(int(seed))
    context_x = np.sort(rng.uniform(-5, 5, int(context_points))).astype(np.float32)
    context_y = amplitude * np.sin(frequency * context_x + phase)
    target_x = np.linspace(-5, 5, 300, dtype=np.float32)
    true_y = amplitude * np.sin(frequency * target_x + phase)
    mean, std = MODEL(
        torch.from_numpy(context_x)[None, :, None],
        torch.from_numpy(context_y.astype(np.float32))[None, :, None],
        torch.from_numpy(target_x)[None, :, None],
    )
    mean = mean[0, :, 0].numpy()
    std = std[0, :, 0].numpy()
    figure = go.Figure()
    figure.add_trace(
        go.Scatter(
            x=np.concatenate([target_x, target_x[::-1]]),
            y=np.concatenate([mean - 1.645 * std, (mean + 1.645 * std)[::-1]]),
            fill="toself",
            name="90% predictive interval",
            line={"color": "rgba(0,0,0,0)"},
        )
    )
    figure.add_trace(go.Scatter(x=target_x, y=true_y, name="True function"))
    figure.add_trace(go.Scatter(x=target_x, y=mean, name="CNP mean"))
    figure.add_trace(
        go.Scatter(x=context_x, y=context_y, mode="markers", name="Context")
    )
    figure.update_layout(template="plotly_dark", xaxis_title="x", yaxis_title="y")
    metrics = {
        "live_rmse": float(np.sqrt(np.mean((mean - true_y) ** 2))),
        "context_points": int(context_points),
        "verified_500_task_rmse": REPORT["benchmark"][
            "conditional_neural_process"
        ]["rmse"],
        "verified_90_percent_coverage": REPORT["benchmark"][
            "conditional_neural_process"
        ]["coverage_90"],
    }
    return figure, metrics


with gr.Blocks(title="Neural Process Pocket") as demo:
    gr.Markdown(
        "# Neural Process Pocket\n"
        "Give the model a few observations and watch it infer a complete function "
        "distribution with a learned uncertainty band."
    )
    with gr.Row():
        amplitude = gr.Slider(0.1, 5.0, value=2.5, step=0.1, label="Amplitude")
        phase = gr.Slider(0, 6.28, value=1.0, step=0.05, label="Phase")
        frequency = gr.Slider(0.8, 1.2, value=1.0, step=0.02, label="Frequency")
        context = gr.Slider(3, 20, value=5, step=1, label="Context points")
        seed = gr.Slider(0, 100_000, value=2141, step=1, label="Seed")
    initial = infer_function(2.5, 1.0, 1.0, 5, 2141)
    chart = gr.Plot(value=initial[0])
    metrics = gr.JSON(value=initial[1])
    button = gr.Button("Infer function distribution", variant="primary")
    button.click(
        infer_function,
        inputs=[amplitude, phase, frequency, context, seed],
        outputs=[chart, metrics],
    )


if __name__ == "__main__":
    demo.launch()