Download scripts/train.py from OneScience-Group/FuXi-DA: direct link, hf CLI and curl.
- Browser
- Download file 4.82 kB
-
https://huggingface.co/OneScience-Group/FuXi-DA/resolve/main/scripts/train.py
- Command line
-
hf download hf://OneScience-Group/FuXi-DA/scripts/train.py
-
curl -L -o train.py https://huggingface.co/OneScience-Group/FuXi-DA/resolve/main/scripts/train.py
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() | |