ARotting's picture
Publish Event-gated hierarchical recurrent memory retest
25021c2 verified
Raw
History Blame Contribute Delete
3.36 kB
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()