ARotting's picture
Publish Parameter-matched RNN, LSTM, and GRU long-lag retest
ddaaeb2 verified
Raw
History Blame Contribute Delete
2.36 kB
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()