OneScience-Group/CRAI-ClimateExtremes
027
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 