"""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()