Team Ai
Modelpublic

OneScience-Group/CRAI-ClimateExtremes

sourceHugging Faceapache-2.0updated 29d agoView on Hugging Face
0likes27downloads
train.py114 linesDownload Raw Back to scripts
1"""Train independent CRAI ensemble members, optionally under torchrun DDP."""2 3from pathlib import Path4import argparse5import json6import os7import random8import sys9import numpy as np10import torch11from torch import distributed as dist12from torch.nn.parallel import DistributedDataParallel13from torch.utils.data import DataLoader, Dataset, DistributedSampler14import yaml15 16ROOT = Path(__file__).resolve().parents[1]17sys.path.insert(0, str(ROOT))18 19from model.crai_climateextremes import CRAIClimateExtremes20 21 22class ClimateDataset(Dataset):23    def __init__(self, archive, limit):24        self.observed = torch.from_numpy(archive["observed"][:limit])25        self.valid = torch.from_numpy(archive["valid_mask"][:limit])26        self.target = torch.from_numpy(archive["target"][:limit])27        land = torch.from_numpy(archive["europe_mask"])[None, None]28        self.missing = land * (1.0 - self.valid)29 30    def __len__(self):31        return len(self.target)32 33    def __getitem__(self, index):34        return torch.cat((self.observed[index], self.valid[index]), 0), self.target[index], self.missing[index]35 36 37def load_config(path):38    with open(path, encoding="utf-8") as handle:39        return yaml.safe_load(handle)40 41 42def main():43    parser = argparse.ArgumentParser()44    parser.add_argument("--config", type=Path, default=ROOT / "conf/config.yaml")45    parser.add_argument("--paper-model", action="store_true")46    args = parser.parse_args()47    config_path = args.config if args.config.is_absolute() else ROOT / args.config48    cfg = load_config(config_path)49    use_paper = args.paper_model or cfg["paper_model"]50    batch_size = cfg["paper_batch_size"] if use_paper else cfg["batch_size"]51    iterations = cfg["paper_iterations"] if use_paper else cfg["max_iterations"]52    members = cfg["paper_ensemble_members"] if use_paper else cfg["ensemble_members"]53    rank, world = int(os.getenv("RANK", 0)), int(os.getenv("WORLD_SIZE", 1))54    if world > 1:55        dist.init_process_group("nccl" if torch.cuda.is_available() else "gloo")56    device = torch.device(f"cuda:{int(os.getenv('LOCAL_RANK', 0))}" if torch.cuda.is_available() else "cpu")57    if device.type == "cuda":58        torch.cuda.set_device(device)59    archive = np.load(ROOT / cfg["data_path"])60    dataset = ClimateDataset(archive, cfg["num_samples"])61    checkpoint_path = ROOT / cfg["checkpoint_path"]62    if rank == 0:63        checkpoint_path.parent.mkdir(parents=True, exist_ok=True)64        (ROOT / "result/training").mkdir(parents=True, exist_ok=True)65    if world > 1:66        dist.barrier()67    records, member_states = [], []68    for member in range(members):69        seed = int(cfg["seed"]) + member70        random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)71        sampler = DistributedSampler(dataset, shuffle=True, seed=seed) if world > 1 else None72        loader = DataLoader(dataset, batch_size=batch_size, shuffle=sampler is None, sampler=sampler)73        model = CRAIClimateExtremes(cfg["base_channels"]).to(device)74        if world > 1:75            model = DistributedDataParallel(model, device_ids=[device.index] if device.type == "cuda" else None)76        optimizer = torch.optim.Adam(model.parameters(), lr=cfg["learning_rate"])77        step, losses = 0, []78        for epoch in range(cfg["epochs"] if not use_paper else 10**9):79            if sampler is not None:80                sampler.set_epoch(epoch)81            for inputs, target, missing in loader:82                inputs, target, missing = inputs.to(device), target.to(device), missing.to(device)83                prediction = model(inputs)84                loss = (torch.abs(prediction - target) * missing).sum() / missing.sum().clamp_min(1)85                optimizer.zero_grad(); loss.backward(); optimizer.step()86                losses.append(float(loss.detach()))87                step += 188                if step >= iterations:89                    break90            if step >= iterations:91                break92        raw_model = model.module if isinstance(model, DistributedDataParallel) else model93        if rank == 0:94            member_states.append({key: value.detach().cpu() for key, value in raw_model.state_dict().items()})95            records.append({"member": member, "iterations": step, "final_missing_mae": losses[-1]})96    if rank == 0:97        torch.save(98            {99                "format_version": "1.0",100                "model_config": {"base_channels": cfg["base_channels"]},101                "model": member_states,102            },103            checkpoint_path,104        )105        payload = {"paper_model": bool(use_paper), "world_size": world, "members": records}106        (ROOT / "result/training/metrics.json").write_text(json.dumps(payload, indent=2) + "\n")107        print(json.dumps(payload))108    if world > 1:109        dist.destroy_process_group()110 111 112if __name__ == "__main__":113    main()114