File size: 3,090 Bytes
380b161
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from pathlib import Path
import sys

import numpy as np
import torch
from torch.nn.parallel import DistributedDataParallel
from torch.utils.data import DataLoader, Dataset, DistributedSampler

ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.ace2 import build_model, hard_correct, init_distributed, load_config, seed_all


class Windows(Dataset):
    def __init__(self, path):
        data = np.load(path)
        self.state = data["state"]
        self.forcing = data["forcing"]
        self.windows = [(n, t) for n in range(len(self.state)) for t in range(self.state.shape[1] - 2)]

    def __len__(self):
        return len(self.windows)

    def __getitem__(self, index):
        n, t = self.windows[index]
        return tuple(torch.from_numpy(x.astype(np.float32)) for x in (
            self.state[n, t], self.state[n, t + 1], self.state[n, t + 2],
            self.forcing[n, t + 1], self.forcing[n, t + 2]))


def main():
    cfg = load_config(ROOT)
    seed_all(cfg["seed"])
    distributed, rank, device = init_distributed()
    dataset = Windows(ROOT / cfg["data"]["path"])
    sampler = DistributedSampler(dataset, shuffle=True) if distributed else None
    loader = DataLoader(dataset, batch_size=cfg["train"]["batch_size"], sampler=sampler,
                        shuffle=sampler is None, num_workers=0)
    model = build_model(cfg).to(device)
    if distributed:
        model = DistributedDataParallel(model, device_ids=[device.index] if device.type == "cuda" else None)
    optimizer = torch.optim.Adam(model.parameters(), lr=cfg["train"]["learning_rate"])
    for epoch in range(cfg["train"]["epochs"]):
        if sampler:
            sampler.set_epoch(epoch)
        total = 0.0
        for x0, y1, y2, f1, f2 in loader:
            x0, y1, y2, f1, f2 = (x.to(device) for x in (x0, y1, y2, f1, f2))
            p1 = hard_correct(x0, model(x0, f1))
            p2 = hard_correct(p1, model(p1, f2))
            loss = torch.mean((p1 - y1) ** 2) + cfg["train"]["two_step_weight"] * torch.mean((p2 - y2) ** 2)
            optimizer.zero_grad(set_to_none=True)
            loss.backward()
            optimizer.step()
            total += loss.item()
        if rank == 0:
            print(f"epoch={epoch + 1} two_step_loss={total / len(loader):.7f}")
    if rank == 0:
        path = ROOT / cfg["train"]["checkpoint"]
        path.parent.mkdir(parents=True, exist_ok=True)
        raw_model = model.module if distributed else model
        torch.save({"model": raw_model.state_dict(), "model_config": cfg["model"], "format_version": cfg["data"]["format_version"]}, path)
        metrics = ROOT / "result/training/metrics.json"
        metrics.parent.mkdir(parents=True, exist_ok=True)
        metrics.write_text(__import__("json").dumps({"history": [{"epoch": cfg["train"]["epochs"], "loss": total / len(loader)}], "world_size": int(__import__("os").environ.get("WORLD_SIZE", "1"))}, indent=2))
        print(f"saved {path}")
    if distributed:
        torch.distributed.destroy_process_group()


if __name__ == "__main__":
    main()