File size: 4,780 Bytes
03573b6 | 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 | """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()
|