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