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)
|