File size: 4,878 Bytes
b20ca9c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Train independent CRAI ensemble members, optionally under torchrun DDP."""

from pathlib import Path
import argparse
import json
import os
import random
import sys
import numpy as np
import torch
from torch import distributed as dist
from torch.nn.parallel import DistributedDataParallel
from torch.utils.data import DataLoader, Dataset, DistributedSampler
import yaml

ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))

from model.crai_climateextremes import CRAIClimateExtremes


class ClimateDataset(Dataset):
    def __init__(self, archive, limit):
        self.observed = torch.from_numpy(archive["observed"][:limit])
        self.valid = torch.from_numpy(archive["valid_mask"][:limit])
        self.target = torch.from_numpy(archive["target"][:limit])
        land = torch.from_numpy(archive["europe_mask"])[None, None]
        self.missing = land * (1.0 - self.valid)

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

    def __getitem__(self, index):
        return torch.cat((self.observed[index], self.valid[index]), 0), self.target[index], self.missing[index]


def load_config(path):
    with open(path, encoding="utf-8") as handle:
        return yaml.safe_load(handle)


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--config", type=Path, default=ROOT / "conf/config.yaml")
    parser.add_argument("--paper-model", action="store_true")
    args = parser.parse_args()
    config_path = args.config if args.config.is_absolute() else ROOT / args.config
    cfg = load_config(config_path)
    use_paper = args.paper_model or cfg["paper_model"]
    batch_size = cfg["paper_batch_size"] if use_paper else cfg["batch_size"]
    iterations = cfg["paper_iterations"] if use_paper else cfg["max_iterations"]
    members = cfg["paper_ensemble_members"] if use_paper else cfg["ensemble_members"]
    rank, world = int(os.getenv("RANK", 0)), int(os.getenv("WORLD_SIZE", 1))
    if world > 1:
        dist.init_process_group("nccl" if torch.cuda.is_available() else "gloo")
    device = torch.device(f"cuda:{int(os.getenv('LOCAL_RANK', 0))}" if torch.cuda.is_available() else "cpu")
    if device.type == "cuda":
        torch.cuda.set_device(device)
    archive = np.load(ROOT / cfg["data_path"])
    dataset = ClimateDataset(archive, cfg["num_samples"])
    checkpoint_path = ROOT / cfg["checkpoint_path"]
    if rank == 0:
        checkpoint_path.parent.mkdir(parents=True, exist_ok=True)
        (ROOT / "result/training").mkdir(parents=True, exist_ok=True)
    if world > 1:
        dist.barrier()
    records, member_states = [], []
    for member in range(members):
        seed = int(cfg["seed"]) + member
        random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
        sampler = DistributedSampler(dataset, shuffle=True, seed=seed) if world > 1 else None
        loader = DataLoader(dataset, batch_size=batch_size, shuffle=sampler is None, sampler=sampler)
        model = CRAIClimateExtremes(cfg["base_channels"]).to(device)
        if world > 1:
            model = DistributedDataParallel(model, device_ids=[device.index] if device.type == "cuda" else None)
        optimizer = torch.optim.Adam(model.parameters(), lr=cfg["learning_rate"])
        step, losses = 0, []
        for epoch in range(cfg["epochs"] if not use_paper else 10**9):
            if sampler is not None:
                sampler.set_epoch(epoch)
            for inputs, target, missing in loader:
                inputs, target, missing = inputs.to(device), target.to(device), missing.to(device)
                prediction = model(inputs)
                loss = (torch.abs(prediction - target) * missing).sum() / missing.sum().clamp_min(1)
                optimizer.zero_grad(); loss.backward(); optimizer.step()
                losses.append(float(loss.detach()))
                step += 1
                if step >= iterations:
                    break
            if step >= iterations:
                break
        raw_model = model.module if isinstance(model, DistributedDataParallel) else model
        if rank == 0:
            member_states.append({key: value.detach().cpu() for key, value in raw_model.state_dict().items()})
            records.append({"member": member, "iterations": step, "final_missing_mae": losses[-1]})
    if rank == 0:
        torch.save(
            {
                "format_version": "1.0",
                "model_config": {"base_channels": cfg["base_channels"]},
                "model": member_states,
            },
            checkpoint_path,
        )
        payload = {"paper_model": bool(use_paper), "world_size": world, "members": records}
        (ROOT / "result/training/metrics.json").write_text(json.dumps(payload, indent=2) + "\n")
        print(json.dumps(payload))
    if world > 1:
        dist.destroy_process_group()


if __name__ == "__main__":
    main()