CoolFace
Modelpublic

OneScience-Group/MassConservingCNN

sourceHugging Faceapache-2.0updated 16d agoView on Hugging Face
0likes28downloads
train.py129 linesDownload Raw Back to scripts
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