Spaces:
Running
Running
| 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), | |
| ) | |
| 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() | |