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