File size: 4,824 Bytes
15aff58
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""DDP-capable FuXi-DA training with analysis and frozen-proxy forecast supervision."""
import argparse
import json
import math
import os
from pathlib import Path

import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel
from torch.utils.data import DataLoader, DistributedSampler

import sys
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
import yaml
from model.fuxi_da import CompactForecastProxy, FuXiDA, ProceduralTileDataset


def latitude_weighted_l1(prediction, target, latitude):
    weight = torch.cos(torch.deg2rad(latitude)).clamp_min(0)
    weight = weight * weight.shape[-1] / weight.sum(dim=-1, keepdim=True)
    return ((prediction - target).abs() * weight[:, None, :, None]).mean()


def learning_rate(step, warmup, total, peak):
    if step < warmup:
        return 1e-8 + (peak - 1e-8) * step / max(1, warmup)
    progress = (step - warmup) / max(1, total - warmup)
    return peak * 0.5 * (1.0 + math.cos(math.pi * progress))


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--config", default="conf/config.yaml")
    parser.add_argument("--iterations", type=int)
    parser.add_argument("--forecast-steps", type=int)
    parser.add_argument("--output-dir")
    args = parser.parse_args()
    cfg = yaml.safe_load((ROOT / args.config).read_text())
    total = args.iterations or cfg["train"]["iterations"]
    forecast_steps = args.forecast_steps or cfg["train"]["forecast_steps"]

    world_size = int(os.environ.get("WORLD_SIZE", "1"))
    local_rank = int(os.environ.get("LOCAL_RANK", "0"))
    use_cuda = torch.cuda.is_available() and torch.cuda.device_count() >= world_size and cfg["runtime"]["device"] != "cpu"
    if world_size > 1:
        dist.init_process_group("nccl" if use_cuda else "gloo")
    device = torch.device(f"cuda:{local_rank}" if use_cuda else "cpu")
    if use_cuda:
        torch.cuda.set_device(device)
    torch.manual_seed(cfg["seed"] + local_rank)

    dataset = ProceduralTileDataset(cfg["data"]["tile_ids"], max(2, total), cfg["data"]["tile_size"], cfg["data"]["missing_probability"])
    sampler = DistributedSampler(dataset, shuffle=True) if world_size > 1 else None
    loader = DataLoader(dataset, batch_size=cfg["train"]["batch_size"], sampler=sampler, shuffle=sampler is None, num_workers=0)
    model = FuXiDA(cfg["model"]["base_channels"]).to(device)
    proxy = CompactForecastProxy().to(device)
    if world_size > 1:
        model = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None)
    optimizer = torch.optim.AdamW(model.parameters(), lr=cfg["train"]["learning_rate"], betas=(0.9, 0.999), weight_decay=cfg["train"]["weight_decay"])

    model.train()
    iterator = iter(loader)
    for step in range(total):
        try:
            batch = next(iterator)
        except StopIteration:
            if sampler is not None:
                sampler.set_epoch(step)
            iterator = iter(loader)
            batch = next(iterator)
        background, obs, target = (batch[key].to(device) for key in ("background", "obs", "target"))
        latitude = batch["latitude"].to(device)
        analysis = model(background, obs)
        analysis_loss = latitude_weighted_l1(analysis, target, latitude)
        state, forecast_loss = analysis, analysis_loss.new_zeros(())
        for lead in range(forecast_steps):
            state = proxy(state)
            forecast_loss = forecast_loss + latitude_weighted_l1(state, batch["forecast_targets"][:, lead].to(device), latitude)
        loss = analysis_loss + forecast_loss / forecast_steps
        optimizer.zero_grad(set_to_none=True)
        loss.backward()
        optimizer.step()
        lr = learning_rate(step + 1, cfg["train"]["warmup_steps"], total, cfg["train"]["learning_rate"])
        for group in optimizer.param_groups:
            group["lr"] = lr
        if local_rank == 0 and (step == 0 or (step + 1) % 100 == 0 or step + 1 == total):
            print(json.dumps({"step": step + 1, "loss": loss.item(), "analysis_l1": analysis_loss.item(), "lr": lr}))

    if local_rank == 0:
        checkpoint = ROOT / cfg["paths"]["checkpoint"]; checkpoint.parent.mkdir(parents=True, exist_ok=True)
        bare_model = model.module if isinstance(model, DistributedDataParallel) else model
        torch.save({"model": bare_model.state_dict(), "model_config": cfg["model"], "format_version": cfg["data"]["format_version"]}, checkpoint)
        metrics = ROOT / cfg["paths"]["training_metrics"]; metrics.parent.mkdir(parents=True, exist_ok=True)
        metrics.write_text(json.dumps({"iterations": total, "final_loss": float(loss), "world_size": world_size}, indent=2))
    if world_size > 1:
        dist.destroy_process_group()


if __name__ == "__main__":
    main()