"""DDP-capable FuXi-DA training with analysis and frozen-proxy forecast supervision.""" import argparse import json import math import os from pathlib import Path import torch import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel from torch.utils.data import DataLoader, DistributedSampler import sys ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) import yaml from model.fuxi_da import CompactForecastProxy, FuXiDA, ProceduralTileDataset def latitude_weighted_l1(prediction, target, latitude): weight = torch.cos(torch.deg2rad(latitude)).clamp_min(0) weight = weight * weight.shape[-1] / weight.sum(dim=-1, keepdim=True) return ((prediction - target).abs() * weight[:, None, :, None]).mean() def learning_rate(step, warmup, total, peak): if step < warmup: return 1e-8 + (peak - 1e-8) * step / max(1, warmup) progress = (step - warmup) / max(1, total - warmup) return peak * 0.5 * (1.0 + math.cos(math.pi * progress)) def main(): parser = argparse.ArgumentParser() parser.add_argument("--config", default="conf/config.yaml") parser.add_argument("--iterations", type=int) parser.add_argument("--forecast-steps", type=int) parser.add_argument("--output-dir") args = parser.parse_args() cfg = yaml.safe_load((ROOT / args.config).read_text()) total = args.iterations or cfg["train"]["iterations"] forecast_steps = args.forecast_steps or cfg["train"]["forecast_steps"] world_size = int(os.environ.get("WORLD_SIZE", "1")) local_rank = int(os.environ.get("LOCAL_RANK", "0")) use_cuda = torch.cuda.is_available() and torch.cuda.device_count() >= world_size and cfg["runtime"]["device"] != "cpu" if world_size > 1: dist.init_process_group("nccl" if use_cuda else "gloo") device = torch.device(f"cuda:{local_rank}" if use_cuda else "cpu") if use_cuda: torch.cuda.set_device(device) torch.manual_seed(cfg["seed"] + local_rank) dataset = ProceduralTileDataset(cfg["data"]["tile_ids"], max(2, total), cfg["data"]["tile_size"], cfg["data"]["missing_probability"]) sampler = DistributedSampler(dataset, shuffle=True) if world_size > 1 else None loader = DataLoader(dataset, batch_size=cfg["train"]["batch_size"], sampler=sampler, shuffle=sampler is None, num_workers=0) model = FuXiDA(cfg["model"]["base_channels"]).to(device) proxy = CompactForecastProxy().to(device) if world_size > 1: model = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None) optimizer = torch.optim.AdamW(model.parameters(), lr=cfg["train"]["learning_rate"], betas=(0.9, 0.999), weight_decay=cfg["train"]["weight_decay"]) model.train() iterator = iter(loader) for step in range(total): try: batch = next(iterator) except StopIteration: if sampler is not None: sampler.set_epoch(step) iterator = iter(loader) batch = next(iterator) background, obs, target = (batch[key].to(device) for key in ("background", "obs", "target")) latitude = batch["latitude"].to(device) analysis = model(background, obs) analysis_loss = latitude_weighted_l1(analysis, target, latitude) state, forecast_loss = analysis, analysis_loss.new_zeros(()) for lead in range(forecast_steps): state = proxy(state) forecast_loss = forecast_loss + latitude_weighted_l1(state, batch["forecast_targets"][:, lead].to(device), latitude) loss = analysis_loss + forecast_loss / forecast_steps optimizer.zero_grad(set_to_none=True) loss.backward() optimizer.step() lr = learning_rate(step + 1, cfg["train"]["warmup_steps"], total, cfg["train"]["learning_rate"]) for group in optimizer.param_groups: group["lr"] = lr if local_rank == 0 and (step == 0 or (step + 1) % 100 == 0 or step + 1 == total): print(json.dumps({"step": step + 1, "loss": loss.item(), "analysis_l1": analysis_loss.item(), "lr": lr})) if local_rank == 0: checkpoint = ROOT / cfg["paths"]["checkpoint"]; checkpoint.parent.mkdir(parents=True, exist_ok=True) bare_model = model.module if isinstance(model, DistributedDataParallel) else model torch.save({"model": bare_model.state_dict(), "model_config": cfg["model"], "format_version": cfg["data"]["format_version"]}, checkpoint) metrics = ROOT / cfg["paths"]["training_metrics"]; metrics.parent.mkdir(parents=True, exist_ok=True) metrics.write_text(json.dumps({"iterations": total, "final_loss": float(loss), "world_size": world_size}, indent=2)) if world_size > 1: dist.destroy_process_group() if __name__ == "__main__": main()