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