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)