File size: 4,146 Bytes
675a89a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 | from __future__ import annotations
import json
from pathlib import Path
import pandas as pd
import torch
import trackio
from data import multiscale_batch
from model import ClockworkRNN, MatchedGRU, PlainRNN, parameter_count
from safetensors.torch import save_file
from torch.nn import functional as F
PROJECT_DIR = Path(__file__).resolve().parent
ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "clockwork-rnn-pocket"
DATA_DIR = PROJECT_DIR / "data"
SEED = 2099
@torch.inference_mode()
def evaluate(model: torch.nn.Module, length: int, seed: int) -> dict:
sequence = multiscale_batch(2_000, length + 1, seed)
prediction = model(sequence[:, :-1])
target = sequence[:, 1:]
error = prediction - target
return {
"rmse": float(error.square().mean().sqrt()),
"mae": float(error.abs().mean()),
"length": length,
"examples": len(sequence),
}
def train_variant(name: str, model: torch.nn.Module) -> tuple[torch.nn.Module, dict]:
optimizer = torch.optim.AdamW(model.parameters(), lr=2e-3, weight_decay=1e-5)
best = float("inf")
best_state = None
best_step = 0
for step in range(1, 2_001):
sequence = multiscale_batch(96, 65, SEED + step)
prediction = model(sequence[:, :-1])
loss = F.mse_loss(prediction, sequence[:, 1:])
optimizer.zero_grad(set_to_none=True)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 5)
optimizer.step()
if step % 100 == 0:
validation = evaluate(model, 64, SEED + 20_000)
trackio.log(
{
"variant": name,
"training_step": step,
"training_mse": float(loss.detach()),
"validation_rmse": validation["rmse"],
}
)
if validation["rmse"] < best:
best = validation["rmse"]
best_step = step
best_state = {
key: value.detach().cpu().clone()
for key, value in model.state_dict().items()
}
assert best_state is not None
model.load_state_dict(best_state)
return model, {"best_step": best_step, "best_validation_rmse": best}
def main() -> None:
torch.manual_seed(SEED)
torch.set_num_threads(1)
models = {
"clockwork_rnn": ClockworkRNN(),
"plain_rnn": PlainRNN(),
"matched_gru": MatchedGRU(),
}
trackio.init(
project="clockwork-rnn-pocket",
name="multiscale-periodic-recurrence-v1",
config={
"training_length": 64,
"training_steps_per_variant": 2_000,
"parameters": {
name: parameter_count(model) for name, model in models.items()
},
},
)
results = {}
ARTIFACT_DIR.mkdir(parents=True, exist_ok=True)
DATA_DIR.mkdir(parents=True, exist_ok=True)
for name, model in models.items():
trained, training = train_variant(name, model)
results[name] = {
"parameters": parameter_count(trained),
"training": training,
"length_64": evaluate(trained, 64, SEED + 30_000),
"length_256_zero_shot": evaluate(trained, 256, SEED + 40_000),
}
save_file(trained.state_dict(), ARTIFACT_DIR / f"{name}.safetensors")
report = {
"experiment": "Modern Clockwork RNN multiscale forecasting retest",
"training_length": 64,
"results": results,
}
(ARTIFACT_DIR / "evaluation.json").write_text(
json.dumps(report, indent=2), encoding="utf-8"
)
pd.DataFrame(
[
{
"variant": name,
"parameters": result["parameters"],
"length_64_rmse": result["length_64"]["rmse"],
"length_256_rmse": result["length_256_zero_shot"]["rmse"],
}
for name, result in results.items()
]
).to_parquet(DATA_DIR / "benchmark_results.parquet", index=False)
trackio.finish()
print(json.dumps(report, indent=2))
if __name__ == "__main__":
main()
|