| from __future__ import annotations |
|
|
| from pathlib import Path |
|
|
| import gradio as gr |
| import plotly.graph_objects as go |
| import torch |
| from data import generate_adding_problem |
| from model import SequenceRegressor |
| from safetensors.torch import load_file |
|
|
| ARTIFACT_DIR = Path(__file__).resolve().parent / "artifacts" / "lstm-time-capsule" |
| MODELS = {} |
| for name, cell in [("Vanilla RNN", "rnn"), ("LSTM", "lstm"), ("GRU", "gru")]: |
| model = SequenceRegressor(cell) |
| filename = name.lower().replace(" ", "_") |
| model.load_state_dict(load_file(ARTIFACT_DIR / f"{filename}.safetensors")) |
| model.eval() |
| MODELS[name] = model |
|
|
|
|
| def inspect_delay(seed: int, length: int) -> tuple[go.Figure, dict]: |
| inputs, targets = generate_adding_problem(1, int(length), int(seed)) |
| sequence = torch.from_numpy(inputs) |
| with torch.inference_mode(): |
| predictions = { |
| name: float(model(sequence)[0]) for name, model in MODELS.items() |
| } |
| markers = inputs[0, :, 1] |
| values = inputs[0, :, 0] |
| colors = ["#f59e0b" if marker else "#334155" for marker in markers] |
| figure = go.Figure( |
| go.Bar(x=list(range(len(values))), y=values, marker_color=colors) |
| ) |
| figure.update_layout( |
| title="Long-lag adding sequence (orange values are remembered)", |
| xaxis_title="Time step", |
| yaxis_title="Input value", |
| template="plotly_dark", |
| ) |
| target = float(targets[0]) |
| return figure, { |
| "target_sum": round(target, 4), |
| **{ |
| name: { |
| "prediction": round(value, 4), |
| "absolute_error": round(abs(value - target), 4), |
| } |
| for name, value in predictions.items() |
| }, |
| } |
|
|
|
|
| with gr.Blocks(title="LSTM Time Capsule") as demo: |
| gr.Markdown( |
| "# LSTM Time Capsule\n" |
| "Place two values hundreds of steps apart and compare recurrent memory." |
| ) |
| with gr.Row(): |
| seed = gr.Number(2043, precision=0, label="Sequence seed") |
| length = gr.Slider(100, 400, 100, step=50, label="Sequence length") |
| run = gr.Button("Test the time lag", variant="primary") |
| sequence = gr.Plot() |
| predictions = gr.JSON() |
| run.click(inspect_delay, [seed, length], [sequence, predictions]) |
| demo.load(inspect_delay, [seed, length], [sequence, predictions]) |
|
|
|
|
| if __name__ == "__main__": |
| demo.launch() |
|
|