from __future__ import annotations import copy import json from pathlib import Path import numpy as np import torch import trackio from data import generate_adding_problem from model import SequenceRegressor, 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" / "lstm-time-capsule" 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 make_loader( dataset: tuple[np.ndarray, np.ndarray], *, shuffle: bool, seed: int, ) -> DataLoader: inputs, targets = dataset return DataLoader( TensorDataset(torch.from_numpy(inputs), torch.from_numpy(targets)), batch_size=256, shuffle=shuffle, generator=torch.Generator().manual_seed(seed), ) @torch.inference_mode() def evaluate(model: nn.Module, loader: DataLoader) -> dict: model.eval() predictions = [] targets = [] for inputs, batch_targets in loader: predictions.append(model(inputs).numpy()) targets.append(batch_targets.numpy()) prediction = np.concatenate(predictions) target = np.concatenate(targets) error = prediction - target return { "rmse": float(np.sqrt(np.mean(error**2))), "mae": float(np.abs(error).mean()), "within_0.1": float(np.mean(np.abs(error) <= 0.1)), } def train_variant( name: str, model: SequenceRegressor, train_loader: DataLoader, validation_loader: DataLoader, ) -> tuple[SequenceRegressor, list[dict]]: optimizer = torch.optim.AdamW(model.parameters(), lr=2e-3, weight_decay=1e-5) criterion = nn.MSELoss() best = copy.deepcopy(model.state_dict()) best_rmse = float("inf") stale = 0 history = [] for epoch in range(1, 61): model.train() losses = [] for inputs, targets in train_loader: prediction = model(inputs) loss = criterion(prediction, 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_mse": float(np.mean(losses)), "validation_rmse": validation["rmse"], } history.append(record) if epoch % 5 == 0: trackio.log(record) if validation["rmse"] < best_rmse - 1e-4: best_rmse = validation["rmse"] best = copy.deepcopy(model.state_dict()) stale = 0 else: stale += 1 if stale >= 15 and epoch >= 30: break model.load_state_dict(best) return model, history def main() -> None: seed_everything(2043) training_length = 100 train_data = generate_adding_problem(16_000, training_length, 2043) validation_data = generate_adding_problem(2_000, training_length, 3043) evaluation_data = { "length_100": generate_adding_problem(4_000, 100, 4043), "length_200_zero_shot": generate_adding_problem(4_000, 200, 5043), "length_400_zero_shot": generate_adding_problem(4_000, 400, 6043), } train_loader = make_loader(train_data, shuffle=True, seed=2043) validation_loader = make_loader(validation_data, shuffle=False, seed=3043) variants = { "vanilla_rnn": SequenceRegressor("rnn"), "lstm": SequenceRegressor("lstm"), "gru": SequenceRegressor("gru"), } trackio.init( project="lstm-time-capsule", name="long-lag-adding-v1", config={ "training_examples": len(train_data[0]), "training_length": training_length, "parameters": { name: parameter_count(model) for name, model in variants.items() }, }, ) ARTIFACT_DIR.mkdir(parents=True, exist_ok=True) results = {} histories = {} completed_epochs = {"vanilla_rnn": 65, "lstm": 100, "gru": 60} 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": completed_epochs.get(name, len(history)), "checkpoint_reused": not bool(history), **{ length: evaluate( trained, make_loader(dataset, shuffle=False, seed=7043) ) for length, dataset in evaluation_data.items() }, } report = { "benchmark": "Long-lag adding problem", "training_examples": len(train_data[0]), "training_length": training_length, "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 / "adding_problem_test.npz", inputs=evaluation_data["length_100"][0], targets=evaluation_data["length_100"][1], ) trackio.log( { "rnn_rmse": results["vanilla_rnn"]["length_100"]["rmse"], "lstm_rmse": results["lstm"]["length_100"]["rmse"], "gru_rmse": results["gru"]["length_100"]["rmse"], "lstm_length_400_rmse": results["lstm"]["length_400_zero_shot"][ "rmse" ], } ) trackio.finish() print(json.dumps(report, indent=2)) if __name__ == "__main__": main()