from __future__ import annotations import copy import json from pathlib import Path import numpy as np import torch import trackio from data import VOCAB_SIZE, generate_selective_memory from model import GRUControl, SelectiveSSM, parameter_count from safetensors.torch import load_file, save_file from torch import nn from torch.utils.data import DataLoader, TensorDataset PROJECT_DIR = Path(__file__).resolve().parent ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "micro-mamba" DATA_DIR = PROJECT_DIR / "data" def seed_everything(seed: int) -> None: np.random.seed(seed) torch.manual_seed(seed) torch.set_num_threads(1) def loader_from( dataset: tuple[np.ndarray, np.ndarray, np.ndarray], batch_size: int, shuffle: bool, seed: int, ) -> DataLoader: tokens, markers, targets = dataset return DataLoader( TensorDataset( torch.from_numpy(tokens), torch.from_numpy(markers), torch.from_numpy(targets), ), batch_size=batch_size, shuffle=shuffle, generator=torch.Generator().manual_seed(seed), ) @torch.inference_mode() def evaluate(model: nn.Module, loader: DataLoader) -> dict: model.eval() correct = 0 examples = 0 losses = [] criterion = nn.CrossEntropyLoss() for tokens, markers, targets in loader: logits = model(tokens, markers) losses.append(float(criterion(logits, targets))) correct += int((logits.argmax(1) == targets).sum()) examples += len(targets) return { "accuracy": correct / examples, "cross_entropy": float(np.mean(losses)), } def train_variant( name: str, model: nn.Module, train_loader: DataLoader, validation_loader: DataLoader, ) -> tuple[nn.Module, list[dict]]: optimizer = torch.optim.AdamW(model.parameters(), lr=2e-3, weight_decay=1e-4) criterion = nn.CrossEntropyLoss() best_state = copy.deepcopy(model.state_dict()) best_accuracy = -1.0 stale = 0 history = [] for epoch in range(1, 26): model.train() losses = [] for tokens, markers, targets in train_loader: logits = model(tokens, markers) loss = criterion(logits, targets) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() losses.append(float(loss.detach())) validation = evaluate(model, validation_loader) record = { "variant": name, "epoch": epoch, "training_loss": float(np.mean(losses)), "validation_accuracy": validation["accuracy"], } history.append(record) trackio.log(record) if validation["accuracy"] > best_accuracy + 1e-4: best_accuracy = validation["accuracy"] best_state = copy.deepcopy(model.state_dict()) stale = 0 else: stale += 1 if stale >= 8 and epoch >= 15: break model.load_state_dict(best_state) return model, history def main() -> None: seed_everything(2043) length = 48 train_data = generate_selective_memory(12_000, length, seed=2043) validation_data = generate_selective_memory(2_000, length, seed=3043) test_data = generate_selective_memory(4_000, length, seed=4043) long_test_data = generate_selective_memory(4_000, 96, seed=5043) train_loader = loader_from(train_data, 256, True, 2043) validation_loader = loader_from(validation_data, 512, False, 3043) test_loader = loader_from(test_data, 512, False, 4043) long_test_loader = loader_from(long_test_data, 512, False, 5043) variants = { "selective_ssm": SelectiveSSM(VOCAB_SIZE, selective=True), "fixed_ssm": SelectiveSSM(VOCAB_SIZE, selective=False), "gru": GRUControl(VOCAB_SIZE), } trackio.init( project="micro-mamba", name="selective-state-space-memory-v1", config={ "training_examples": len(train_data[0]), "sequence_length": length, "marked_items": 4, "variants": { name: parameter_count(model) for name, model in variants.items() }, }, ) histories = {} results = {} ARTIFACT_DIR.mkdir(parents=True, exist_ok=True) for name, model in variants.items(): checkpoint = ARTIFACT_DIR / f"{name}.safetensors" if checkpoint.exists(): model.load_state_dict(load_file(checkpoint)) trained, history = model, [] else: trained, history = train_variant( name, model, train_loader, validation_loader ) save_file(trained.state_dict(), checkpoint) histories[name] = history results[name] = { "parameters": parameter_count(trained), "training_epochs": 25, "epochs_in_current_run": len(history), "checkpoint_reused": not bool(history), "length_48": evaluate(trained, test_loader), "length_96_zero_shot": evaluate(trained, long_test_loader), } report = { "benchmark": "Selective ordinal memory", "training_examples": len(train_data[0]), "training_sequence_length": length, "test_examples_per_length": len(test_data[0]), "results": results, "training_history": histories, } (ARTIFACT_DIR / "evaluation.json").write_text( json.dumps(report, indent=2), encoding="utf-8" ) DATA_DIR.mkdir(parents=True, exist_ok=True) np.savez_compressed( DATA_DIR / "selective_memory_test.npz", tokens=test_data[0], markers=test_data[1], targets=test_data[2], ) trackio.log( { "selective_ssm_test_accuracy": results["selective_ssm"][ "length_48" ]["accuracy"], "fixed_ssm_test_accuracy": results["fixed_ssm"]["length_48"][ "accuracy" ], "gru_test_accuracy": results["gru"]["length_48"]["accuracy"], "selective_ssm_long_accuracy": results["selective_ssm"][ "length_96_zero_shot" ]["accuracy"], } ) trackio.finish() print(json.dumps(report, indent=2)) if __name__ == "__main__": main()