File size: 4,017 Bytes
8791cfe
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
import argparse
import json
from pathlib import Path

import numpy as np
import torch
from torch.utils.data import DataLoader, TensorDataset

from model.nncam import build_model, fit_normalizer, normalize_input, scale_output


ROOT = Path(__file__).resolve().parents[1]


def main():
    parser = argparse.ArgumentParser(description="Train NNCAM on the prepared NPZ dataset.")
    parser.add_argument("--data", type=Path, default=ROOT / "data/nncam_fake.npz")
    parser.add_argument("--checkpoint", type=Path, default=ROOT / "result/checkpoints/nncam.pt")
    parser.add_argument("--metrics", type=Path, default=ROOT / "result/training/metrics.json")
    parser.add_argument("--epochs", type=int)
    parser.add_argument("--batch-size", type=int)
    parser.add_argument("--width", type=int)
    parser.add_argument("--depth", type=int)
    parser.add_argument("--paper-model", action="store_true", help="Explicitly use depth=9, width=256, epochs=18, batch_size=1024 (567361 parameters).")
    parser.add_argument("--lr", type=float, default=1e-3)
    parser.add_argument("--seed", type=int, default=42)
    args = parser.parse_args()
    if not args.data.is_file():
        raise FileNotFoundError(f"missing dataset {args.data}; run python scripts/fake_data.py first")
    defaults = {"depth": 9, "width": 256, "epochs": 18, "batch_size": 1024} if args.paper_model else {"depth": 4, "width": 32, "epochs": 3, "batch_size": 64}
    depth, width = args.depth or defaults["depth"], args.width or defaults["width"]
    epochs, batch_size = args.epochs or defaults["epochs"], args.batch_size or defaults["batch_size"]
    torch.manual_seed(args.seed)
    with np.load(args.data) as data:
        x, y = data["x"].astype(np.float32), data["y"].astype(np.float32)
    if x.ndim != 2 or x.shape[1] != 94 or y.shape != (x.shape[0], 65):
        raise ValueError(f"expected x=[N,94], y=[N,65], got {x.shape}, {y.shape}")
    input_mean, input_scale = fit_normalizer(x)
    scaled_y = scale_output(y)
    target_mean, target_scale = fit_normalizer(scaled_y)
    dataset = TensorDataset(torch.from_numpy(normalize_input(x, input_mean, input_scale)), torch.from_numpy(normalize_input(scaled_y, target_mean, target_scale)))
    loader = DataLoader(dataset, batch_size=batch_size, shuffle=True)
    model = build_model(width=width, depth=depth)
    optimizer = torch.optim.Adam(model.parameters(), lr=args.lr)
    scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.2)
    history = []
    for epoch in range(epochs):
        total = 0.0
        for xb, yb in loader:
            optimizer.zero_grad(set_to_none=True)
            loss = torch.nn.functional.mse_loss(model(xb), yb)
            loss.backward()
            optimizer.step()
            total += loss.item() * len(xb)
        history.append(total / len(dataset))
        scheduler.step()
        print(f"epoch={epoch + 1:02d} loss={history[-1]:.6f}")
    parameter_count = sum(parameter.numel() for parameter in model.parameters())
    checkpoint = {
        "format_version": 1,
        "model": model.state_dict(),
        "model_config": model.model_config,
        "normalization": {"input_mean": torch.from_numpy(input_mean), "input_scale": torch.from_numpy(input_scale), "target_mean": torch.from_numpy(target_mean), "target_scale": torch.from_numpy(target_scale)},
        "training": {"epochs": epochs, "batch_size": batch_size, "learning_rate": args.lr, "paper_model": args.paper_model, "parameters": parameter_count},
    }
    args.checkpoint.parent.mkdir(parents=True, exist_ok=True)
    args.metrics.parent.mkdir(parents=True, exist_ok=True)
    torch.save(checkpoint, args.checkpoint)
    args.metrics.write_text(json.dumps({"loss": history, "final_loss": history[-1], "parameters": parameter_count, "model_config": model.model_config}, indent=2), encoding="utf-8")
    print(f"saved {args.checkpoint}; parameters={parameter_count}; paper_model={args.paper_model}")


if __name__ == "__main__":
    main()