FuXi-DA / scripts /train.py
zhangrenchao's picture
Publish FuXi-DA reproduction
15aff58 verified
Raw
History Blame Contribute Delete
4.82 kB
"""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()