| 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() |
|
|