OneScience-Group/MassConservingCNN
028
1"""Train MassConservingCNN with optional torchrun DDP."""2 3import json4import os5import random6import sys7from pathlib import Path8 9import numpy as np10import torch11import yaml12from torch.nn.parallel import DistributedDataParallel13from torch.utils.data import DataLoader, Dataset, DistributedSampler14 15 16ROOT = Path(__file__).resolve().parents[1]17sys.path.insert(0, str(ROOT))18from model.massconservingcnn import MassConservingCNN19 20 21class MSWDataset(Dataset):22 def __init__(self, path, config):23 self.data = np.load(path)24 if str(self.data["format_version"]) != config["data"]["format_version"]:25 raise ValueError("incompatible data format version")26 count = len(self.data["inputs"])27 if self.data["inputs"].shape != (count, 4, 250):28 raise ValueError("inputs must have shape [B,4,250]")29 if self.data["targets"].shape != (count, 3, 250):30 raise ValueError("targets must have shape [B,3,250]")31 if self.data["inputs"].dtype != np.float32 or self.data["targets"].dtype != np.float32:32 raise TypeError("inputs and targets must be float32")33 if not np.isfinite(self.data["inputs"]).all() or not np.isfinite(self.data["targets"]).all():34 raise ValueError("data must be finite")35 if not np.isin(self.data["radar"], (0.0, 1.0)).all():36 raise ValueError("radar indicator must be binary")37 if (self.data["inputs"][:, 2] < 0).any() or (self.data["targets"][:, 2] < 0).any():38 raise ValueError("normalized rain must remain non-negative")39 40 def __len__(self):41 return len(self.data["inputs"])42 43 def __getitem__(self, index):44 return torch.from_numpy(self.data["inputs"][index]), torch.from_numpy(self.data["targets"][index])45 46 47def paper_j(prediction, target):48 return torch.sqrt(torch.mean((prediction - target) ** 2, dim=2) + 1e-12).mean(dim=1)49 50 51def mass_aware_loss(prediction, target, eta):52 base = paper_j(prediction, target)53 mass = eta / prediction.shape[2] * torch.abs(prediction[:, 1].sum(1) - target[:, 1].sum(1))54 return (base + mass).mean(), base.mean(), mass.mean()55 56 57def device_from_config(config, local_rank=0):58 requested = config["runtime"]["device"]59 if requested == "auto":60 return torch.device("cuda", local_rank) if torch.cuda.is_available() else torch.device("cpu")61 return torch.device(requested)62 63 64def main():65 config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())66 seed = int(config["seed"])67 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)68 distributed = int(os.environ.get("WORLD_SIZE", "1")) > 169 local_rank = int(os.environ.get("LOCAL_RANK", "0"))70 if distributed:71 torch.distributed.init_process_group("nccl" if torch.cuda.is_available() else "gloo")72 rank = torch.distributed.get_rank() if distributed else 073 device = device_from_config(config, local_rank)74 if device.type == "cuda":75 torch.cuda.set_device(device); torch.cuda.manual_seed_all(seed)76 train_set = MSWDataset(ROOT / config["data"]["root"] / "train.npz", config)77 valid_set = MSWDataset(ROOT / config["data"]["root"] / "validation.npz", config)78 sampler = DistributedSampler(train_set, shuffle=True, seed=seed) if distributed else None79 loader = DataLoader(train_set, batch_size=int(config["train"]["batch_size"]),80 shuffle=sampler is None, sampler=sampler,81 num_workers=int(config["train"]["num_workers"]))82 valid_loader = DataLoader(valid_set, batch_size=int(config["train"]["batch_size"]), shuffle=False)83 model = MassConservingCNN(**config["model"]).to(device)84 wrapped = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None) if distributed else model85 bare = wrapped.module if distributed else wrapped86 optimizer = torch.optim.Adam(wrapped.parameters(), lr=float(config["train"]["learning_rate"]))87 history = []88 for epoch in range(int(config["train"]["epochs"])):89 if sampler is not None:90 sampler.set_epoch(epoch)91 wrapped.train(); total = 0.0; seen = 092 for inputs, targets in loader:93 prediction = wrapped(inputs.to(device)); loss, _, _ = mass_aware_loss(prediction, targets.to(device), float(config["train"]["eta"]))94 optimizer.zero_grad(set_to_none=True); loss.backward(); optimizer.step()95 total += float(loss.detach()) * len(inputs); seen += len(inputs)96 totals = torch.tensor([total, seen], dtype=torch.float64, device=device)97 if distributed:98 torch.distributed.all_reduce(totals)99 wrapped.eval(); valid_total = valid_j = valid_mass = 0.0; valid_seen = 0100 if rank == 0:101 with torch.no_grad():102 for inputs, targets in valid_loader:103 loss, base, mass = mass_aware_loss(bare(inputs.to(device)), targets.to(device), float(config["train"]["eta"]))104 valid_total += float(loss) * len(inputs); valid_j += float(base) * len(inputs)105 valid_mass += float(mass) * len(inputs); valid_seen += len(inputs)106 history.append({"epoch": epoch + 1, "train_loss": float(totals[0] / totals[1]),107 "validation_loss": valid_total / valid_seen, "validation_J": valid_j / valid_seen,108 "validation_mass_penalty": valid_mass / valid_seen})109 if rank == 0:110 checkpoint_path = ROOT / config["paths"]["checkpoint"]111 metrics_path = ROOT / config["paths"]["training_metrics"]112 checkpoint_path.parent.mkdir(parents=True, exist_ok=True); metrics_path.parent.mkdir(parents=True, exist_ok=True)113 model_state = bare.state_dict()114 torch.save({"model": model_state, "model_state_dict": model_state,115 "optimizer_state_dict": optimizer.state_dict(),116 "model_config": config["model"], "epoch": int(config["train"]["epochs"]),117 "eta": float(config["train"]["eta"]), "format_version": config["data"]["format_version"],118 "variable_order": ["u", "h", "r"], "normalization": "u,h: center/scale; r: scale only",119 "climate_mean_uh": train_set.data["climate_mean_uh"],120 "climate_std_uhr": train_set.data["climate_std_uhr"], "seed": seed}, checkpoint_path)121 metrics_path.write_text(json.dumps({"history": history}, indent=2) + "\n")122 print(f"checkpoint={checkpoint_path.relative_to(ROOT)} validation_loss={history[-1]['validation_loss']:.6f}")123 if distributed:124 torch.distributed.destroy_process_group()125 126 127if __name__ == "__main__":128 main()129 