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