FourCastNet_v2 / scripts /train.py
yzt15806542928's picture
Upload folder using huggingface_hub
eca4864 verified
Raw
History Blame Contribute Delete
10.8 kB
from __future__ import annotations
import argparse
import json
import os
import sys
import tempfile
from contextlib import nullcontext
from pathlib import Path
from typing import Any, Iterable
import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel
SCRIPT_DIR = Path(__file__).resolve().parent
MODEL_DIR = SCRIPT_DIR.parent / "model"
for module_dir in (SCRIPT_DIR, MODEL_DIR):
if str(module_dir) not in sys.path:
sys.path.insert(0, str(module_dir))
from common import (
DEFAULT_CONFIG,
active_model_config,
load_config,
resolve_path,
seed_everything,
)
from data import build_loader
from fourcastnet_v2 import (
FourCastNetV2,
load_checkpoint,
)
def initialize_distributed(backend: str) -> tuple[torch.device, int, int, int]:
world_size = int(os.environ.get("WORLD_SIZE", "1"))
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
if torch.cuda.is_available():
device_count = torch.cuda.device_count()
if not 0 <= local_rank < device_count:
raise RuntimeError(
f"LOCAL_RANK={local_rank} is not available; "
f"this process can see {device_count} CUDA devices "
f"(CUDA_VISIBLE_DEVICES={os.environ.get('CUDA_VISIBLE_DEVICES', '<unset>')})"
)
# Select the rank-local device before NCCL initialization. Otherwise
# every process starts with the default device (usually cuda:0).
torch.cuda.set_device(local_rank)
device = torch.device("cuda", local_rank)
else:
device = torch.device("cpu")
if world_size > 1 and not dist.is_initialized():
dist.init_process_group(backend=backend, init_method="env://")
rank = dist.get_rank() if dist.is_initialized() else 0
return device, rank, local_rank, world_size
def spherical_relative_l2(prediction: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
height = target.shape[-2]
latitude = torch.linspace(
torch.pi / 2,
-torch.pi / 2,
height,
device=target.device,
dtype=target.dtype,
)
weights = torch.cos(latitude).clamp_min(0)
weights = weights / weights.mean()
weights = weights.view(1, 1, height, 1)
error = ((prediction - target).square() * weights).sum(dim=(-2, -1))
reference = (target.square() * weights).sum(dim=(-2, -1)).clamp_min(1e-12)
return torch.sqrt(error / reference).mean()
def autoregressive_loss(
model: torch.nn.Module,
inputs: torch.Tensor,
targets: torch.Tensor,
steps: int,
) -> torch.Tensor:
if steps == 1:
targets = targets if targets.ndim == 4 else targets[:, 0]
elif targets.ndim != 5 or targets.shape[1] != steps:
raise ValueError(f"Expected targets [B,{steps},C,H,W], got {targets.shape}")
state = inputs
losses = []
for step in range(steps):
state = model(state)
target = targets if steps == 1 else targets[:, step]
losses.append(spherical_relative_l2(state, target))
return torch.stack(losses).mean()
def limited_batches(loader: Iterable, maximum: int | None):
for index, batch in enumerate(loader):
if maximum is not None and index >= maximum:
break
yield batch
def reduce_average(total: float, count: int, device: torch.device) -> float:
values = torch.tensor([total, count], dtype=torch.float64, device=device)
if dist.is_initialized():
dist.all_reduce(values, op=dist.ReduceOp.SUM)
return (values[0] / values[1].clamp_min(1)).item()
def run_epoch(
model: torch.nn.Module,
loader,
device: torch.device,
*,
steps: int,
optimizer: torch.optim.Optimizer | None,
amp: bool,
max_batches: int | None,
max_grad_norm: float,
) -> float:
training = optimizer is not None
model.train(training)
total = 0.0
count = 0
context = nullcontext if training else torch.no_grad
with context():
for batch in limited_batches(loader, max_batches):
inputs = batch[0].to(device, non_blocking=True)
targets = batch[1].to(device, non_blocking=True)
if training:
optimizer.zero_grad(set_to_none=True)
with torch.autocast(
device_type=device.type,
dtype=torch.float16,
enabled=amp and device.type == "cuda",
):
loss = autoregressive_loss(model, inputs, targets, steps)
if training:
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_grad_norm)
optimizer.step()
total += loss.detach().item()
count += 1
return reduce_average(total, count, device)
def save_checkpoint_atomic(path: Path, state: dict[str, Any]) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with tempfile.NamedTemporaryFile(dir=path.parent, delete=False) as stream:
temporary_path = Path(stream.name)
try:
torch.save(state, temporary_path)
os.replace(temporary_path, path)
finally:
temporary_path.unlink(missing_ok=True)
def main() -> None:
parser = argparse.ArgumentParser(description="Train FourCastNet v2")
parser.add_argument("--config", default=str(DEFAULT_CONFIG))
parser.add_argument("--stage", choices=("one_step", "finetune"))
parser.add_argument("--resume")
args = parser.parse_args()
config = load_config(args.config)
seed_everything(config["project"]["seed"])
training = config["training"]
stage = args.stage or training["stage"]
if stage == "one_step" and args.resume:
raise ValueError("One-step training always starts from random initialization")
steps = 1 if stage == "one_step" else training["finetune"]["autoregressive_steps"]
epochs = training["epochs"] if stage == "one_step" else training["finetune"]["epochs"]
learning_rate = (
training["learning_rate"]
if stage == "one_step"
else training["finetune"]["learning_rate"]
)
device, rank, local_rank, world_size = initialize_distributed(
config["distributed"]["backend"]
)
print(
f"rank={rank}/{world_size} local_rank={local_rank} "
f"device={device} visible_devices={torch.cuda.device_count()}",
flush=True,
)
train_loader, train_sampler = build_loader(
config,
config["data"]["train_years"],
train=True,
distributed=world_size > 1,
output_steps=steps,
)
val_loader, val_sampler = build_loader(
config,
config["data"]["val_years"],
train=False,
distributed=world_size > 1,
output_steps=steps,
)
model = FourCastNetV2(active_model_config(config)).to(device)
if stage == "finetune":
resume_path = args.resume or config["checkpoint"]["finetune_from"]
result = load_checkpoint(
model,
resolve_path(config, resume_path),
expected_profile=config["model"]["profile"],
expected_variables=config["data"]["variables"],
allowed_stages={"one_step"},
allowed_initializations={"random"},
strict=config["checkpoint"]["strict"],
map_location=device,
)
optimizer = torch.optim.AdamW(
model.parameters(),
lr=learning_rate,
betas=tuple(training["optimizer_betas"]),
weight_decay=training["weight_decay"],
)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)
if world_size > 1:
ddp_options = (
{
"device_ids": [device.index],
"output_device": device.index,
# SFNO's SHT buffers are immutable coefficients. DDP's
# per-forward buffer broadcast mutates them in-place and
# invalidates the graph in multi-step autoregressive training.
"broadcast_buffers": False,
}
if device.type == "cuda"
else {"broadcast_buffers": False}
)
model = DistributedDataParallel(model, **ddp_options)
checkpoint_dir = (
resolve_path(config, config["project"]["checkpoint_dir"]) / stage
)
checkpoint_prefix = config["checkpoint"].get("prefix", "model_bak")
best_loss = float("inf")
history = []
for epoch in range(epochs):
if train_sampler is not None:
train_sampler.set_epoch(epoch)
if val_sampler is not None:
val_sampler.set_epoch(epoch)
train_loss = run_epoch(
model,
train_loader,
device,
steps=steps,
optimizer=optimizer,
amp=training["amp"],
max_batches=training["max_train_batches"],
max_grad_norm=training["max_grad_norm"],
)
val_loss = run_epoch(
model,
val_loader,
device,
steps=steps,
optimizer=None,
amp=training["amp"],
max_batches=training["max_val_batches"],
max_grad_norm=training["max_grad_norm"],
)
scheduler.step()
history.append({"epoch": epoch, "train_loss": train_loss, "val_loss": val_loss})
if rank == 0:
print(
f"epoch={epoch + 1}/{epochs} train_loss={train_loss:.6f} "
f"val_loss={val_loss:.6f}"
)
raw_model = model.module if hasattr(model, "module") else model
state = {
"checkpoint_format": "fourcastnet_v2_project",
"scratch_lineage": True,
"model_state_dict": raw_model.state_dict(),
"optimizer_state_dict": optimizer.state_dict(),
"scheduler_state_dict": scheduler.state_dict(),
"epoch": epoch,
"stage": stage,
"initialization": (
"random" if stage == "one_step" else "one_step_checkpoint"
),
"model_profile": config["model"]["profile"],
"variables": config["data"]["variables"],
}
save_checkpoint_atomic(checkpoint_dir / f"{checkpoint_prefix}_last.pt", state)
if val_loss < best_loss:
best_loss = val_loss
save_checkpoint_atomic(checkpoint_dir / f"{checkpoint_prefix}.pt", state)
checkpoint_dir.mkdir(parents=True, exist_ok=True)
(checkpoint_dir / "history.json").write_text(
json.dumps(history, indent=2), encoding="utf-8"
)
if dist.is_initialized():
dist.destroy_process_group()
if __name__ == "__main__":
main()