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