File size: 5,018 Bytes
950fc23
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Train independently initialized DD members with optional DDP."""

import copy
import json
import os
import sys
from pathlib import Path

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


ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.precipdd import DATA_FORMAT_VERSION, PrecipDD, load_config, seed_all, validate_archive


def main():
    config = load_config(ROOT / "conf/config.yaml")
    runtime, settings = config["runtime"], config["training"]
    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"))
    use_cuda = (torch.cuda.is_available() and runtime["device"] != "cpu"
                and torch.cuda.device_count() >= local_world_size)
    if distributed:
        torch.distributed.init_process_group(runtime["ddp_backend_gpu"] if use_cuda else runtime["ddp_backend_cpu"])
    rank = torch.distributed.get_rank() if distributed else 0
    device = torch.device(f"cuda:{local_rank}" if use_cuda else "cpu")
    if use_cuda:
        torch.cuda.set_device(device)
    data = np.load(ROOT / config["paths"]["data"])
    validate_archive(data)
    train_mask, val_mask = data["split"] == 0, data["split"] == 1
    dataset = TensorDataset(torch.from_numpy(data["precipitation"][train_mask]).float(), torch.from_numpy(data["agmt"][train_mask]).float())
    val_x = torch.from_numpy(data["precipitation"][val_mask]).float().to(device)
    val_y = torch.from_numpy(data["agmt"][val_mask]).float().to(device)
    ensemble_states, histories = [], []
    for member in range(settings["ensemble_members"]):
        seed_all(config["project"]["seed"] + member)
        sampler = DistributedSampler(dataset, shuffle=True, seed=member) if distributed else None
        loader = DataLoader(dataset, batch_size=settings["batch_size"], shuffle=sampler is None, sampler=sampler,
                            num_workers=settings["num_workers"])
        model = PrecipDD(config["model"]["filters"], config["model"]["dense_units"]).to(device)
        wrapped = DistributedDataParallel(model, device_ids=[local_rank] if use_cuda else None) if distributed else model
        optimizer = torch.optim.Adam(wrapped.parameters(), lr=settings["learning_rate"], weight_decay=settings["l2_weight_decay"])
        history, best_loss, best_state = [], float("inf"), None
        for epoch in range(settings["epochs"]):
            if sampler is not None:
                sampler.set_epoch(member * settings["epochs"] + epoch)
            wrapped.train()
            total, count = 0.0, 0
            for fields, target in loader:
                fields, target = fields.to(device), target.to(device)
                optimizer.zero_grad(set_to_none=True)
                loss = nn.functional.l1_loss(wrapped(fields), target)
                loss.backward()
                optimizer.step()
                total += float(loss.detach()) * len(fields)
                count += len(fields)
            summary = torch.tensor([total, count], dtype=torch.float64, device=device)
            if distributed:
                torch.distributed.all_reduce(summary)
            wrapped.eval()
            with torch.no_grad():
                val_loss = float(nn.functional.l1_loss(model(val_x), val_y))
            history.append({"epoch": epoch + 1, "train_mae": float(summary[0] / summary[1]), "validation_mae": val_loss})
            if val_loss < best_loss:
                best_loss = val_loss
                best_state = {key: value.detach().cpu().clone() for key, value in model.state_dict().items()}
        ensemble_states.append(best_state)
        histories.append(history)
    if rank == 0:
        checkpoint_path = ROOT / config["paths"]["checkpoint"]
        checkpoint_path.parent.mkdir(parents=True, exist_ok=True)
        torch.save({"format_version": DATA_FORMAT_VERSION, "ensemble_states": ensemble_states,
                    "model_config": {"filters": config["model"]["filters"], "dense_units": config["model"]["dense_units"],
                                     "input_shape": [1, 55, 160], "feature_shape": [16, 14, 40], "flattened_features": 8960},
                    "training_config": settings, "histories": histories, "world_size": torch.distributed.get_world_size() if distributed else 1}, checkpoint_path)
        metrics_path = ROOT / config["paths"]["training_metrics"]
        metrics_path.parent.mkdir(parents=True, exist_ok=True)
        metrics_path.write_text(json.dumps({"ensemble_members": len(histories), "history": histories}, indent=2) + "\n", encoding="utf-8")
        print(f"checkpoint={checkpoint_path.relative_to(ROOT)} members={len(ensemble_states)}")
    if distributed:
        torch.distributed.barrier()
        torch.distributed.destroy_process_group()


if __name__ == "__main__":
    main()