File size: 2,363 Bytes
675a89a | 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 | from __future__ import annotations
import json
from pathlib import Path
import gradio as gr
import plotly.graph_objects as go
import torch
from data import multiscale_batch
from model import ClockworkRNN, MatchedGRU, PlainRNN
from safetensors.torch import load_file
ARTIFACT_DIR = Path(__file__).resolve().parent / "artifacts" / "clockwork-rnn-pocket"
REPORT = json.loads((ARTIFACT_DIR / "evaluation.json").read_text(encoding="utf-8"))
MODELS = {
"Clockwork RNN": (ClockworkRNN(), "clockwork_rnn"),
"Plain RNN": (PlainRNN(), "plain_rnn"),
"Matched GRU": (MatchedGRU(), "matched_gru"),
}
for model, key in MODELS.values():
model.load_state_dict(load_file(ARTIFACT_DIR / f"{key}.safetensors"))
model.eval()
@torch.inference_mode()
def compare(seed: int, length: int) -> tuple[go.Figure, dict]:
sequence = multiscale_batch(1, int(length) + 1, int(seed))
target = sequence[0, 1:, 0]
figure = go.Figure()
figure.add_trace(go.Scatter(y=target, name="Target", line={"width": 4}))
metrics = {}
for label, (model, key) in MODELS.items():
prediction = model(sequence[:, :-1])[0, :, 0]
figure.add_trace(go.Scatter(y=prediction, name=label))
metrics[label] = {
"live_rmse": float((prediction - target).square().mean().sqrt()),
"verified_length_256_rmse": REPORT["results"][key][
"length_256_zero_shot"
]["rmse"],
}
figure.update_layout(
template="plotly_dark",
title="Teacher-forced next-step multiscale forecast",
xaxis_title="Time",
yaxis_title="Signal",
)
return figure, metrics
with gr.Blocks(title="Clockwork RNN Pocket") as demo:
gr.Markdown(
"# Clockwork RNN Pocket\n"
"Compare periodic hidden-state updates with parameter-matched recurrent "
"controls on a multiscale signal."
)
with gr.Row():
seed = gr.Slider(0, 100_000, value=2099, step=1, label="Signal seed")
length = gr.Slider(32, 256, value=128, step=16, label="Sequence length")
initial = compare(2099, 128)
chart = gr.Plot(value=initial[0])
metrics = gr.JSON(value=initial[1])
button = gr.Button("Run recurrent retest", variant="primary")
button.click(compare, inputs=[seed, length], outputs=[chart, metrics])
if __name__ == "__main__":
demo.launch()
|