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