| |
| """Train FengWu-W2S on ERA5-style HDF5 windows (single or torchrun DDP).""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import math |
| import os |
| import random |
| import sys |
| from pathlib import Path |
| from typing import Any, Dict |
|
|
| import numpy as np |
| import torch |
| import yaml |
| from torch import nn |
| from torch.nn.parallel import DistributedDataParallel as DDP |
|
|
| PROJECT_ROOT = Path(__file__).resolve().parents[1] |
| sys.path.insert(0, str(PROJECT_ROOT)) |
|
|
| from model.fengwu_w2s import FengWuW2S |
| from scripts.data_loader import make_dataloader, resolve_data_dir |
|
|
|
|
| def parse_args() -> argparse.Namespace: |
| parser = argparse.ArgumentParser(description=__doc__) |
| parser.add_argument("--config", default=str(PROJECT_ROOT / "conf/config.yaml")) |
| parser.add_argument("--data-dir") |
| parser.add_argument("--checkpoint") |
| parser.add_argument("--max-epoch", type=int) |
| parser.add_argument("--rollout-steps", type=int) |
| parser.add_argument("--batch-size", type=int) |
| parser.add_argument("--num-workers", type=int) |
| parser.add_argument("--device", choices=("auto", "cpu", "cuda"), help="Execution device; auto prefers DCU/GPU") |
| parser.add_argument("--finetune", action="store_true", help="Resume with the lower finetune learning rate") |
| parser.add_argument("--seed", type=int, default=42) |
| return parser.parse_args() |
|
|
|
|
| def _device(config: Dict[str, Any], local_rank: int, override: str | None = None) -> torch.device: |
| requested = str(override or config.get("runtime", {}).get("device", "auto")).lower() |
| if requested not in {"auto", "cpu", "cuda"}: |
| raise ValueError(f"unsupported device {requested!r}; expected auto, cpu, or cuda") |
| cuda_available = bool(torch.cuda.is_available()) |
| cuda_count = int(torch.cuda.device_count()) if cuda_available else 0 |
| if requested == "cpu": |
| return torch.device("cpu") |
| if not cuda_available or cuda_count < 1: |
| if requested == "cuda": |
| raise RuntimeError( |
| "CUDA/DCU was explicitly requested but torch.cuda is unavailable; " |
| "load the accelerator module before starting training" |
| ) |
| return torch.device("cpu") |
| index = local_rank |
| if index >= cuda_count: |
| raise RuntimeError(f"local rank {index} exceeds visible accelerator count {cuda_count}") |
| torch.cuda.set_device(index) |
| return torch.device("cuda", index) |
|
|
|
|
| def _distributed_context() -> tuple[bool, int, int]: |
| world_size = int(os.environ.get("WORLD_SIZE", "1")) |
| distributed = world_size > 1 |
| rank = int(os.environ.get("RANK", "0")) |
| local_rank = int(os.environ.get("LOCAL_RANK", "0")) |
| return distributed, rank, local_rank |
|
|
|
|
| def _model_from_config(config: Dict[str, Any]) -> FengWuW2S: |
| model_cfg = config["model"] |
| data_cfg = config["data"] |
| groups = {name: tuple(values) for name, values in data_cfg["groups"].items()} |
| return FengWuW2S( |
| in_channels=len(data_cfg["channels"]), |
| group_indices=groups, |
| hidden_channels=int(model_cfg.get("hidden_channels", 32)), |
| latent_channels=int(model_cfg.get("latent_channels", 64)), |
| patch_size=int(model_cfg.get("patch_size", 4)), |
| num_blocks=int(model_cfg.get("num_blocks", 2)), |
| perturbation_scale=float(model_cfg.get("perturbation_scale", 0.01)), |
| ) |
|
|
|
|
| def _load_checkpoint(path: Path, model: nn.Module, optimizer: torch.optim.Optimizer | None, device: torch.device) -> Dict[str, Any]: |
| state = torch.load(path, map_location=device, weights_only=False) |
| model.load_state_dict(state["model_state_dict"]) |
| if optimizer is not None and "optimizer_state_dict" in state: |
| optimizer.load_state_dict(state["optimizer_state_dict"]) |
| return state |
|
|
|
|
| def _epoch( |
| model: nn.Module, |
| loader, |
| device: torch.device, |
| optimizer: torch.optim.Optimizer | None, |
| rollout_steps: int, |
| distributed: bool, |
| probability_loss: bool, |
| kl_weight: float, |
| perturbation_enabled: bool, |
| ) -> float: |
| training = optimizer is not None |
| model.train(training) |
| total = 0.0 |
| batches = 0 |
| for inputs, targets, _ in loader: |
| with torch.set_grad_enabled(training): |
| state = inputs[:, -1].to(device=device, dtype=torch.float32, non_blocking=True) |
| targets = targets.to(device=device, dtype=torch.float32, non_blocking=True) |
| if optimizer is not None: |
| optimizer.zero_grad(set_to_none=True) |
| loss = state.new_zeros(()) |
| steps = max(1, min(int(rollout_steps), targets.shape[1])) |
| for step in range(steps): |
| prediction, aux = model( |
| state, |
| stochastic=training and perturbation_enabled, |
| return_aux=True, |
| ) |
| error = prediction - targets[:, step] |
| if probability_loss: |
| min_log_std = -6.0 |
| log_std = aux["log_std"].clamp(min_log_std, 2.0) |
| |
| |
| nll = ( |
| 0.5 * torch.exp(-2.0 * log_std) * error.square() |
| + log_std |
| - min_log_std |
| ) |
| loss = loss + torch.mean(nll) |
| if training and perturbation_enabled: |
| loss = loss + float(kl_weight) * aux["kl"] |
| else: |
| loss = loss + 0.8 * torch.mean(torch.abs(error)) + 0.2 * torch.mean(error.square()) |
| state = prediction |
| loss = loss / steps |
| loss_value = float(loss.detach().cpu()) |
| if not math.isfinite(loss_value) or loss_value < -1.0e-7: |
| raise RuntimeError(f"invalid loss {loss_value} on {device}") |
| if optimizer is not None: |
| loss.backward() |
| torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) |
| optimizer.step() |
| total += loss_value |
| batches += 1 |
| if batches == 0: |
| return float("nan") |
| value = torch.tensor(total / batches, device=device) |
| if distributed: |
| torch.distributed.all_reduce(value, op=torch.distributed.ReduceOp.SUM) |
| value = value / torch.distributed.get_world_size() |
| return float(value.cpu()) |
|
|
|
|
| def main() -> None: |
| args = parse_args() |
| with Path(args.config).open(encoding="utf-8") as source: |
| config = yaml.safe_load(source) |
| random.seed(args.seed) |
| np.random.seed(args.seed) |
| torch.manual_seed(args.seed) |
| distributed, rank, local_rank = _distributed_context() |
| device = _device(config, local_rank, args.device) |
| if distributed: |
| backend = "nccl" if device.type == "cuda" else "gloo" |
| torch.distributed.init_process_group(backend=backend, init_method="env://") |
| model_cfg = config["model"] |
| data_cfg = config["data"] |
| data_dir = resolve_data_dir(args.data_dir or data_cfg["data_dir"], PROJECT_ROOT) |
| input_steps = int(model_cfg.get("input_steps", 2)) |
| rollout_steps = int(args.rollout_steps or model_cfg.get("rollout_steps", 1)) |
| batch_size = int(args.batch_size or data_cfg.get("dataloader", {}).get("batch_size", 1)) |
| num_workers = int(args.num_workers if args.num_workers is not None else data_cfg.get("dataloader", {}).get("num_workers", 0)) |
| train_loader, train_sampler = make_dataloader( |
| data_dir, data_cfg["train_years"], data_cfg["channels"], input_steps, rollout_steps, |
| batch_size, num_workers, distributed, True, data_cfg.get("dataloader", {}).get("pin_memory", False), |
| ) |
| val_loader, val_sampler = make_dataloader( |
| data_dir, data_cfg["val_years"], data_cfg["channels"], input_steps, rollout_steps, |
| batch_size, num_workers, distributed, False, data_cfg.get("dataloader", {}).get("pin_memory", False), |
| ) |
| model = _model_from_config(config).to(device) |
| if not args.finetune: |
| for parameter in model.perturbation.parameters(): |
| parameter.requires_grad_(False) |
| parameter_device = next(model.parameters()).device |
| if parameter_device != device: |
| raise RuntimeError(f"model parameters are on {parameter_device}, expected {device}") |
| learning_rate = float(model_cfg.get("learning_rate", 2e-4)) |
| if args.finetune: |
| learning_rate = float(model_cfg.get("finetune_learning_rate", learning_rate * 0.1)) |
| optimizer = torch.optim.AdamW( |
| model.parameters(), |
| lr=learning_rate, |
| betas=( |
| float(model_cfg.get("adam_beta1", 0.9)), |
| float(model_cfg.get("adam_beta2", 0.999)), |
| ), |
| weight_decay=float(model_cfg.get("weight_decay", 1e-5)), |
| ) |
| scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode="min", factor=0.5, patience=3) |
| checkpoint_dir = resolve_data_dir(model_cfg.get("checkpoint_dir", "./data/checkpoints"), PROJECT_ROOT) |
| checkpoint_dir.mkdir(parents=True, exist_ok=True) |
| checkpoint_path = Path(args.checkpoint) if args.checkpoint else checkpoint_dir / "model_bak.pth" |
| if not checkpoint_path.is_absolute(): |
| checkpoint_path = (PROJECT_ROOT / checkpoint_path).resolve() |
| start_epoch = 0 |
| best_valid = float("inf") |
| train_losses: list[float] = [] |
| valid_losses: list[float] = [] |
| if args.finetune: |
| if not checkpoint_path.exists(): |
| raise FileNotFoundError( |
| f"Finetune checkpoint not found: {checkpoint_path}; run base training first" |
| ) |
| state = _load_checkpoint(checkpoint_path, model, optimizer, device) |
| start_epoch = int(state.get("epoch", -1)) + 1 |
| best_valid = float(state.get("best_valid_loss", best_valid)) |
| |
| |
| for group in optimizer.param_groups: |
| group["lr"] = learning_rate |
| if distributed: |
| model = DDP(model, device_ids=[local_rank] if device.type == "cuda" else None) |
| requested_epochs = int(args.max_epoch or (model_cfg.get("finetune_epoch") if args.finetune else model_cfg.get("max_epoch", 1))) |
| |
| |
| |
| max_epoch = start_epoch + requested_epochs if args.finetune else requested_epochs |
| patience = int(model_cfg.get("patience", 10)) |
| probability_loss = bool(model_cfg.get("probability_loss", True)) |
| kl_weight = float(model_cfg.get("kl_weight", 1.0e-4)) |
| stale = 0 |
| if rank == 0: |
| cuda_state = { |
| "available": bool(torch.cuda.is_available()), |
| "count": int(torch.cuda.device_count()) if torch.cuda.is_available() else 0, |
| "name": torch.cuda.get_device_name(device) if device.type == "cuda" else None, |
| "allocated_bytes": torch.cuda.memory_allocated(device) if device.type == "cuda" else 0, |
| } |
| print( |
| f"FengWu-W2S training on {device}; model_device={parameter_device}; " |
| f"cuda={cuda_state}; samples={len(train_loader.dataset)}; rollout={rollout_steps}", |
| flush=True, |
| ) |
| print(f"Parameters: {sum(parameter.numel() for parameter in model.parameters()):,}", flush=True) |
| for epoch in range(start_epoch, max_epoch): |
| if distributed: |
| train_sampler.set_epoch(epoch) |
| val_sampler.set_epoch(epoch) |
| train_loss = _epoch( |
| model, train_loader, device, optimizer, rollout_steps, distributed, |
| probability_loss, kl_weight, args.finetune, |
| ) |
| valid_loss = _epoch( |
| model, val_loader, device, None, rollout_steps, distributed, |
| probability_loss, kl_weight, False, |
| ) |
| scheduler.step(valid_loss) |
| train_losses.append(train_loss) |
| valid_losses.append(valid_loss) |
| improved = valid_loss < best_valid |
| if improved: |
| best_valid = valid_loss |
| stale = 0 |
| else: |
| stale += 1 |
| if rank == 0: |
| model_to_save = model.module if hasattr(model, "module") else model |
| state = { |
| "model_state_dict": model_to_save.state_dict(), |
| "optimizer_state_dict": optimizer.state_dict(), |
| "scheduler_state_dict": scheduler.state_dict(), |
| "epoch": epoch, |
| "best_valid_loss": best_valid, |
| "channels": data_cfg["channels"], |
| "group_indices": data_cfg["groups"], |
| "config": config, |
| } |
| if improved or epoch == max_epoch - 1: |
| torch.save(state, checkpoint_dir / "model_bak.pth") |
| np.save(checkpoint_dir / "trloss.npy", np.asarray(train_losses, dtype=np.float32)) |
| np.save(checkpoint_dir / "valoss.npy", np.asarray(valid_losses, dtype=np.float32)) |
| print( |
| f"Epoch {epoch + 1}/{max_epoch}: train={train_loss:.6f} valid={valid_loss:.6f}", |
| flush=True, |
| ) |
| if stale > patience: |
| break |
| if distributed: |
| torch.distributed.barrier() |
| torch.distributed.destroy_process_group() |
| if rank == 0: |
| (checkpoint_dir / "training_summary.json").write_text( |
| json.dumps({"best_valid_loss": best_valid, "epochs": len(train_losses), "device": str(device)}, indent=2), |
| encoding="utf-8", |
| ) |
| print(f"Saved checkpoint to {checkpoint_dir / 'model_bak.pth'}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|