Download scripts/train.py from OneScience-Group/UnetDif: direct link, hf CLI and curl.
- Browser
- Download file 5.8 kB
-
https://huggingface.co/OneScience-Group/UnetDif/resolve/main/scripts/train.py
- Command line
-
hf download hf://OneScience-Group/UnetDif/scripts/train.py
-
curl -L -o train.py https://huggingface.co/OneScience-Group/UnetDif/resolve/main/scripts/train.py
5.8 kB
| """Train UnetDif with single-process or torchrun DDP execution.""" | |
| import json | |
| import os | |
| import sys | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import yaml | |
| from torch.nn.parallel import DistributedDataParallel as DDP | |
| from torch.utils.data import DataLoader, Dataset, DistributedSampler | |
| ROOT = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(ROOT)) | |
| from model.unetdif import UnetDif, loss_components | |
| class RainfallDataset(Dataset): | |
| def __init__(self, path: Path, config): | |
| archive = np.load(path, allow_pickle=False) | |
| expected = (int(config["time_steps"]), int(config["channels"]), int(config["height"]), int(config["width"])) | |
| if archive["inputs"].shape[1:] != expected or archive["targets"].shape[1:] != (expected[0], expected[2], expected[3]): | |
| raise ValueError(f"expected day shapes {expected} and {(expected[0], expected[2], expected[3])}") | |
| if len(np.unique(archive["group_id"])) != len(archive["group_id"]): | |
| raise ValueError("each rain day must belong to one split group") | |
| if archive["lead_hours"].tolist() != config["lead_hours"]: | |
| raise ValueError("NPZ lead_hours do not match the configured eight 3-hour steps") | |
| self.inputs = archive["inputs"].astype(np.float32, copy=False).reshape(-1, *expected[1:]) | |
| self.targets = archive["targets"].astype(np.float32, copy=False).reshape(-1, expected[2], expected[3]) | |
| def __len__(self): | |
| return len(self.inputs) | |
| def __getitem__(self, index): | |
| return torch.from_numpy(self.inputs[index]), torch.from_numpy(self.targets[index]) | |
| def get_device(local_rank=0, local_world_size=1): | |
| use_cuda = torch.cuda.is_available() and torch.cuda.device_count() >= local_world_size | |
| return torch.device("cuda", local_rank) if use_cuda else torch.device("cpu") | |
| def reduce_values(values, device, distributed): | |
| tensor = torch.tensor(values, device=device) | |
| if distributed: | |
| torch.distributed.all_reduce(tensor) | |
| return tensor.cpu().tolist() | |
| def main(): | |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) | |
| torch.manual_seed(int(config["seed"])) | |
| distributed = int(os.environ.get("WORLD_SIZE", "1")) > 1 | |
| local_rank = int(os.environ.get("LOCAL_RANK", "0")) | |
| local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1")) | |
| device = get_device(local_rank, local_world_size) | |
| if distributed: | |
| torch.distributed.init_process_group("nccl" if device.type == "cuda" else "gloo") | |
| rank = torch.distributed.get_rank() if distributed else 0 | |
| dataset = RainfallDataset(ROOT / config["data"]["root"] / "train.npz", config["data"]) | |
| if distributed and len(dataset) < torch.distributed.get_world_size(): | |
| raise ValueError("DDP requires at least one training item per process") | |
| sampler = DistributedSampler(dataset) if distributed else None | |
| loader = DataLoader(dataset, batch_size=int(config["train"]["batch_size"]), shuffle=sampler is None, | |
| sampler=sampler, num_workers=int(config["train"]["num_workers"])) | |
| model = UnetDif(**config["model"]).to(device) | |
| if distributed: | |
| model = DDP(model, device_ids=[local_rank] if device.type == "cuda" else None) | |
| optimizer = torch.optim.Adam(model.parameters(), lr=float(config["train"]["learning_rate"])) | |
| names = ["dry_focal", "false_alarm_focal", "positive_mse", "negative_mse", "dry_mae", "all_mae"] | |
| normalizers = {name: 1.0 for name in names} | |
| history = [] | |
| for epoch in range(int(config["train"]["epochs"])): | |
| if sampler is not None: | |
| sampler.set_epoch(epoch) | |
| model.train(); sums = {name: 0.0 for name in names}; batches = 0 | |
| for inputs, target in loader: | |
| inputs, target = inputs.to(device), target.to(device) | |
| heads = model(inputs) | |
| raw = loss_components(heads, inputs[:, int(config["data"]["precipitation_channel"])], target, | |
| float(config["data"]["rain_threshold_mm_3h"]), | |
| float(config["loss"]["focal_alpha"]), float(config["loss"]["focal_gamma"])) | |
| loss = sum(raw[name] * normalizers[name] for name in names) | |
| optimizer.zero_grad(set_to_none=True); loss.backward(); optimizer.step() | |
| for name in names: | |
| sums[name] += float(raw[name].detach()) | |
| batches += 1 | |
| reduced = reduce_values([sums[name] for name in names] + [batches], device, distributed) | |
| total_batches = max(reduced[-1], 1) | |
| averages = {name: reduced[index] / total_batches for index, name in enumerate(names)} | |
| if epoch == 0: | |
| normalizers = {name: 1.0 / max(value, 1e-6) for name, value in averages.items()} | |
| record = {"epoch": epoch + 1, "components": averages, "normalizers": normalizers} | |
| history.append(record) | |
| if rank == 0: | |
| print(json.dumps(record)) | |
| checkpoint = ROOT / config["paths"]["checkpoint"] | |
| checkpoint.parent.mkdir(parents=True, exist_ok=True) | |
| state = model.module.state_dict() if distributed else model.state_dict() | |
| torch.save({"epoch": epoch + 1, "model": state, "optimizer": optimizer.state_dict(), | |
| "model_config": config["model"], "loss_normalizers": normalizers, | |
| "loss_components": names, | |
| "format_version": config["data"]["format_version"]}, checkpoint) | |
| if rank == 0: | |
| metrics = ROOT / config["paths"]["training_metrics"] | |
| metrics.parent.mkdir(parents=True, exist_ok=True) | |
| metrics.write_text(json.dumps({"history": history}, indent=2) + "\n") | |
| if distributed: | |
| torch.distributed.destroy_process_group() | |
| if __name__ == "__main__": | |
| main() | |