GenCast / scripts /train.py
yzt15806542928's picture
Upload folder using huggingface_hub
7180154 verified
Raw
History Blame Contribute Delete
12.8 kB
#!/usr/bin/env python3
"""使用官方单步 EDM 去噪目标训练 GenCast。"""
from __future__ import annotations
import argparse
import itertools
import sys
import warnings
from pathlib import Path
import xarray
# Mesh adjacency construction triggers one-time scipy CSR restructure warning.
warnings.filterwarnings("ignore", message="Changing the sparsity structure")
PROJECT_ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(PROJECT_ROOT))
from model.common import configure_jax, load_config, load_stats, resolve_path
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--config", default=str(PROJECT_ROOT / "conf/config.yaml"))
parser.add_argument("--max-steps", type=int)
parser.add_argument("--resume")
parser.add_argument("--parallel-mode", choices=("single", "pmap"))
parser.add_argument("--num-devices", type=int)
parser.add_argument("--global-batch-size", type=int)
parser.add_argument("--checkpoint")
parser.add_argument("--seed", type=int)
return parser.parse_args()
def _adam_init(params):
import jax
import jax.numpy as jnp
zeros = jax.tree_util.tree_map(jnp.zeros_like, params)
return {"count": jnp.asarray(0, dtype=jnp.int32), "mu": zeros, "nu": zeros}
def _adam_update(params, grads, state, learning_rate, beta1, beta2, eps):
import jax
import jax.numpy as jnp
count = state["count"] + 1
mu = jax.tree_util.tree_map(
lambda old, grad: beta1 * old + (1.0 - beta1) * grad,
state["mu"], grads,
)
nu = jax.tree_util.tree_map(
lambda old, grad: beta2 * old + (1.0 - beta2) * jnp.square(grad),
state["nu"], grads,
)
mu_hat = jax.tree_util.tree_map(lambda value: value / (1.0 - beta1**count), mu)
nu_hat = jax.tree_util.tree_map(lambda value: value / (1.0 - beta2**count), nu)
params = jax.tree_util.tree_map(
lambda value, first, second: value - learning_rate * first / (jnp.sqrt(second) + eps),
params, mu_hat, nu_hat,
)
return params, {"count": count, "mu": mu, "nu": nu}
def _replicate(tree, devices):
import jax
return jax.device_put_replicated(tree, devices)
def _unreplicate(tree):
import jax
return jax.tree_util.tree_map(lambda value: value[0], tree)
def _device_batch(batch, device_count):
"""Add a leading device dimension to each GenCast xarray input."""
result = []
for value in batch:
if not isinstance(value, xarray.Dataset):
raise TypeError("GenCast batches must contain xarray.Dataset values")
value = value.transpose("batch", ...)
if "batch" not in value.dims:
value = value.expand_dims("batch")
if value.sizes["batch"] % device_count:
raise ValueError("Batch size must be divisible by the device count")
local_batch = value.sizes["batch"] // device_count
shards = [
value.isel(batch=slice(index * local_batch, (index + 1) * local_batch))
for index in range(device_count)
]
result.append(xarray.concat(shards, dim="device"))
return tuple(result)
def main() -> None:
args = parse_args()
config = load_config(args.config)
parallel = config.setdefault("parallel", {})
if args.parallel_mode is not None:
parallel["mode"] = args.parallel_mode
if args.num_devices is not None:
parallel["num_devices"] = args.num_devices
if args.global_batch_size is not None:
parallel["global_batch_size"] = args.global_batch_size
if args.checkpoint is not None:
config["checkpoint"]["trainer"] = args.checkpoint
if args.seed is not None:
config["training"]["seed"] = args.seed
configure_jax(config["runtime"].get("platform", "auto"))
import jax
import jax.numpy as jnp
from model.gencast import GenCastModel, parameter_count
from model.common import (
load_trainer_checkpoint, save_trainer_checkpoint,
validate_checkpoint_config,
)
from model.data_loader import GenCastERA5Dataset, batch_iterator
mode = str(parallel.get("mode", "single")).lower()
if mode not in ("single", "pmap"):
raise ValueError("parallel.mode must be 'single' or 'pmap'")
devices = list(jax.local_devices())
requested_devices = int(parallel.get("num_devices", 1))
if requested_devices < 1:
raise ValueError("parallel.num_devices must be positive")
if mode == "pmap":
if requested_devices > len(devices):
raise ValueError(
f"Requested {requested_devices} devices, only {len(devices)} available"
)
devices = devices[:requested_devices]
else:
requested_devices = 1
devices = devices[:1]
global_batch_size = int(parallel.get("global_batch_size", requested_devices))
if global_batch_size < 1 or global_batch_size % requested_devices:
raise ValueError("global_batch_size must be divisible by the device count")
stats = load_stats(config["data"]["stats_dir"])
model = GenCastModel.from_config_and_stats(config, stats)
dataset = GenCastERA5Dataset(
resolve_path(config["data"]["data_dir"]),
list(config["data"]["train_years"]),
static_dir=resolve_path(config["data"]["static_dir"]),
prediction_steps=1,
stride=int(config["data"].get("train_stride", 1)),
precipitation_interval_hours=int(
config["data"]["precipitation_interval_hours"]
),
)
first_batch = dataset[0]
seed = int(config["training"]["seed"])
start_step = 0
resume = args.resume or config["checkpoint"].get("resume")
if resume:
params, state, optimizer_state, start_step, saved_config = \
load_trainer_checkpoint(resume)
validate_checkpoint_config(config, saved_config)
else:
params, state = model.init(
jax.random.fold_in(jax.random.PRNGKey(seed), -1), *first_batch
)
optimizer_state = _adam_init(params)
learning_rate = float(config["training"]["learning_rate"])
beta1, beta2 = (float(value) for value in config["training"]["betas"])
epsilon = float(config["training"].get("epsilon", 1e-8))
def train_step(params, state, optimizer_state, rng, inputs, targets, forcings):
def objective(current_params, current_state):
(loss, diagnostics), next_state = model.loss(
current_params, current_state, rng, inputs, targets, forcings
)
return loss, (diagnostics, next_state)
(loss, (diagnostics, next_state)), grads = jax.value_and_grad(
objective, has_aux=True
)(params, state)
finite = jnp.logical_and(
jnp.isfinite(loss),
jnp.all(jnp.asarray([jnp.all(jnp.isfinite(x)) for x in jax.tree_util.tree_leaves(grads)])),
)
new_params, new_optimizer_state = _adam_update(
params, grads, optimizer_state, learning_rate, beta1, beta2, epsilon
)
params = jax.tree_util.tree_map(
lambda new, old: jnp.where(finite, new, old), new_params, params
)
next_state = jax.tree_util.tree_map(
lambda new, old: jnp.where(finite, new, old), next_state, state
)
new_optimizer_state = jax.tree_util.tree_map(
lambda new, old: jnp.where(finite, new, old),
new_optimizer_state,
optimizer_state,
)
return params, next_state, new_optimizer_state, loss, diagnostics, finite
if mode == "pmap":
axis_name = str(parallel.get("axis_name", "devices"))
def parallel_train_step(
params, state, optimizer_state, rng, inputs, targets, forcings
):
rng = jax.random.fold_in(rng, jax.lax.axis_index(axis_name))
def objective(current_params, current_state):
(loss, diagnostics), next_state = model.loss(
current_params, current_state, rng, inputs, targets, forcings
)
return loss, (diagnostics, next_state)
(loss, (diagnostics, next_state)), grads = jax.value_and_grad(
objective, has_aux=True
)(params, state)
grads = jax.lax.pmean(grads, axis_name)
loss = jax.lax.pmean(loss, axis_name)
diagnostics = jax.tree_util.tree_map(
lambda value: jax.lax.pmean(value, axis_name), diagnostics
)
next_state = jax.tree_util.tree_map(
lambda value: jax.lax.pmean(value, axis_name), next_state
)
finite = jnp.logical_and(
jnp.isfinite(loss),
jnp.all(jnp.asarray([
jnp.all(jnp.isfinite(x))
for x in jax.tree_util.tree_leaves(grads)
])),
)
finite = jax.lax.pmin(finite, axis_name)
new_params, new_optimizer_state = _adam_update(
params, grads, optimizer_state, learning_rate, beta1, beta2, epsilon
)
params = jax.tree_util.tree_map(
lambda new, old: jnp.where(finite, new, old), new_params, params
)
next_state = jax.tree_util.tree_map(
lambda new, old: jnp.where(finite, new, old), next_state, state
)
new_optimizer_state = jax.tree_util.tree_map(
lambda new, old: jnp.where(finite, new, old),
new_optimizer_state,
optimizer_state,
)
return params, next_state, new_optimizer_state, loss, diagnostics, finite
from model.graphcast import xarray_jax
train_step = xarray_jax.pmap(
parallel_train_step, dim="device", axis_name=axis_name, devices=devices
)
else:
train_step = jax.jit(train_step)
max_steps = int(args.max_steps or config["training"]["max_steps"])
save_interval = int(config["training"].get("save_interval", max_steps))
checkpoint_path = config["checkpoint"]["trainer"]
print(f"Training samples: {len(dataset)}; parameters: {parameter_count(params):,}")
if mode == "pmap":
params = _replicate(params, devices)
state = _replicate(state, devices)
optimizer_state = _replicate(optimizer_state, devices)
print(
f"Parallel mode: pmap; devices: {requested_devices}; "
f"global batch: {global_batch_size}"
)
step = start_step
batches_per_epoch = len(dataset) // global_batch_size
if batches_per_epoch < 1:
raise ValueError(
f"Dataset has {len(dataset)} samples, fewer than global_batch_size "
f"{global_batch_size}"
)
while step < max_steps:
epoch = step // batches_per_epoch
offset = step % batches_per_epoch
epoch_batches = batch_iterator(
dataset,
shuffle=True,
seed=seed + epoch,
batch_size=global_batch_size,
)
for batch in itertools.islice(epoch_batches, offset, None):
if step >= max_steps:
break
step_rng = jax.random.fold_in(jax.random.PRNGKey(seed), step)
if mode == "pmap":
batch = _device_batch(batch, requested_devices)
step_rng = jax.numpy.broadcast_to(
step_rng, (requested_devices, *step_rng.shape)
)
params, state, optimizer_state, loss, _, finite = train_step(
params, state, optimizer_state, step_rng, *batch
)
loss, finite = loss[0], finite[0]
else:
params, state, optimizer_state, loss, _, finite = train_step(
params, state, optimizer_state, step_rng, *batch
)
step += 1
print(f"step={step} loss={float(loss):.8f} finite={bool(finite)}")
if not bool(finite):
raise FloatingPointError(f"Non-finite GenCast loss at step {step}")
if step % save_interval == 0 or step == max_steps:
checkpoint_trees = (params, state, optimizer_state)
if mode == "pmap":
checkpoint_trees = tuple(map(_unreplicate, checkpoint_trees))
save_trainer_checkpoint(
checkpoint_path,
params=checkpoint_trees[0],
state=checkpoint_trees[1],
optimizer_state=checkpoint_trees[2],
step=step,
config=config,
)
print(f"Saved checkpoint to {resolve_path(checkpoint_path)}")
if __name__ == "__main__":
main()