File size: 4,780 Bytes
03573b6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Train PrecipitationSRCNN with optional torch DistributedDataParallel."""

import json
import os
import random
import sys
from pathlib import Path

import numpy as np
import torch
import yaml
from torch.nn.parallel import DistributedDataParallel


ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.precipitationsrcnn import FORMAT_VERSION, MODEL_NAME, bilinear_input, build_model, processed_loss


def load_data(config):
    data = np.load(ROOT / config["data"]["root"] / "daily_precipitation.npz")
    expected = (216, 488)
    if str(data["format_version"]) != FORMAT_VERSION or tuple(data["target_grid"]) != expected:
        raise ValueError("data version or target grid mismatch")
    target = data["target_precipitation"]
    elevation = data["elevation"]
    if target.ndim != 4 or target.shape[1:] != (1, *expected) or elevation.shape != (1, 1, *expected):
        raise ValueError("invalid precipitation/elevation shape")
    return data


def main():
    config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
    seed = int(config["seed"])
    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
    distributed = int(os.environ.get("WORLD_SIZE", "1")) > 1
    local_rank = int(os.environ.get("LOCAL_RANK", "0"))
    requested = config["runtime"]["device"]
    use_cuda = requested == "cuda" and torch.cuda.is_available()
    if distributed:
        backend = "nccl" if use_cuda else "gloo"
        torch.distributed.init_process_group(backend=backend)
    rank = torch.distributed.get_rank() if distributed else 0
    world_size = torch.distributed.get_world_size() if distributed else 1
    device = torch.device(f"cuda:{local_rank}" if use_cuda else "cpu")
    if device.type == "cuda":
        torch.cuda.set_device(device)
    data = load_data(config)
    count = int(config["data"]["train_samples"])
    indices = list(range(rank, count, world_size))
    if not indices:
        raise ValueError("each DDP rank requires at least one training sample")
    model = build_model(config["model"]).to(device)
    train_model = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None) if distributed else model
    optimizer = torch.optim.Adam(train_model.parameters(), lr=float(config["training"]["learning_rate"]))
    losses = []
    elevation = torch.from_numpy(data["elevation"]).to(device)
    for _ in range(int(config["training"]["epochs"])):
        for index in indices:
            coarse = torch.from_numpy(data["coarse_precipitation"][index:index + 1]).to(device)
            target = torch.from_numpy(data["target_precipitation"][index:index + 1]).to(device)
            inputs = bilinear_input(coarse, elevation, tuple(config["data"]["target_grid"]))
            optimizer.zero_grad(set_to_none=True)
            prediction = train_model(inputs)
            loss = processed_loss(prediction, target, config["training"]["loss"],
                                  float(config["training"]["exponential_alpha"]),
                                  float(config["training"]["exponential_weight"]),
                                  float(config["training"]["quantile"]))
            if not torch.isfinite(loss):
                raise FloatingPointError("non-finite training loss")
            loss.backward(); optimizer.step(); losses.append(float(loss.detach()))
    if distributed:
        gathered = [None] * world_size
        torch.distributed.all_gather_object(gathered, losses)
        losses = [value for part in gathered for value in part]
    if rank == 0:
        checkpoint_path = ROOT / config["paths"]["checkpoint"]
        checkpoint_path.parent.mkdir(parents=True, exist_ok=True)
        checkpoint = {"format_version": FORMAT_VERSION, "model_name": MODEL_NAME,
                      "model_config": config["model"], "paper_model": config["paper_model"],
                      "target_grid": config["data"]["target_grid"], "model": model.state_dict(),
                      "loss": config["training"]["loss"], "seed": seed}
        torch.save(checkpoint, checkpoint_path)
        metrics_path = ROOT / config["paths"]["training_metrics"]
        metrics_path.parent.mkdir(parents=True, exist_ok=True)
        metrics = {"format_version": FORMAT_VERSION, "epochs": int(config["training"]["epochs"]),
                   "world_size": world_size, "loss": config["training"]["loss"],
                   "loss_history": losses, "final_loss": losses[-1]}
        metrics_path.write_text(json.dumps(metrics, indent=2) + "\n")
        print(f"checkpoint={checkpoint_path.relative_to(ROOT)} final_loss={losses[-1]:.6f}")
    if distributed:
        torch.distributed.barrier(); torch.distributed.destroy_process_group()


if __name__ == "__main__":
    main()