| """Train MassConservingCNN with optional torchrun DDP.""" |
|
|
| 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 |
| from torch.utils.data import DataLoader, Dataset, DistributedSampler |
|
|
|
|
| ROOT = Path(__file__).resolve().parents[1] |
| sys.path.insert(0, str(ROOT)) |
| from model.massconservingcnn import MassConservingCNN |
|
|
|
|
| class MSWDataset(Dataset): |
| def __init__(self, path, config): |
| self.data = np.load(path) |
| if str(self.data["format_version"]) != config["data"]["format_version"]: |
| raise ValueError("incompatible data format version") |
| count = len(self.data["inputs"]) |
| if self.data["inputs"].shape != (count, 4, 250): |
| raise ValueError("inputs must have shape [B,4,250]") |
| if self.data["targets"].shape != (count, 3, 250): |
| raise ValueError("targets must have shape [B,3,250]") |
| if self.data["inputs"].dtype != np.float32 or self.data["targets"].dtype != np.float32: |
| raise TypeError("inputs and targets must be float32") |
| if not np.isfinite(self.data["inputs"]).all() or not np.isfinite(self.data["targets"]).all(): |
| raise ValueError("data must be finite") |
| if not np.isin(self.data["radar"], (0.0, 1.0)).all(): |
| raise ValueError("radar indicator must be binary") |
| if (self.data["inputs"][:, 2] < 0).any() or (self.data["targets"][:, 2] < 0).any(): |
| raise ValueError("normalized rain must remain non-negative") |
|
|
| def __len__(self): |
| return len(self.data["inputs"]) |
|
|
| def __getitem__(self, index): |
| return torch.from_numpy(self.data["inputs"][index]), torch.from_numpy(self.data["targets"][index]) |
|
|
|
|
| def paper_j(prediction, target): |
| return torch.sqrt(torch.mean((prediction - target) ** 2, dim=2) + 1e-12).mean(dim=1) |
|
|
|
|
| def mass_aware_loss(prediction, target, eta): |
| base = paper_j(prediction, target) |
| mass = eta / prediction.shape[2] * torch.abs(prediction[:, 1].sum(1) - target[:, 1].sum(1)) |
| return (base + mass).mean(), base.mean(), mass.mean() |
|
|
|
|
| def device_from_config(config, local_rank=0): |
| requested = config["runtime"]["device"] |
| if requested == "auto": |
| return torch.device("cuda", local_rank) if torch.cuda.is_available() else torch.device("cpu") |
| return torch.device(requested) |
|
|
|
|
| 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")) |
| if distributed: |
| torch.distributed.init_process_group("nccl" if torch.cuda.is_available() else "gloo") |
| rank = torch.distributed.get_rank() if distributed else 0 |
| device = device_from_config(config, local_rank) |
| if device.type == "cuda": |
| torch.cuda.set_device(device); torch.cuda.manual_seed_all(seed) |
| train_set = MSWDataset(ROOT / config["data"]["root"] / "train.npz", config) |
| valid_set = MSWDataset(ROOT / config["data"]["root"] / "validation.npz", config) |
| sampler = DistributedSampler(train_set, shuffle=True, seed=seed) if distributed else None |
| loader = DataLoader(train_set, batch_size=int(config["train"]["batch_size"]), |
| shuffle=sampler is None, sampler=sampler, |
| num_workers=int(config["train"]["num_workers"])) |
| valid_loader = DataLoader(valid_set, batch_size=int(config["train"]["batch_size"]), shuffle=False) |
| model = MassConservingCNN(**config["model"]).to(device) |
| wrapped = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None) if distributed else model |
| bare = wrapped.module if distributed else wrapped |
| optimizer = torch.optim.Adam(wrapped.parameters(), lr=float(config["train"]["learning_rate"])) |
| history = [] |
| for epoch in range(int(config["train"]["epochs"])): |
| if sampler is not None: |
| sampler.set_epoch(epoch) |
| wrapped.train(); total = 0.0; seen = 0 |
| for inputs, targets in loader: |
| prediction = wrapped(inputs.to(device)); loss, _, _ = mass_aware_loss(prediction, targets.to(device), float(config["train"]["eta"])) |
| optimizer.zero_grad(set_to_none=True); loss.backward(); optimizer.step() |
| total += float(loss.detach()) * len(inputs); seen += len(inputs) |
| totals = torch.tensor([total, seen], dtype=torch.float64, device=device) |
| if distributed: |
| torch.distributed.all_reduce(totals) |
| wrapped.eval(); valid_total = valid_j = valid_mass = 0.0; valid_seen = 0 |
| if rank == 0: |
| with torch.no_grad(): |
| for inputs, targets in valid_loader: |
| loss, base, mass = mass_aware_loss(bare(inputs.to(device)), targets.to(device), float(config["train"]["eta"])) |
| valid_total += float(loss) * len(inputs); valid_j += float(base) * len(inputs) |
| valid_mass += float(mass) * len(inputs); valid_seen += len(inputs) |
| history.append({"epoch": epoch + 1, "train_loss": float(totals[0] / totals[1]), |
| "validation_loss": valid_total / valid_seen, "validation_J": valid_j / valid_seen, |
| "validation_mass_penalty": valid_mass / valid_seen}) |
| if rank == 0: |
| checkpoint_path = ROOT / config["paths"]["checkpoint"] |
| metrics_path = ROOT / config["paths"]["training_metrics"] |
| checkpoint_path.parent.mkdir(parents=True, exist_ok=True); metrics_path.parent.mkdir(parents=True, exist_ok=True) |
| model_state = bare.state_dict() |
| torch.save({"model": model_state, "model_state_dict": model_state, |
| "optimizer_state_dict": optimizer.state_dict(), |
| "model_config": config["model"], "epoch": int(config["train"]["epochs"]), |
| "eta": float(config["train"]["eta"]), "format_version": config["data"]["format_version"], |
| "variable_order": ["u", "h", "r"], "normalization": "u,h: center/scale; r: scale only", |
| "climate_mean_uh": train_set.data["climate_mean_uh"], |
| "climate_std_uhr": train_set.data["climate_std_uhr"], "seed": seed}, checkpoint_path) |
| metrics_path.write_text(json.dumps({"history": history}, indent=2) + "\n") |
| print(f"checkpoint={checkpoint_path.relative_to(ROOT)} validation_loss={history[-1]['validation_loss']:.6f}") |
| if distributed: |
| torch.distributed.destroy_process_group() |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|