File size: 5,209 Bytes
2e913c2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Train four independent paper CNN-LSTM branches with RMSprop and MSE."""

import json
import os
import sys
from pathlib import Path

import numpy as np
import torch
import yaml
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.climatebench import ClimateBench, ClimateBenchBranch, TARGETS


class ClimateDataset(Dataset):
    def __init__(self, path: Path, config: dict):
        self.data = np.load(path)
        expected = config["data"]
        shape = (int(expected["time_steps"]), 4, int(expected["height"]), int(expected["width"]))
        target_shape = (4, int(expected["height"]), int(expected["width"]))
        if str(self.data["format_version"]) != expected["format_version"] or str(self.data["storage_layout"]) != "NTCHW":
            raise ValueError("incompatible ClimateBench NPZ format or storage layout")
        if self.data["inputs"].shape[1:] != shape or self.data["targets"].shape[1:] != target_shape:
            raise ValueError(f"expected inputs [N,{shape}] and targets [N,{target_shape}]")
        if self.data["inputs"].dtype != np.float32 or self.data["targets"].dtype != np.float32:
            raise TypeError("inputs and targets must be float32")
        if tuple(self.data["channel_names"].tolist()) != tuple(expected["channels"]):
            raise ValueError("forcing channels must be [co2_cumulative,ch4,so2,bc]")

    def __len__(self) -> int:
        return len(self.data["inputs"])

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


def device_from_config(config: dict, rank: int = 0) -> torch.device:
    requested = config["runtime"]["device"]
    if requested == "auto":
        return torch.device("cuda", rank) if torch.cuda.is_available() else torch.device("cpu")
    return torch.device(requested)


def main() -> None:
    config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
    np.random.seed(int(config["seed"]))
    torch.manual_seed(int(config["seed"]))
    distributed = int(os.environ.get("WORLD_SIZE", "1")) > 1
    local_rank = int(os.environ.get("LOCAL_RANK", "0"))
    if distributed:
        torch.distributed.init_process_group("nccl" if torch.cuda.is_available() else "gloo")
    rank = torch.distributed.get_rank() if distributed else 0
    device = device_from_config(config, local_rank)
    if device.type == "cuda":
        torch.cuda.set_device(device)
    dataset = ClimateDataset(ROOT / config["data"]["root"] / "train.npz", config)
    sampler = DistributedSampler(dataset, shuffle=True) if distributed else None
    loader = DataLoader(dataset, batch_size=int(config["train"]["batch_size"]), sampler=sampler,
                        shuffle=sampler is None, num_workers=int(config["train"]["num_workers"]))
    model = ClimateBench(int(config["data"]["height"]), int(config["data"]["width"])).to(device)
    counts = model.parameter_counts()
    expected_count = int(config["model"]["parameters_per_target"])
    if any(count != expected_count for count in counts.values()):
        raise RuntimeError(f"parameter count mismatch: {counts}")
    wrapped = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None) if distributed else model
    optimizer = torch.optim.RMSprop(wrapped.parameters(), lr=float(config["train"]["learning_rate"]))
    history = []
    for epoch in range(int(config["train"]["epochs"])):
        if sampler is not None:
            sampler.set_epoch(epoch)
        total, batches = 0.0, 0
        for inputs, targets in loader:
            prediction = wrapped(inputs.to(device))
            loss = torch.nn.functional.mse_loss(prediction, targets.to(device))
            optimizer.zero_grad(set_to_none=True)
            loss.backward()
            torch.nn.utils.clip_grad_norm_(wrapped.parameters(), float(config["train"]["gradient_clip_norm"]))
            optimizer.step()
            total += float(loss.detach())
            batches += 1
        history.append({"epoch": epoch + 1, "mse": total / max(batches, 1), "batches": batches})
    model = wrapped.module if distributed else wrapped
    if rank == 0:
        checkpoint_path = ROOT / config["paths"]["checkpoint"]
        metrics_path = ROOT / config["paths"]["training_metrics"]
        checkpoint_path.parent.mkdir(parents=True, exist_ok=True)
        metrics_path.parent.mkdir(parents=True, exist_ok=True)
        torch.save({"model": model.state_dict(), "parameter_counts": counts, "targets": TARGETS,
                    "format_version": config["data"]["format_version"], "storage_layout": "NTCHW",
                    "model_config": {"height": model.height, "width": model.width}}, checkpoint_path)
        metrics_path.write_text(json.dumps({"history": history, "parameter_counts": counts}, indent=2) + "\n")
        print(f"checkpoint={checkpoint_path.relative_to(ROOT)} batches={history[-1]['batches']} parameters={counts}")
    if distributed:
        torch.distributed.destroy_process_group()


if __name__ == "__main__":
    main()