ARotting's picture
Publish Parameter-matched RNN, LSTM, and GRU long-lag retest
ddaaeb2 verified
Raw
History Blame Contribute Delete
6.19 kB
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()