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 HistoryCompressor, PlainRNN from safetensors.torch import load_file from train import sample_batch PROJECT_DIR = Path(__file__).resolve().parent ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "history-compressor-pocket" COMPRESSOR = HistoryCompressor() COMPRESSOR.load_state_dict(load_file(ARTIFACT_DIR / "compressor.safetensors")) COMPRESSOR.eval() PLAIN = PlainRNN() PLAIN.load_state_dict(load_file(ARTIFACT_DIR / "plain_rnn.safetensors")) PLAIN.eval() REPORT = json.loads((ARTIFACT_DIR / "evaluation.json").read_text(encoding="utf-8")) @torch.inference_mode() def inspect_history(gap: int, seed: int) -> tuple[go.Figure, dict]: generator = torch.Generator().manual_seed(int(seed)) tokens, targets, mask = sample_batch( 1, blocks=4, gap=int(gap), generator=generator, ) compressor_logits, compressor_states = COMPRESSOR(tokens, return_states=True) plain_logits, plain_states = PLAIN(tokens, return_states=True) compressor_norm = compressor_states[0].norm(dim=1).numpy() plain_norm = plain_states[0].norm(dim=1).numpy() figure = go.Figure() figure.add_scatter(y=compressor_norm, name="Event-gated slow state") figure.add_scatter(y=plain_norm, name="Plain recurrent state", opacity=0.7) for boundary in range(int(gap) - 1, len(tokens[0]), int(gap)): figure.add_vline(x=boundary, line_dash="dot", line_color="#ffd166") figure.update_layout( template="plotly_dark", title="Hidden-state norm across control events and distractors", xaxis_title="Sequence step", yaxis_title="State norm", ) expected = targets[mask].tolist() compressor_predictions = compressor_logits[mask].argmax(1).tolist() plain_predictions = plain_logits[mask].argmax(1).tolist() metrics = { "gap": int(gap), "maximum_training_gap": 32, "expected_boundary_states": expected, "compressor_predictions": compressor_predictions, "plain_rnn_predictions": plain_predictions, "verified_gap_128_compressor_accuracy": REPORT["results"][ "history_compressor" ]["accuracy_mean"]["gap_128"], "verified_gap_128_plain_accuracy": REPORT["results"]["plain_rnn"][ "accuracy_mean" ]["gap_128"], "visible_distractor_fraction": float( np.mean(tokens[0].numpy() >= 8) ), } return figure, metrics with gr.Blocks(title="History Compressor Pocket") as demo: gr.Markdown( "# History Compressor Pocket\n" "Watch a slow recurrent level update only on informative control events " "while a parameter-matched plain RNN is perturbed by every distractor." ) with gr.Row(): gap = gr.Slider(8, 128, value=64, step=8, label="Distractor gap") seed = gr.Slider(0, 10_000, value=71, step=1, label="Sequence seed") initial = inspect_history(64, 71) chart = gr.Plot(value=initial[0]) metrics = gr.JSON(value=initial[1]) button = gr.Button("Compress another history", variant="primary") button.click(inspect_history, inputs=[gap, seed], outputs=[chart, metrics]) if __name__ == "__main__": demo.launch()