File size: 6,823 Bytes
8bfc737
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
from __future__ import annotations

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

import numpy as np
import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel

if __package__ in (None, ""):
    sys.path.insert(0, str(Path(__file__).resolve().parents[1]))

from model.earthformer import Earthformer
from script.data_loader import make_loader
from script.metrics import metric_sums, metric_sums_light, metrics_from_light_sums, metrics_from_sums, mse
from script.utils import DEFAULT_CONFIG, atomic_torch_save, clean_state_dict, load_checkpoint_payload, load_config, resolve_device


def seed_everything(seed: int) -> None:
    random.seed(seed)
    np.random.seed(seed)
    torch.manual_seed(seed)
    if torch.cuda.is_available():
        torch.cuda.manual_seed_all(seed)


def validate(
    model: torch.nn.Module,
    loader,
    device: torch.device,
    max_batches: int | None = None,
    full_metrics: bool = False,
) -> dict[str, float]:
    """Run validation with only cheap MSE/MAE by default.

    The training loss is already the per-sample MSE, so recomputing MSE during
    training is redundant monitoring. SSIM and per-threshold CSI are expensive
    on full-resolution 384x384 data, so they are computed only when
    `train.compute_full_metrics` is enabled; the authoritative full evaluation
    lives in script/result.py.
    """
    model.eval()
    sums = torch.zeros(22 if full_metrics else 3, dtype=torch.float64, device=device)
    with torch.no_grad():
        for batch_index, (inputs, targets) in enumerate(loader):
            predictions = model(inputs.to(device, non_blocking=True)).clamp(0.0, 1.0)
            targets_device = targets.to(device, non_blocking=True)
            if full_metrics:
                sums += metric_sums(predictions, targets_device)
            else:
                sums += metric_sums_light(predictions, targets_device)
            if max_batches is not None and batch_index + 1 >= max_batches:
                break
    if dist.is_initialized():
        dist.all_reduce(sums, op=dist.ReduceOp.SUM)
    count_index = 3 if full_metrics else 2
    if sums[count_index].item() == 0:
        raise ValueError("validation loader is empty")
    return (metrics_from_sums if full_metrics else metrics_from_light_sums)(sums.cpu())


def run_training(config: dict, requested_device: str = "auto", resume: str | None = None) -> tuple[Path, dict[str, float]]:
    world_size = int(os.environ.get("WORLD_SIZE", "1"))
    rank = int(os.environ.get("RANK", "0"))
    local_rank = int(os.environ.get("LOCAL_RANK", "0"))
    distributed = world_size > 1
    device = resolve_device(requested_device, local_rank)
    backend = config["distributed"].get("backend", "auto")
    if backend == "auto":
        backend = "nccl" if device.type == "cuda" else "gloo"
    if distributed:
        dist.init_process_group(backend=backend, rank=rank, world_size=world_size)
    try:
        seed_everything(int(config["train"]["seed"]))
        model = Earthformer(config).to(device)
        optimizer = torch.optim.AdamW(
            model.parameters(),
            lr=float(config["train"]["learning_rate"]),
            weight_decay=float(config["train"]["weight_decay"]),
        )
        start_epoch, step = 0, 0
        if resume:
            payload = load_checkpoint_payload(resume, device)
            model.load_state_dict(clean_state_dict(payload["model"]))
            if "optimizer" in payload:
                optimizer.load_state_dict(payload["optimizer"])
            start_epoch = int(payload.get("epoch", -1)) + 1
            step = int(payload.get("step", 0))
        if distributed:
            model = DistributedDataParallel(
                model,
                device_ids=[local_rank] if device.type == "cuda" else None,
                find_unused_parameters=True,
            )
        train_loader, train_sampler = make_loader(config, "train", distributed, rank, world_size)
        val_loader, _ = make_loader(config, "val", distributed, rank, world_size, shuffle=False)
        full_metrics = bool(config["train"].get("compute_full_metrics", False))
        last_epoch = max(start_epoch - 1, 0)
        metrics: dict[str, float] = {}
        for epoch in range(start_epoch, int(config["train"]["epochs"])):
            last_epoch = epoch
            if train_sampler is not None:
                train_sampler.set_epoch(epoch)
            model.train()
            epoch_loss_sum, epoch_steps = 0.0, 0
            for inputs, targets in train_loader:
                optimizer.zero_grad(set_to_none=True)
                loss = mse(model(inputs.to(device, non_blocking=True)), targets.to(device, non_blocking=True))
                if not torch.isfinite(loss):
                    raise FloatingPointError("training loss is not finite")
                loss.backward()
                optimizer.step()
                epoch_loss_sum += float(loss.detach().item())
                epoch_steps += 1
                step += 1
            if epoch_steps == 0:
                break
            metrics = validate(model, val_loader, device, int(config["train"].get("validation_steps", 1)), full_metrics)
            if rank == 0:
                epoch_line = {"epoch": epoch, "train_loss": epoch_loss_sum / epoch_steps, "validation": metrics}
                print(json.dumps(epoch_line))
        checkpoint = Path(config["train"]["output_dir"]) / "earthformer.pt"
        if rank == 0:
            raw_model = model.module if isinstance(model, DistributedDataParallel) else model
            atomic_torch_save(
                {
                    "model": raw_model.state_dict(),
                    "optimizer": optimizer.state_dict(),
                    "config": config,
                    "metrics": metrics,
                    "epoch": last_epoch,
                    "step": step,
                    "world_size": world_size,
                },
                checkpoint,
            )
            print(json.dumps({"checkpoint": str(checkpoint), "step": step, "world_size": world_size, "metrics": metrics}, indent=2))
        if distributed:
            dist.barrier()
        return checkpoint, metrics
    finally:
        if dist.is_initialized():
            dist.destroy_process_group()


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="Train Earthformer with one process or DDP")
    parser.add_argument("--config", default=str(DEFAULT_CONFIG))
    parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
    parser.add_argument("--resume")
    return parser.parse_args()


if __name__ == "__main__":
    args = parse_args()
    run_training(load_config(args.config), args.device, args.resume)