File size: 5,800 Bytes
30f7852
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Train UnetDif with single-process or torchrun DDP execution."""

import json
import os
import sys
from pathlib import Path

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


ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.unetdif import UnetDif, loss_components


class RainfallDataset(Dataset):
    def __init__(self, path: Path, config):
        archive = np.load(path, allow_pickle=False)
        expected = (int(config["time_steps"]), int(config["channels"]), int(config["height"]), int(config["width"]))
        if archive["inputs"].shape[1:] != expected or archive["targets"].shape[1:] != (expected[0], expected[2], expected[3]):
            raise ValueError(f"expected day shapes {expected} and {(expected[0], expected[2], expected[3])}")
        if len(np.unique(archive["group_id"])) != len(archive["group_id"]):
            raise ValueError("each rain day must belong to one split group")
        if archive["lead_hours"].tolist() != config["lead_hours"]:
            raise ValueError("NPZ lead_hours do not match the configured eight 3-hour steps")
        self.inputs = archive["inputs"].astype(np.float32, copy=False).reshape(-1, *expected[1:])
        self.targets = archive["targets"].astype(np.float32, copy=False).reshape(-1, expected[2], expected[3])

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

    def __getitem__(self, index):
        return torch.from_numpy(self.inputs[index]), torch.from_numpy(self.targets[index])


def get_device(local_rank=0, local_world_size=1):
    use_cuda = torch.cuda.is_available() and torch.cuda.device_count() >= local_world_size
    return torch.device("cuda", local_rank) if use_cuda else torch.device("cpu")


def reduce_values(values, device, distributed):
    tensor = torch.tensor(values, device=device)
    if distributed:
        torch.distributed.all_reduce(tensor)
    return tensor.cpu().tolist()


def main():
    config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
    torch.manual_seed(int(config["seed"]))
    distributed = int(os.environ.get("WORLD_SIZE", "1")) > 1
    local_rank = int(os.environ.get("LOCAL_RANK", "0"))
    local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1"))
    device = get_device(local_rank, local_world_size)
    if distributed:
        torch.distributed.init_process_group("nccl" if device.type == "cuda" else "gloo")
    rank = torch.distributed.get_rank() if distributed else 0
    dataset = RainfallDataset(ROOT / config["data"]["root"] / "train.npz", config["data"])
    if distributed and len(dataset) < torch.distributed.get_world_size():
        raise ValueError("DDP requires at least one training item per process")
    sampler = DistributedSampler(dataset) if distributed else None
    loader = DataLoader(dataset, batch_size=int(config["train"]["batch_size"]), shuffle=sampler is None,
                        sampler=sampler, num_workers=int(config["train"]["num_workers"]))
    model = UnetDif(**config["model"]).to(device)
    if distributed:
        model = DDP(model, device_ids=[local_rank] if device.type == "cuda" else None)
    optimizer = torch.optim.Adam(model.parameters(), lr=float(config["train"]["learning_rate"]))
    names = ["dry_focal", "false_alarm_focal", "positive_mse", "negative_mse", "dry_mae", "all_mae"]
    normalizers = {name: 1.0 for name in names}
    history = []
    for epoch in range(int(config["train"]["epochs"])):
        if sampler is not None:
            sampler.set_epoch(epoch)
        model.train(); sums = {name: 0.0 for name in names}; batches = 0
        for inputs, target in loader:
            inputs, target = inputs.to(device), target.to(device)
            heads = model(inputs)
            raw = loss_components(heads, inputs[:, int(config["data"]["precipitation_channel"])], target,
                                  float(config["data"]["rain_threshold_mm_3h"]),
                                  float(config["loss"]["focal_alpha"]), float(config["loss"]["focal_gamma"]))
            loss = sum(raw[name] * normalizers[name] for name in names)
            optimizer.zero_grad(set_to_none=True); loss.backward(); optimizer.step()
            for name in names:
                sums[name] += float(raw[name].detach())
            batches += 1
        reduced = reduce_values([sums[name] for name in names] + [batches], device, distributed)
        total_batches = max(reduced[-1], 1)
        averages = {name: reduced[index] / total_batches for index, name in enumerate(names)}
        if epoch == 0:
            normalizers = {name: 1.0 / max(value, 1e-6) for name, value in averages.items()}
        record = {"epoch": epoch + 1, "components": averages, "normalizers": normalizers}
        history.append(record)
        if rank == 0:
            print(json.dumps(record))
            checkpoint = ROOT / config["paths"]["checkpoint"]
            checkpoint.parent.mkdir(parents=True, exist_ok=True)
            state = model.module.state_dict() if distributed else model.state_dict()
            torch.save({"epoch": epoch + 1, "model": state, "optimizer": optimizer.state_dict(),
                         "model_config": config["model"], "loss_normalizers": normalizers,
                        "loss_components": names,
                         "format_version": config["data"]["format_version"]}, checkpoint)
    if rank == 0:
        metrics = ROOT / config["paths"]["training_metrics"]
        metrics.parent.mkdir(parents=True, exist_ok=True)
        metrics.write_text(json.dumps({"history": history}, indent=2) + "\n")
    if distributed:
        torch.distributed.destroy_process_group()


if __name__ == "__main__":
    main()