ARotting's picture
Publish Event-gated hierarchical recurrent memory retest
25021c2 verified
Raw
History Blame Contribute Delete
6.61 kB
from __future__ import annotations
import json
from pathlib import Path
import numpy as np
import torch
import trackio
from model import (
CONTROL_STATES,
HistoryCompressor,
PlainRNN,
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" / "history-compressor-pocket"
DATA_DIR = PROJECT_DIR / "data"
SEEDS = [2371, 2377, 2381]
TRAIN_GAPS = [8, 16, 24, 32]
STEPS = 500
BATCH_SIZE = 64
def next_state(state: torch.Tensor) -> torch.Tensor:
return (state * 5 + 1).remainder(CONTROL_STATES)
def sample_batch(
batch_size: int,
*,
blocks: int,
gap: int,
generator: torch.Generator,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
sequence_length = blocks * gap
tokens = torch.randint(
CONTROL_STATES,
CONTROL_STATES * 2,
(batch_size, sequence_length),
generator=generator,
)
targets = torch.zeros(batch_size, sequence_length, dtype=torch.long)
mask = torch.zeros(batch_size, sequence_length, dtype=torch.bool)
state = torch.randint(
CONTROL_STATES,
(batch_size,),
generator=generator,
)
for block in range(blocks):
start = block * gap
tokens[:, start] = state
state = next_state(state)
boundary = start + gap - 1
targets[:, boundary] = state
mask[:, boundary] = True
return tokens, targets, mask
@torch.inference_mode()
def evaluate(
model: torch.nn.Module,
*,
gap: int,
seed: int,
examples: int = 1_024,
) -> dict:
generator = torch.Generator().manual_seed(seed)
model.eval()
correct = 0
total = 0
for start in range(0, examples, 256):
size = min(256, examples - start)
tokens, targets, mask = sample_batch(
size,
blocks=6,
gap=gap,
generator=generator,
)
predictions = model(tokens).argmax(-1)
correct += int(predictions[mask].eq(targets[mask]).sum())
total += int(mask.sum())
return {"boundary_accuracy": correct / total, "decisions": total}
def train_one(
constructor: type[HistoryCompressor] | type[PlainRNN],
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):
gap = TRAIN_GAPS[
int(torch.randint(len(TRAIN_GAPS), (), generator=generator))
]
tokens, targets, mask = sample_batch(
BATCH_SIZE,
blocks=4,
gap=gap,
generator=generator,
)
logits = model(tokens)
loss = F.cross_entropy(logits[mask], targets[mask])
optimizer.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
if step % 200 == 0:
trackio.log(
{
"training_step": step,
"variant": constructor.__name__,
"training_gap": gap,
"training_loss": float(loss.detach()),
}
)
return model
def write_dataset() -> None:
generator = torch.Generator().manual_seed(23_903)
tokens, targets, mask = sample_batch(
256,
blocks=6,
gap=128,
generator=generator,
)
records = []
for index in range(len(tokens)):
records.append(
json.dumps(
{
"tokens": tokens[index].tolist(),
"boundary_targets": targets[index][mask[index]].tolist(),
"gap": 128,
}
)
)
DATA_DIR.mkdir(parents=True, exist_ok=True)
(DATA_DIR / "gap_128_eval.jsonl").write_text(
"\n".join(records) + "\n",
encoding="utf-8",
)
def main() -> None:
torch.set_num_threads(1)
assert parameter_count(HistoryCompressor()) == parameter_count(PlainRNN()) == 3_464
trackio.init(
project="history-compressor-pocket",
name="event-gated-versus-plain-recurrence-v1",
config={
"parameters_per_model": 3_464,
"training_gaps": TRAIN_GAPS,
"seeds": SEEDS,
},
)
constructors = {
"history_compressor": HistoryCompressor,
"plain_rnn": PlainRNN,
}
runs = {name: [] for name in constructors}
saved = {}
for seed in SEEDS:
for name, constructor in constructors.items():
model = train_one(constructor, seed)
run = {
"seed": seed,
**{
f"gap_{gap}": evaluate(
model,
gap=gap,
seed=seed + gap,
)
for gap in [32, 64, 128]
},
}
runs[name].append(run)
if seed == SEEDS[0]:
saved[name] = model.state_dict()
results = {}
for name, model_runs in runs.items():
results[name] = {
"parameters": 3_464,
"runs": model_runs,
"accuracy_mean": {
f"gap_{gap}": float(
np.mean(
[
run[f"gap_{gap}"]["boundary_accuracy"]
for run in model_runs
]
)
)
for gap in [32, 64, 128]
},
}
report = {
"experiment": "Event-gated history compression versus plain recurrence",
"maximum_training_gap": max(TRAIN_GAPS),
"results": results,
}
ARTIFACT_DIR.mkdir(parents=True, exist_ok=True)
save_file(saved["history_compressor"], ARTIFACT_DIR / "compressor.safetensors")
save_file(saved["plain_rnn"], ARTIFACT_DIR / "plain_rnn.safetensors")
(ARTIFACT_DIR / "evaluation.json").write_text(
json.dumps(report, indent=2),
encoding="utf-8",
)
write_dataset()
trackio.log(
{
"compressor_gap_128": results["history_compressor"]["accuracy_mean"][
"gap_128"
],
"plain_gap_128": results["plain_rnn"]["accuracy_mean"]["gap_128"],
}
)
trackio.finish()
print(json.dumps(report, indent=2))
if __name__ == "__main__":
main()