Download scripts/train.py from OneScience-Group/CRAI-ClimateExtremes: direct link, hf CLI and curl.
- Browser
- Download file 4.88 kB
-
https://huggingface.co/OneScience-Group/CRAI-ClimateExtremes/resolve/main/scripts/train.py
- Command line
-
hf download hf://OneScience-Group/CRAI-ClimateExtremes/scripts/train.py
-
curl -L -o train.py https://huggingface.co/OneScience-Group/CRAI-ClimateExtremes/resolve/main/scripts/train.py
4.88 kB
| """Train independent CRAI ensemble members, optionally under torchrun DDP.""" | |
| from pathlib import Path | |
| import argparse | |
| import json | |
| import os | |
| import random | |
| import sys | |
| import numpy as np | |
| import torch | |
| from torch import distributed as dist | |
| from torch.nn.parallel import DistributedDataParallel | |
| from torch.utils.data import DataLoader, Dataset, DistributedSampler | |
| import yaml | |
| ROOT = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(ROOT)) | |
| from model.crai_climateextremes import CRAIClimateExtremes | |
| class ClimateDataset(Dataset): | |
| def __init__(self, archive, limit): | |
| self.observed = torch.from_numpy(archive["observed"][:limit]) | |
| self.valid = torch.from_numpy(archive["valid_mask"][:limit]) | |
| self.target = torch.from_numpy(archive["target"][:limit]) | |
| land = torch.from_numpy(archive["europe_mask"])[None, None] | |
| self.missing = land * (1.0 - self.valid) | |
| def __len__(self): | |
| return len(self.target) | |
| def __getitem__(self, index): | |
| return torch.cat((self.observed[index], self.valid[index]), 0), self.target[index], self.missing[index] | |
| def load_config(path): | |
| with open(path, encoding="utf-8") as handle: | |
| return yaml.safe_load(handle) | |
| def main(): | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--config", type=Path, default=ROOT / "conf/config.yaml") | |
| parser.add_argument("--paper-model", action="store_true") | |
| args = parser.parse_args() | |
| config_path = args.config if args.config.is_absolute() else ROOT / args.config | |
| cfg = load_config(config_path) | |
| use_paper = args.paper_model or cfg["paper_model"] | |
| batch_size = cfg["paper_batch_size"] if use_paper else cfg["batch_size"] | |
| iterations = cfg["paper_iterations"] if use_paper else cfg["max_iterations"] | |
| members = cfg["paper_ensemble_members"] if use_paper else cfg["ensemble_members"] | |
| rank, world = int(os.getenv("RANK", 0)), int(os.getenv("WORLD_SIZE", 1)) | |
| if world > 1: | |
| dist.init_process_group("nccl" if torch.cuda.is_available() else "gloo") | |
| device = torch.device(f"cuda:{int(os.getenv('LOCAL_RANK', 0))}" if torch.cuda.is_available() else "cpu") | |
| if device.type == "cuda": | |
| torch.cuda.set_device(device) | |
| archive = np.load(ROOT / cfg["data_path"]) | |
| dataset = ClimateDataset(archive, cfg["num_samples"]) | |
| checkpoint_path = ROOT / cfg["checkpoint_path"] | |
| if rank == 0: | |
| checkpoint_path.parent.mkdir(parents=True, exist_ok=True) | |
| (ROOT / "result/training").mkdir(parents=True, exist_ok=True) | |
| if world > 1: | |
| dist.barrier() | |
| records, member_states = [], [] | |
| for member in range(members): | |
| seed = int(cfg["seed"]) + member | |
| random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) | |
| sampler = DistributedSampler(dataset, shuffle=True, seed=seed) if world > 1 else None | |
| loader = DataLoader(dataset, batch_size=batch_size, shuffle=sampler is None, sampler=sampler) | |
| model = CRAIClimateExtremes(cfg["base_channels"]).to(device) | |
| if world > 1: | |
| model = DistributedDataParallel(model, device_ids=[device.index] if device.type == "cuda" else None) | |
| optimizer = torch.optim.Adam(model.parameters(), lr=cfg["learning_rate"]) | |
| step, losses = 0, [] | |
| for epoch in range(cfg["epochs"] if not use_paper else 10**9): | |
| if sampler is not None: | |
| sampler.set_epoch(epoch) | |
| for inputs, target, missing in loader: | |
| inputs, target, missing = inputs.to(device), target.to(device), missing.to(device) | |
| prediction = model(inputs) | |
| loss = (torch.abs(prediction - target) * missing).sum() / missing.sum().clamp_min(1) | |
| optimizer.zero_grad(); loss.backward(); optimizer.step() | |
| losses.append(float(loss.detach())) | |
| step += 1 | |
| if step >= iterations: | |
| break | |
| if step >= iterations: | |
| break | |
| raw_model = model.module if isinstance(model, DistributedDataParallel) else model | |
| if rank == 0: | |
| member_states.append({key: value.detach().cpu() for key, value in raw_model.state_dict().items()}) | |
| records.append({"member": member, "iterations": step, "final_missing_mae": losses[-1]}) | |
| if rank == 0: | |
| torch.save( | |
| { | |
| "format_version": "1.0", | |
| "model_config": {"base_channels": cfg["base_channels"]}, | |
| "model": member_states, | |
| }, | |
| checkpoint_path, | |
| ) | |
| payload = {"paper_model": bool(use_paper), "world_size": world, "members": records} | |
| (ROOT / "result/training/metrics.json").write_text(json.dumps(payload, indent=2) + "\n") | |
| print(json.dumps(payload)) | |
| if world > 1: | |
| dist.destroy_process_group() | |
| if __name__ == "__main__": | |
| main() | |