Download scripts/train.py from OneScience-Group/MassConservingCNN: direct link, hf CLI and curl.
- Browser
- Download file 6.69 kB
-
https://huggingface.co/OneScience-Group/MassConservingCNN/resolve/main/scripts/train.py
- Command line
-
hf download hf://OneScience-Group/MassConservingCNN/scripts/train.py
-
curl -L -o train.py https://huggingface.co/OneScience-Group/MassConservingCNN/resolve/main/scripts/train.py
6.69 kB
| """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() | |