Download train.py from ARotting/memory-tape-pocket-lab: direct link, hf CLI and curl.
- Browser
- Download file 6.99 kB
-
https://huggingface.co/ARotting/memory-tape-pocket-lab/resolve/main/train.py
- Command line
-
hf download hf://ARotting/memory-tape-pocket-lab/train.py
-
curl -L -o train.py https://huggingface.co/ARotting/memory-tape-pocket-lab/resolve/main/train.py
6.99 kB
| from __future__ import annotations | |
| import json | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import trackio | |
| from model import ( | |
| VOCAB_SIZE, | |
| ContentAddressedMemory, | |
| FixedStateGRU, | |
| parameter_count, | |
| ) | |
| from safetensors.torch import save_file | |
| from torch.nn import functional as F | |
| PROJECT_DIR = Path(__file__).resolve().parent | |
| ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "memory-tape-pocket" | |
| DATA_DIR = PROJECT_DIR / "data" | |
| TRAIN_SLOT_RANGE = (2, 8) | |
| STEPS = 2_500 | |
| BATCH_SIZE = 256 | |
| SEEDS = [2281, 2287, 2293] | |
| def sample_batch( | |
| batch_size: int, | |
| slots: int, | |
| generator: torch.Generator, | |
| ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: | |
| keys = torch.stack( | |
| [torch.randperm(VOCAB_SIZE, generator=generator)[:slots] for _ in range(batch_size)] | |
| ) | |
| values = torch.randint( | |
| VOCAB_SIZE, | |
| (batch_size, slots), | |
| generator=generator, | |
| ) | |
| query_positions = torch.randint(slots, (batch_size,), generator=generator) | |
| rows = torch.arange(batch_size) | |
| query = keys[rows, query_positions] | |
| target = values[rows, query_positions] | |
| return keys, values, query, target | |
| def evaluate( | |
| model: torch.nn.Module, | |
| *, | |
| slots: int, | |
| seed: int, | |
| examples: int = 4_096, | |
| ) -> dict: | |
| generator = torch.Generator().manual_seed(seed) | |
| model.eval() | |
| correct = 0 | |
| attention_mass = [] | |
| for start in range(0, examples, 256): | |
| size = min(256, examples - start) | |
| keys, values, query, target = sample_batch(size, slots, generator) | |
| if isinstance(model, ContentAddressedMemory): | |
| logits, attention = model( | |
| keys, | |
| values, | |
| query, | |
| return_attention=True, | |
| ) | |
| match = keys.eq(query[:, None]) | |
| attention_mass.extend(attention[match].tolist()) | |
| else: | |
| logits = model(keys, values, query) | |
| correct += int(logits.argmax(1).eq(target).sum()) | |
| report = {"accuracy": correct / examples, "examples": examples} | |
| if attention_mass: | |
| report["mean_attention_on_correct_slot"] = float(np.mean(attention_mass)) | |
| return report | |
| def train_one( | |
| constructor: type[ContentAddressedMemory] | type[FixedStateGRU], | |
| seed: int, | |
| ) -> torch.nn.Module: | |
| torch.manual_seed(seed) | |
| generator = torch.Generator().manual_seed(seed + 1) | |
| model = constructor() | |
| optimizer = torch.optim.AdamW(model.parameters(), lr=3e-3, weight_decay=1e-5) | |
| for step in range(1, STEPS + 1): | |
| slots = int( | |
| torch.randint( | |
| TRAIN_SLOT_RANGE[0], | |
| TRAIN_SLOT_RANGE[1] + 1, | |
| (), | |
| generator=generator, | |
| ) | |
| ) | |
| keys, values, query, target = sample_batch(BATCH_SIZE, slots, generator) | |
| loss = F.cross_entropy(model(keys, values, query), target) | |
| optimizer.zero_grad(set_to_none=True) | |
| loss.backward() | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) | |
| optimizer.step() | |
| if step % 250 == 0: | |
| trackio.log( | |
| { | |
| "training_step": step, | |
| "variant": constructor.__name__, | |
| "training_loss": float(loss.detach()), | |
| } | |
| ) | |
| return model | |
| def write_dataset() -> None: | |
| generator = torch.Generator().manual_seed(23_117) | |
| keys, values, queries, targets = sample_batch(512, 32, generator) | |
| lines = [] | |
| for index in range(len(keys)): | |
| lines.append( | |
| json.dumps( | |
| { | |
| "keys": keys[index].tolist(), | |
| "values": values[index].tolist(), | |
| "query": int(queries[index]), | |
| "target": int(targets[index]), | |
| } | |
| ) | |
| ) | |
| DATA_DIR.mkdir(parents=True, exist_ok=True) | |
| (DATA_DIR / "associative_recall_eval.jsonl").write_text( | |
| "\n".join(lines) + "\n", | |
| encoding="utf-8", | |
| ) | |
| def main() -> None: | |
| torch.set_num_threads(1) | |
| trackio.init( | |
| project="memory-tape-pocket", | |
| name="content-addressing-vs-fixed-state-v1", | |
| config={ | |
| "training_slots": list(TRAIN_SLOT_RANGE), | |
| "steps": STEPS, | |
| "seeds": SEEDS, | |
| }, | |
| ) | |
| constructors = { | |
| "memory": ContentAddressedMemory, | |
| "gru": FixedStateGRU, | |
| } | |
| runs = {name: [] for name in constructors} | |
| saved_models = {} | |
| for seed in SEEDS: | |
| for name, constructor in constructors.items(): | |
| model = train_one(constructor, seed) | |
| run = { | |
| "seed": seed, | |
| "slots_8": evaluate(model, slots=8, seed=seed + 100), | |
| "slots_16": evaluate(model, slots=16, seed=seed + 200), | |
| "slots_32": evaluate(model, slots=32, seed=seed + 300), | |
| } | |
| runs[name].append(run) | |
| if seed == SEEDS[0]: | |
| saved_models[name] = model | |
| results = {} | |
| for name, model_runs in runs.items(): | |
| results[name] = { | |
| "parameters": parameter_count(saved_models[name]), | |
| "runs": model_runs, | |
| "accuracy_mean": { | |
| f"slots_{slots}": float( | |
| np.mean( | |
| [ | |
| run[f"slots_{slots}"]["accuracy"] | |
| for run in model_runs | |
| ] | |
| ) | |
| ) | |
| for slots in [8, 16, 32] | |
| }, | |
| } | |
| if name == "memory": | |
| results[name]["correct_slot_attention_mean"] = { | |
| f"slots_{slots}": float( | |
| np.mean( | |
| [ | |
| run[f"slots_{slots}"][ | |
| "mean_attention_on_correct_slot" | |
| ] | |
| for run in model_runs | |
| ] | |
| ) | |
| ) | |
| for slots in [8, 16, 32] | |
| } | |
| report = { | |
| "experiment": "Differentiable content addressing versus fixed-state recall", | |
| "training_slots": list(TRAIN_SLOT_RANGE), | |
| "results": results, | |
| } | |
| ARTIFACT_DIR.mkdir(parents=True, exist_ok=True) | |
| save_file( | |
| saved_models["memory"].state_dict(), | |
| ARTIFACT_DIR / "content_memory.safetensors", | |
| ) | |
| save_file( | |
| saved_models["gru"].state_dict(), | |
| ARTIFACT_DIR / "fixed_gru.safetensors", | |
| ) | |
| (ARTIFACT_DIR / "evaluation.json").write_text( | |
| json.dumps(report, indent=2), | |
| encoding="utf-8", | |
| ) | |
| write_dataset() | |
| trackio.log( | |
| { | |
| "memory_slots_32_mean": results["memory"]["accuracy_mean"]["slots_32"], | |
| "gru_slots_32_mean": results["gru"]["accuracy_mean"]["slots_32"], | |
| } | |
| ) | |
| trackio.finish() | |
| print(json.dumps(report, indent=2)) | |
| if __name__ == "__main__": | |
| main() | |