"""Train PrecipitationSRCNN with optional torch DistributedDataParallel.""" import json import os import random import sys from pathlib import Path import numpy as np import torch import yaml from torch.nn.parallel import DistributedDataParallel ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) from model.precipitationsrcnn import FORMAT_VERSION, MODEL_NAME, bilinear_input, build_model, processed_loss def load_data(config): data = np.load(ROOT / config["data"]["root"] / "daily_precipitation.npz") expected = (216, 488) if str(data["format_version"]) != FORMAT_VERSION or tuple(data["target_grid"]) != expected: raise ValueError("data version or target grid mismatch") target = data["target_precipitation"] elevation = data["elevation"] if target.ndim != 4 or target.shape[1:] != (1, *expected) or elevation.shape != (1, 1, *expected): raise ValueError("invalid precipitation/elevation shape") return data def main(): config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) seed = int(config["seed"]) random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) distributed = int(os.environ.get("WORLD_SIZE", "1")) > 1 local_rank = int(os.environ.get("LOCAL_RANK", "0")) requested = config["runtime"]["device"] use_cuda = requested == "cuda" and torch.cuda.is_available() if distributed: backend = "nccl" if use_cuda else "gloo" torch.distributed.init_process_group(backend=backend) rank = torch.distributed.get_rank() if distributed else 0 world_size = torch.distributed.get_world_size() if distributed else 1 device = torch.device(f"cuda:{local_rank}" if use_cuda else "cpu") if device.type == "cuda": torch.cuda.set_device(device) data = load_data(config) count = int(config["data"]["train_samples"]) indices = list(range(rank, count, world_size)) if not indices: raise ValueError("each DDP rank requires at least one training sample") model = build_model(config["model"]).to(device) train_model = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None) if distributed else model optimizer = torch.optim.Adam(train_model.parameters(), lr=float(config["training"]["learning_rate"])) losses = [] elevation = torch.from_numpy(data["elevation"]).to(device) for _ in range(int(config["training"]["epochs"])): for index in indices: coarse = torch.from_numpy(data["coarse_precipitation"][index:index + 1]).to(device) target = torch.from_numpy(data["target_precipitation"][index:index + 1]).to(device) inputs = bilinear_input(coarse, elevation, tuple(config["data"]["target_grid"])) optimizer.zero_grad(set_to_none=True) prediction = train_model(inputs) loss = processed_loss(prediction, target, config["training"]["loss"], float(config["training"]["exponential_alpha"]), float(config["training"]["exponential_weight"]), float(config["training"]["quantile"])) if not torch.isfinite(loss): raise FloatingPointError("non-finite training loss") loss.backward(); optimizer.step(); losses.append(float(loss.detach())) if distributed: gathered = [None] * world_size torch.distributed.all_gather_object(gathered, losses) losses = [value for part in gathered for value in part] if rank == 0: checkpoint_path = ROOT / config["paths"]["checkpoint"] checkpoint_path.parent.mkdir(parents=True, exist_ok=True) checkpoint = {"format_version": FORMAT_VERSION, "model_name": MODEL_NAME, "model_config": config["model"], "paper_model": config["paper_model"], "target_grid": config["data"]["target_grid"], "model": model.state_dict(), "loss": config["training"]["loss"], "seed": seed} torch.save(checkpoint, checkpoint_path) metrics_path = ROOT / config["paths"]["training_metrics"] metrics_path.parent.mkdir(parents=True, exist_ok=True) metrics = {"format_version": FORMAT_VERSION, "epochs": int(config["training"]["epochs"]), "world_size": world_size, "loss": config["training"]["loss"], "loss_history": losses, "final_loss": losses[-1]} metrics_path.write_text(json.dumps(metrics, indent=2) + "\n") print(f"checkpoint={checkpoint_path.relative_to(ROOT)} final_loss={losses[-1]:.6f}") if distributed: torch.distributed.barrier(); torch.distributed.destroy_process_group() if __name__ == "__main__": main()