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