Publish Compact content-addressed differentiable memory
Browse files- README.md +38 -0
- content_memory.safetensors +3 -0
- evaluation.json +133 -0
- fixed_gru.safetensors +3 -0
- source/app.py +85 -0
- source/model.py +76 -0
- source/requirements.txt +6 -0
- source/train.py +226 -0
README.md
ADDED
|
@@ -0,0 +1,38 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
title: Memory Tape Pocket
|
| 3 |
+
emoji: 🧠
|
| 4 |
+
colorFrom: indigo
|
| 5 |
+
colorTo: pink
|
| 6 |
+
sdk: gradio
|
| 7 |
+
sdk_version: "6.5.1"
|
| 8 |
+
app_file: app.py
|
| 9 |
+
pinned: false
|
| 10 |
+
---
|
| 11 |
+
|
| 12 |
+
# Memory Tape Pocket
|
| 13 |
+
|
| 14 |
+
Memory Tape Pocket is a compact differentiable-memory retest inspired by the
|
| 15 |
+
content-addressing mechanism of Neural Turing Machines. It learns random
|
| 16 |
+
key-value associative recall on tapes containing two to eight slots, then faces
|
| 17 |
+
unseen tapes with 16 and 32 slots.
|
| 18 |
+
|
| 19 |
+
The control is a larger fixed-state GRU trained on the same batches. The
|
| 20 |
+
interactive Space exposes the complete external tape and the learned read
|
| 21 |
+
weight assigned to every slot.
|
| 22 |
+
|
| 23 |
+
## Verified result
|
| 24 |
+
|
| 25 |
+
Across three independent training seeds, the 4,673-parameter content-addressed
|
| 26 |
+
model achieved **100% exact recall** on 8-, 16-, and 32-slot tapes. At 32 slots,
|
| 27 |
+
four times the maximum training length, its read head placed **99.974%** of its
|
| 28 |
+
attention on the correct slot.
|
| 29 |
+
|
| 30 |
+
The larger 5,584-parameter fixed-state GRU reached 13.51% accuracy at eight
|
| 31 |
+
slots, 7.66% at 16 slots, and **4.60% at 32 slots**. This benchmark isolates the
|
| 32 |
+
inductive bias of external content addressing; it does not claim the tiny model
|
| 33 |
+
implements every component of a full Neural Turing Machine.
|
| 34 |
+
|
| 35 |
+
```bash
|
| 36 |
+
uv run python projects/memory-tape-pocket/train.py
|
| 37 |
+
uv run pytest tests/test_memory_tape_pocket.py
|
| 38 |
+
```
|
content_memory.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:2057913703db35ea2ede700d11fc55cfec73073c7c776638fd9b4a097aa75639
|
| 3 |
+
size 19084
|
evaluation.json
ADDED
|
@@ -0,0 +1,133 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"experiment": "Differentiable content addressing versus fixed-state recall",
|
| 3 |
+
"training_slots": [
|
| 4 |
+
2,
|
| 5 |
+
8
|
| 6 |
+
],
|
| 7 |
+
"results": {
|
| 8 |
+
"memory": {
|
| 9 |
+
"parameters": 4673,
|
| 10 |
+
"runs": [
|
| 11 |
+
{
|
| 12 |
+
"seed": 2281,
|
| 13 |
+
"slots_8": {
|
| 14 |
+
"accuracy": 1.0,
|
| 15 |
+
"examples": 4096,
|
| 16 |
+
"mean_attention_on_correct_slot": 0.9999377218191512
|
| 17 |
+
},
|
| 18 |
+
"slots_16": {
|
| 19 |
+
"accuracy": 1.0,
|
| 20 |
+
"examples": 4096,
|
| 21 |
+
"mean_attention_on_correct_slot": 0.9998665036546299
|
| 22 |
+
},
|
| 23 |
+
"slots_32": {
|
| 24 |
+
"accuracy": 1.0,
|
| 25 |
+
"examples": 4096,
|
| 26 |
+
"mean_attention_on_correct_slot": 0.9997205645777285
|
| 27 |
+
}
|
| 28 |
+
},
|
| 29 |
+
{
|
| 30 |
+
"seed": 2287,
|
| 31 |
+
"slots_8": {
|
| 32 |
+
"accuracy": 1.0,
|
| 33 |
+
"examples": 4096,
|
| 34 |
+
"mean_attention_on_correct_slot": 0.9999373428727267
|
| 35 |
+
},
|
| 36 |
+
"slots_16": {
|
| 37 |
+
"accuracy": 1.0,
|
| 38 |
+
"examples": 4096,
|
| 39 |
+
"mean_attention_on_correct_slot": 0.9998682647856185
|
| 40 |
+
},
|
| 41 |
+
"slots_32": {
|
| 42 |
+
"accuracy": 1.0,
|
| 43 |
+
"examples": 4096,
|
| 44 |
+
"mean_attention_on_correct_slot": 0.999723744156654
|
| 45 |
+
}
|
| 46 |
+
},
|
| 47 |
+
{
|
| 48 |
+
"seed": 2293,
|
| 49 |
+
"slots_8": {
|
| 50 |
+
"accuracy": 1.0,
|
| 51 |
+
"examples": 4096,
|
| 52 |
+
"mean_attention_on_correct_slot": 0.9999510854540858
|
| 53 |
+
},
|
| 54 |
+
"slots_16": {
|
| 55 |
+
"accuracy": 1.0,
|
| 56 |
+
"examples": 4096,
|
| 57 |
+
"mean_attention_on_correct_slot": 0.9998970205051592
|
| 58 |
+
},
|
| 59 |
+
"slots_32": {
|
| 60 |
+
"accuracy": 1.0,
|
| 61 |
+
"examples": 4096,
|
| 62 |
+
"mean_attention_on_correct_slot": 0.9997834917012369
|
| 63 |
+
}
|
| 64 |
+
}
|
| 65 |
+
],
|
| 66 |
+
"accuracy_mean": {
|
| 67 |
+
"slots_8": 1.0,
|
| 68 |
+
"slots_16": 1.0,
|
| 69 |
+
"slots_32": 1.0
|
| 70 |
+
},
|
| 71 |
+
"correct_slot_attention_mean": {
|
| 72 |
+
"slots_8": 0.9999420500486546,
|
| 73 |
+
"slots_16": 0.9998772629818026,
|
| 74 |
+
"slots_32": 0.9997426001452064
|
| 75 |
+
}
|
| 76 |
+
},
|
| 77 |
+
"gru": {
|
| 78 |
+
"parameters": 5584,
|
| 79 |
+
"runs": [
|
| 80 |
+
{
|
| 81 |
+
"seed": 2281,
|
| 82 |
+
"slots_8": {
|
| 83 |
+
"accuracy": 0.132080078125,
|
| 84 |
+
"examples": 4096
|
| 85 |
+
},
|
| 86 |
+
"slots_16": {
|
| 87 |
+
"accuracy": 0.083984375,
|
| 88 |
+
"examples": 4096
|
| 89 |
+
},
|
| 90 |
+
"slots_32": {
|
| 91 |
+
"accuracy": 0.044677734375,
|
| 92 |
+
"examples": 4096
|
| 93 |
+
}
|
| 94 |
+
},
|
| 95 |
+
{
|
| 96 |
+
"seed": 2287,
|
| 97 |
+
"slots_8": {
|
| 98 |
+
"accuracy": 0.135009765625,
|
| 99 |
+
"examples": 4096
|
| 100 |
+
},
|
| 101 |
+
"slots_16": {
|
| 102 |
+
"accuracy": 0.0703125,
|
| 103 |
+
"examples": 4096
|
| 104 |
+
},
|
| 105 |
+
"slots_32": {
|
| 106 |
+
"accuracy": 0.046142578125,
|
| 107 |
+
"examples": 4096
|
| 108 |
+
}
|
| 109 |
+
},
|
| 110 |
+
{
|
| 111 |
+
"seed": 2293,
|
| 112 |
+
"slots_8": {
|
| 113 |
+
"accuracy": 0.13818359375,
|
| 114 |
+
"examples": 4096
|
| 115 |
+
},
|
| 116 |
+
"slots_16": {
|
| 117 |
+
"accuracy": 0.075439453125,
|
| 118 |
+
"examples": 4096
|
| 119 |
+
},
|
| 120 |
+
"slots_32": {
|
| 121 |
+
"accuracy": 0.047119140625,
|
| 122 |
+
"examples": 4096
|
| 123 |
+
}
|
| 124 |
+
}
|
| 125 |
+
],
|
| 126 |
+
"accuracy_mean": {
|
| 127 |
+
"slots_8": 0.13509114583333334,
|
| 128 |
+
"slots_16": 0.07657877604166667,
|
| 129 |
+
"slots_32": 0.045979817708333336
|
| 130 |
+
}
|
| 131 |
+
}
|
| 132 |
+
}
|
| 133 |
+
}
|
fixed_gru.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:1f4a929a706231c551122c07cdf5788751da9145032e0b18592512e3ba398b95
|
| 3 |
+
size 22880
|
source/app.py
ADDED
|
@@ -0,0 +1,85 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import json
|
| 4 |
+
from pathlib import Path
|
| 5 |
+
|
| 6 |
+
import gradio as gr
|
| 7 |
+
import plotly.graph_objects as go
|
| 8 |
+
import torch
|
| 9 |
+
from model import ContentAddressedMemory, FixedStateGRU
|
| 10 |
+
from safetensors.torch import load_file
|
| 11 |
+
from train import sample_batch
|
| 12 |
+
|
| 13 |
+
PROJECT_DIR = Path(__file__).resolve().parent
|
| 14 |
+
ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "memory-tape-pocket"
|
| 15 |
+
MEMORY = ContentAddressedMemory()
|
| 16 |
+
MEMORY.load_state_dict(load_file(ARTIFACT_DIR / "content_memory.safetensors"))
|
| 17 |
+
MEMORY.eval()
|
| 18 |
+
GRU = FixedStateGRU()
|
| 19 |
+
GRU.load_state_dict(load_file(ARTIFACT_DIR / "fixed_gru.safetensors"))
|
| 20 |
+
GRU.eval()
|
| 21 |
+
REPORT = json.loads((ARTIFACT_DIR / "evaluation.json").read_text(encoding="utf-8"))
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
@torch.inference_mode()
|
| 25 |
+
def inspect_tape(slots: int, seed: int) -> tuple[go.Figure, dict]:
|
| 26 |
+
generator = torch.Generator().manual_seed(int(seed))
|
| 27 |
+
keys, values, query, target = sample_batch(1, int(slots), generator)
|
| 28 |
+
memory_logits, attention = MEMORY(
|
| 29 |
+
keys,
|
| 30 |
+
values,
|
| 31 |
+
query,
|
| 32 |
+
return_attention=True,
|
| 33 |
+
)
|
| 34 |
+
gru_logits = GRU(keys, values, query)
|
| 35 |
+
weights = attention[0].numpy()
|
| 36 |
+
labels = [
|
| 37 |
+
f"slot {index}: {int(key)} → {int(value)}"
|
| 38 |
+
for index, (key, value) in enumerate(zip(keys[0], values[0], strict=True))
|
| 39 |
+
]
|
| 40 |
+
figure = go.Figure(go.Bar(x=labels, y=weights))
|
| 41 |
+
figure.update_layout(
|
| 42 |
+
template="plotly_dark",
|
| 43 |
+
title=f"Content-addressed read weights for query key {int(query)}",
|
| 44 |
+
xaxis_title="External memory tape",
|
| 45 |
+
yaxis_title="Attention weight",
|
| 46 |
+
yaxis_range=[0, 1],
|
| 47 |
+
)
|
| 48 |
+
correct_slot = int(keys[0].eq(query[0]).nonzero()[0])
|
| 49 |
+
result = {
|
| 50 |
+
"query_key": int(query),
|
| 51 |
+
"target_value": int(target),
|
| 52 |
+
"content_memory_prediction": int(memory_logits.argmax(1)),
|
| 53 |
+
"fixed_gru_prediction": int(gru_logits.argmax(1)),
|
| 54 |
+
"correct_slot": correct_slot,
|
| 55 |
+
"attention_on_correct_slot": float(attention[0, correct_slot]),
|
| 56 |
+
"training_tape_length": "2 to 8 slots",
|
| 57 |
+
"verified_32_slot_memory_accuracy": REPORT["results"]["memory"][
|
| 58 |
+
"accuracy_mean"
|
| 59 |
+
]["slots_32"],
|
| 60 |
+
"verified_32_slot_gru_accuracy": REPORT["results"]["gru"][
|
| 61 |
+
"accuracy_mean"
|
| 62 |
+
]["slots_32"],
|
| 63 |
+
}
|
| 64 |
+
return figure, result
|
| 65 |
+
|
| 66 |
+
|
| 67 |
+
with gr.Blocks(title="Memory Tape Pocket") as demo:
|
| 68 |
+
gr.Markdown(
|
| 69 |
+
"# Memory Tape Pocket\n"
|
| 70 |
+
"A differentiable content-addressed tape retrieves random key-value "
|
| 71 |
+
"bindings. Compare its read head with a larger GRU that compresses the "
|
| 72 |
+
"whole tape into one fixed state."
|
| 73 |
+
)
|
| 74 |
+
with gr.Row():
|
| 75 |
+
slots = gr.Slider(2, 32, value=16, step=1, label="Memory slots")
|
| 76 |
+
seed = gr.Slider(0, 10_000, value=42, step=1, label="Episode seed")
|
| 77 |
+
initial = inspect_tape(16, 42)
|
| 78 |
+
chart = gr.Plot(value=initial[0], label="Differentiable read head")
|
| 79 |
+
metrics = gr.JSON(value=initial[1])
|
| 80 |
+
button = gr.Button("Generate a new tape", variant="primary")
|
| 81 |
+
button.click(inspect_tape, inputs=[slots, seed], outputs=[chart, metrics])
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
if __name__ == "__main__":
|
| 85 |
+
demo.launch()
|
source/model.py
ADDED
|
@@ -0,0 +1,76 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import math
|
| 4 |
+
|
| 5 |
+
import torch
|
| 6 |
+
from torch import nn
|
| 7 |
+
from torch.nn import functional as F
|
| 8 |
+
|
| 9 |
+
VOCAB_SIZE = 64
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
class ContentAddressedMemory(nn.Module):
|
| 13 |
+
"""A learned key-value tape with differentiable content addressing."""
|
| 14 |
+
|
| 15 |
+
def __init__(self, width: int = 24) -> None:
|
| 16 |
+
super().__init__()
|
| 17 |
+
self.key_embedding = nn.Embedding(VOCAB_SIZE, width)
|
| 18 |
+
self.value_embedding = nn.Embedding(VOCAB_SIZE, width)
|
| 19 |
+
self.output = nn.Linear(width, VOCAB_SIZE)
|
| 20 |
+
self.log_beta = nn.Parameter(torch.tensor(math.log(10.0)))
|
| 21 |
+
|
| 22 |
+
def forward(
|
| 23 |
+
self,
|
| 24 |
+
keys: torch.Tensor,
|
| 25 |
+
values: torch.Tensor,
|
| 26 |
+
query: torch.Tensor,
|
| 27 |
+
*,
|
| 28 |
+
return_attention: bool = False,
|
| 29 |
+
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
|
| 30 |
+
memory_keys = F.normalize(self.key_embedding(keys), dim=-1)
|
| 31 |
+
query_key = F.normalize(self.key_embedding(query), dim=-1)
|
| 32 |
+
beta = self.log_beta.exp().clamp(1.0, 30.0)
|
| 33 |
+
scores = torch.einsum("bsd,bd->bs", memory_keys, query_key) * beta
|
| 34 |
+
attention = scores.softmax(dim=-1)
|
| 35 |
+
read = torch.einsum(
|
| 36 |
+
"bs,bsd->bd",
|
| 37 |
+
attention,
|
| 38 |
+
self.value_embedding(values),
|
| 39 |
+
)
|
| 40 |
+
logits = self.output(read)
|
| 41 |
+
if return_attention:
|
| 42 |
+
return logits, attention
|
| 43 |
+
return logits
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
class FixedStateGRU(nn.Module):
|
| 47 |
+
"""A larger recurrent control that compresses the tape into one state."""
|
| 48 |
+
|
| 49 |
+
def __init__(self, embedding_dim: int = 8, hidden_dim: int = 24) -> None:
|
| 50 |
+
super().__init__()
|
| 51 |
+
self.embedding = nn.Embedding(VOCAB_SIZE * 3, embedding_dim)
|
| 52 |
+
self.gru = nn.GRU(embedding_dim, hidden_dim, batch_first=True)
|
| 53 |
+
self.output = nn.Linear(hidden_dim, VOCAB_SIZE)
|
| 54 |
+
|
| 55 |
+
def forward(
|
| 56 |
+
self,
|
| 57 |
+
keys: torch.Tensor,
|
| 58 |
+
values: torch.Tensor,
|
| 59 |
+
query: torch.Tensor,
|
| 60 |
+
) -> torch.Tensor:
|
| 61 |
+
batch, slots = keys.shape
|
| 62 |
+
tape = torch.empty(
|
| 63 |
+
batch,
|
| 64 |
+
slots * 2 + 1,
|
| 65 |
+
dtype=torch.long,
|
| 66 |
+
device=keys.device,
|
| 67 |
+
)
|
| 68 |
+
tape[:, 0 : slots * 2 : 2] = keys
|
| 69 |
+
tape[:, 1 : slots * 2 : 2] = values + VOCAB_SIZE
|
| 70 |
+
tape[:, -1] = query + VOCAB_SIZE * 2
|
| 71 |
+
_, state = self.gru(self.embedding(tape))
|
| 72 |
+
return self.output(state[-1])
|
| 73 |
+
|
| 74 |
+
|
| 75 |
+
def parameter_count(model: nn.Module) -> int:
|
| 76 |
+
return sum(parameter.numel() for parameter in model.parameters())
|
source/requirements.txt
ADDED
|
@@ -0,0 +1,6 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
gradio
|
| 2 |
+
numpy
|
| 3 |
+
plotly
|
| 4 |
+
safetensors
|
| 5 |
+
torch
|
| 6 |
+
trackio
|
source/train.py
ADDED
|
@@ -0,0 +1,226 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from __future__ import annotations
|
| 2 |
+
|
| 3 |
+
import json
|
| 4 |
+
from pathlib import Path
|
| 5 |
+
|
| 6 |
+
import numpy as np
|
| 7 |
+
import torch
|
| 8 |
+
import trackio
|
| 9 |
+
from model import (
|
| 10 |
+
VOCAB_SIZE,
|
| 11 |
+
ContentAddressedMemory,
|
| 12 |
+
FixedStateGRU,
|
| 13 |
+
parameter_count,
|
| 14 |
+
)
|
| 15 |
+
from safetensors.torch import save_file
|
| 16 |
+
from torch.nn import functional as F
|
| 17 |
+
|
| 18 |
+
PROJECT_DIR = Path(__file__).resolve().parent
|
| 19 |
+
ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "memory-tape-pocket"
|
| 20 |
+
DATA_DIR = PROJECT_DIR / "data"
|
| 21 |
+
TRAIN_SLOT_RANGE = (2, 8)
|
| 22 |
+
STEPS = 2_500
|
| 23 |
+
BATCH_SIZE = 256
|
| 24 |
+
SEEDS = [2281, 2287, 2293]
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def sample_batch(
|
| 28 |
+
batch_size: int,
|
| 29 |
+
slots: int,
|
| 30 |
+
generator: torch.Generator,
|
| 31 |
+
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
| 32 |
+
keys = torch.stack(
|
| 33 |
+
[torch.randperm(VOCAB_SIZE, generator=generator)[:slots] for _ in range(batch_size)]
|
| 34 |
+
)
|
| 35 |
+
values = torch.randint(
|
| 36 |
+
VOCAB_SIZE,
|
| 37 |
+
(batch_size, slots),
|
| 38 |
+
generator=generator,
|
| 39 |
+
)
|
| 40 |
+
query_positions = torch.randint(slots, (batch_size,), generator=generator)
|
| 41 |
+
rows = torch.arange(batch_size)
|
| 42 |
+
query = keys[rows, query_positions]
|
| 43 |
+
target = values[rows, query_positions]
|
| 44 |
+
return keys, values, query, target
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
@torch.inference_mode()
|
| 48 |
+
def evaluate(
|
| 49 |
+
model: torch.nn.Module,
|
| 50 |
+
*,
|
| 51 |
+
slots: int,
|
| 52 |
+
seed: int,
|
| 53 |
+
examples: int = 4_096,
|
| 54 |
+
) -> dict:
|
| 55 |
+
generator = torch.Generator().manual_seed(seed)
|
| 56 |
+
model.eval()
|
| 57 |
+
correct = 0
|
| 58 |
+
attention_mass = []
|
| 59 |
+
for start in range(0, examples, 256):
|
| 60 |
+
size = min(256, examples - start)
|
| 61 |
+
keys, values, query, target = sample_batch(size, slots, generator)
|
| 62 |
+
if isinstance(model, ContentAddressedMemory):
|
| 63 |
+
logits, attention = model(
|
| 64 |
+
keys,
|
| 65 |
+
values,
|
| 66 |
+
query,
|
| 67 |
+
return_attention=True,
|
| 68 |
+
)
|
| 69 |
+
match = keys.eq(query[:, None])
|
| 70 |
+
attention_mass.extend(attention[match].tolist())
|
| 71 |
+
else:
|
| 72 |
+
logits = model(keys, values, query)
|
| 73 |
+
correct += int(logits.argmax(1).eq(target).sum())
|
| 74 |
+
report = {"accuracy": correct / examples, "examples": examples}
|
| 75 |
+
if attention_mass:
|
| 76 |
+
report["mean_attention_on_correct_slot"] = float(np.mean(attention_mass))
|
| 77 |
+
return report
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
def train_one(
|
| 81 |
+
constructor: type[ContentAddressedMemory] | type[FixedStateGRU],
|
| 82 |
+
seed: int,
|
| 83 |
+
) -> torch.nn.Module:
|
| 84 |
+
torch.manual_seed(seed)
|
| 85 |
+
generator = torch.Generator().manual_seed(seed + 1)
|
| 86 |
+
model = constructor()
|
| 87 |
+
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-3, weight_decay=1e-5)
|
| 88 |
+
for step in range(1, STEPS + 1):
|
| 89 |
+
slots = int(
|
| 90 |
+
torch.randint(
|
| 91 |
+
TRAIN_SLOT_RANGE[0],
|
| 92 |
+
TRAIN_SLOT_RANGE[1] + 1,
|
| 93 |
+
(),
|
| 94 |
+
generator=generator,
|
| 95 |
+
)
|
| 96 |
+
)
|
| 97 |
+
keys, values, query, target = sample_batch(BATCH_SIZE, slots, generator)
|
| 98 |
+
loss = F.cross_entropy(model(keys, values, query), target)
|
| 99 |
+
optimizer.zero_grad(set_to_none=True)
|
| 100 |
+
loss.backward()
|
| 101 |
+
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
|
| 102 |
+
optimizer.step()
|
| 103 |
+
if step % 250 == 0:
|
| 104 |
+
trackio.log(
|
| 105 |
+
{
|
| 106 |
+
"training_step": step,
|
| 107 |
+
"variant": constructor.__name__,
|
| 108 |
+
"training_loss": float(loss.detach()),
|
| 109 |
+
}
|
| 110 |
+
)
|
| 111 |
+
return model
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
def write_dataset() -> None:
|
| 115 |
+
generator = torch.Generator().manual_seed(23_117)
|
| 116 |
+
keys, values, queries, targets = sample_batch(512, 32, generator)
|
| 117 |
+
lines = []
|
| 118 |
+
for index in range(len(keys)):
|
| 119 |
+
lines.append(
|
| 120 |
+
json.dumps(
|
| 121 |
+
{
|
| 122 |
+
"keys": keys[index].tolist(),
|
| 123 |
+
"values": values[index].tolist(),
|
| 124 |
+
"query": int(queries[index]),
|
| 125 |
+
"target": int(targets[index]),
|
| 126 |
+
}
|
| 127 |
+
)
|
| 128 |
+
)
|
| 129 |
+
DATA_DIR.mkdir(parents=True, exist_ok=True)
|
| 130 |
+
(DATA_DIR / "associative_recall_eval.jsonl").write_text(
|
| 131 |
+
"\n".join(lines) + "\n",
|
| 132 |
+
encoding="utf-8",
|
| 133 |
+
)
|
| 134 |
+
|
| 135 |
+
|
| 136 |
+
def main() -> None:
|
| 137 |
+
torch.set_num_threads(1)
|
| 138 |
+
trackio.init(
|
| 139 |
+
project="memory-tape-pocket",
|
| 140 |
+
name="content-addressing-vs-fixed-state-v1",
|
| 141 |
+
config={
|
| 142 |
+
"training_slots": list(TRAIN_SLOT_RANGE),
|
| 143 |
+
"steps": STEPS,
|
| 144 |
+
"seeds": SEEDS,
|
| 145 |
+
},
|
| 146 |
+
)
|
| 147 |
+
constructors = {
|
| 148 |
+
"memory": ContentAddressedMemory,
|
| 149 |
+
"gru": FixedStateGRU,
|
| 150 |
+
}
|
| 151 |
+
runs = {name: [] for name in constructors}
|
| 152 |
+
saved_models = {}
|
| 153 |
+
for seed in SEEDS:
|
| 154 |
+
for name, constructor in constructors.items():
|
| 155 |
+
model = train_one(constructor, seed)
|
| 156 |
+
run = {
|
| 157 |
+
"seed": seed,
|
| 158 |
+
"slots_8": evaluate(model, slots=8, seed=seed + 100),
|
| 159 |
+
"slots_16": evaluate(model, slots=16, seed=seed + 200),
|
| 160 |
+
"slots_32": evaluate(model, slots=32, seed=seed + 300),
|
| 161 |
+
}
|
| 162 |
+
runs[name].append(run)
|
| 163 |
+
if seed == SEEDS[0]:
|
| 164 |
+
saved_models[name] = model
|
| 165 |
+
results = {}
|
| 166 |
+
for name, model_runs in runs.items():
|
| 167 |
+
results[name] = {
|
| 168 |
+
"parameters": parameter_count(saved_models[name]),
|
| 169 |
+
"runs": model_runs,
|
| 170 |
+
"accuracy_mean": {
|
| 171 |
+
f"slots_{slots}": float(
|
| 172 |
+
np.mean(
|
| 173 |
+
[
|
| 174 |
+
run[f"slots_{slots}"]["accuracy"]
|
| 175 |
+
for run in model_runs
|
| 176 |
+
]
|
| 177 |
+
)
|
| 178 |
+
)
|
| 179 |
+
for slots in [8, 16, 32]
|
| 180 |
+
},
|
| 181 |
+
}
|
| 182 |
+
if name == "memory":
|
| 183 |
+
results[name]["correct_slot_attention_mean"] = {
|
| 184 |
+
f"slots_{slots}": float(
|
| 185 |
+
np.mean(
|
| 186 |
+
[
|
| 187 |
+
run[f"slots_{slots}"][
|
| 188 |
+
"mean_attention_on_correct_slot"
|
| 189 |
+
]
|
| 190 |
+
for run in model_runs
|
| 191 |
+
]
|
| 192 |
+
)
|
| 193 |
+
)
|
| 194 |
+
for slots in [8, 16, 32]
|
| 195 |
+
}
|
| 196 |
+
report = {
|
| 197 |
+
"experiment": "Differentiable content addressing versus fixed-state recall",
|
| 198 |
+
"training_slots": list(TRAIN_SLOT_RANGE),
|
| 199 |
+
"results": results,
|
| 200 |
+
}
|
| 201 |
+
ARTIFACT_DIR.mkdir(parents=True, exist_ok=True)
|
| 202 |
+
save_file(
|
| 203 |
+
saved_models["memory"].state_dict(),
|
| 204 |
+
ARTIFACT_DIR / "content_memory.safetensors",
|
| 205 |
+
)
|
| 206 |
+
save_file(
|
| 207 |
+
saved_models["gru"].state_dict(),
|
| 208 |
+
ARTIFACT_DIR / "fixed_gru.safetensors",
|
| 209 |
+
)
|
| 210 |
+
(ARTIFACT_DIR / "evaluation.json").write_text(
|
| 211 |
+
json.dumps(report, indent=2),
|
| 212 |
+
encoding="utf-8",
|
| 213 |
+
)
|
| 214 |
+
write_dataset()
|
| 215 |
+
trackio.log(
|
| 216 |
+
{
|
| 217 |
+
"memory_slots_32_mean": results["memory"]["accuracy_mean"]["slots_32"],
|
| 218 |
+
"gru_slots_32_mean": results["gru"]["accuracy_mean"]["slots_32"],
|
| 219 |
+
}
|
| 220 |
+
)
|
| 221 |
+
trackio.finish()
|
| 222 |
+
print(json.dumps(report, indent=2))
|
| 223 |
+
|
| 224 |
+
|
| 225 |
+
if __name__ == "__main__":
|
| 226 |
+
main()
|