Download app.py from ARotting/memory-tape-pocket-lab: direct link, hf CLI and curl.
- Browser
- Download file 3.1 kB
-
https://huggingface.co/ARotting/memory-tape-pocket-lab/resolve/main/app.py
- Command line
-
hf download hf://ARotting/memory-tape-pocket-lab/app.py
-
curl -L -o app.py https://huggingface.co/ARotting/memory-tape-pocket-lab/resolve/main/app.py
3.1 kB
| from __future__ import annotations | |
| import json | |
| from pathlib import Path | |
| import gradio as gr | |
| import plotly.graph_objects as go | |
| import torch | |
| from model import ContentAddressedMemory, FixedStateGRU | |
| from safetensors.torch import load_file | |
| from train import sample_batch | |
| PROJECT_DIR = Path(__file__).resolve().parent | |
| ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "memory-tape-pocket" | |
| MEMORY = ContentAddressedMemory() | |
| MEMORY.load_state_dict(load_file(ARTIFACT_DIR / "content_memory.safetensors")) | |
| MEMORY.eval() | |
| GRU = FixedStateGRU() | |
| GRU.load_state_dict(load_file(ARTIFACT_DIR / "fixed_gru.safetensors")) | |
| GRU.eval() | |
| REPORT = json.loads((ARTIFACT_DIR / "evaluation.json").read_text(encoding="utf-8")) | |
| def inspect_tape(slots: int, seed: int) -> tuple[go.Figure, dict]: | |
| generator = torch.Generator().manual_seed(int(seed)) | |
| keys, values, query, target = sample_batch(1, int(slots), generator) | |
| memory_logits, attention = MEMORY( | |
| keys, | |
| values, | |
| query, | |
| return_attention=True, | |
| ) | |
| gru_logits = GRU(keys, values, query) | |
| weights = attention[0].numpy() | |
| labels = [ | |
| f"slot {index}: {int(key)} → {int(value)}" | |
| for index, (key, value) in enumerate(zip(keys[0], values[0], strict=True)) | |
| ] | |
| figure = go.Figure(go.Bar(x=labels, y=weights)) | |
| figure.update_layout( | |
| template="plotly_dark", | |
| title=f"Content-addressed read weights for query key {int(query)}", | |
| xaxis_title="External memory tape", | |
| yaxis_title="Attention weight", | |
| yaxis_range=[0, 1], | |
| ) | |
| correct_slot = int(keys[0].eq(query[0]).nonzero()[0]) | |
| result = { | |
| "query_key": int(query), | |
| "target_value": int(target), | |
| "content_memory_prediction": int(memory_logits.argmax(1)), | |
| "fixed_gru_prediction": int(gru_logits.argmax(1)), | |
| "correct_slot": correct_slot, | |
| "attention_on_correct_slot": float(attention[0, correct_slot]), | |
| "training_tape_length": "2 to 8 slots", | |
| "verified_32_slot_memory_accuracy": REPORT["results"]["memory"][ | |
| "accuracy_mean" | |
| ]["slots_32"], | |
| "verified_32_slot_gru_accuracy": REPORT["results"]["gru"][ | |
| "accuracy_mean" | |
| ]["slots_32"], | |
| } | |
| return figure, result | |
| with gr.Blocks(title="Memory Tape Pocket") as demo: | |
| gr.Markdown( | |
| "# Memory Tape Pocket\n" | |
| "A differentiable content-addressed tape retrieves random key-value " | |
| "bindings. Compare its read head with a larger GRU that compresses the " | |
| "whole tape into one fixed state." | |
| ) | |
| with gr.Row(): | |
| slots = gr.Slider(2, 32, value=16, step=1, label="Memory slots") | |
| seed = gr.Slider(0, 10_000, value=42, step=1, label="Episode seed") | |
| initial = inspect_tape(16, 42) | |
| chart = gr.Plot(value=initial[0], label="Differentiable read head") | |
| metrics = gr.JSON(value=initial[1]) | |
| button = gr.Button("Generate a new tape", variant="primary") | |
| button.click(inspect_tape, inputs=[slots, seed], outputs=[chart, metrics]) | |
| if __name__ == "__main__": | |
| demo.launch() | |