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