Samudra / scripts /train.py
yzt15806542928's picture
Upload folder using huggingface_hub
929e312 verified
Raw
History Blame Contribute Delete
8.39 kB
"""Single-device and torchrun-based distributed training for Samudra v1."""
try:
from ._bootstrap import ROOT
except ImportError:
from _bootstrap import ROOT
import argparse
import os
import random
from pathlib import Path
import numpy as np
import torch
import yaml
from torch import nn
from torch.nn.parallel import DistributedDataParallel
from torch.utils.data import DataLoader, Dataset, DistributedSampler
from model.samudra import build_model
STATE_CHANNELS = 77
BOUNDARY_CHANNELS = 4
def load_data(path: str | Path) -> tuple[np.ndarray, np.ndarray]:
"""Load and validate native Samudra time-major arrays."""
with np.load(path) as data:
prognostic = np.asarray(data["prognostic"], dtype=np.float32)
boundary = np.asarray(data["boundary"], dtype=np.float32)
if prognostic.ndim != 4 or prognostic.shape[1] != STATE_CHANNELS:
raise ValueError("prognostic must have shape [time, 77, lat, lon]")
if boundary.ndim != 4 or boundary.shape[1] != BOUNDARY_CHANNELS:
raise ValueError("boundary must have shape [time, 4, lat, lon]")
if prognostic.shape[0] != boundary.shape[0] or prognostic.shape[2:] != boundary.shape[2:]:
raise ValueError("prognostic and boundary time/grid dimensions must match")
return prognostic, boundary
class SamudraDataset(Dataset):
"""Build recurrent training windows from native Samudra arrays."""
def __init__(self, path: str | Path, recurrent_passes: int):
self.prognostic, self.boundary = load_data(path)
self.recurrent_passes = recurrent_passes
self.end = self.prognostic.shape[0] - 2 * recurrent_passes
if self.end <= 1:
raise ValueError("dataset does not contain enough samples for recurrent training")
def __len__(self) -> int:
return self.end - 1
def __getitem__(self, index: int) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
t = index + 1
history = np.stack((self.prognostic[t - 1], self.prognostic[t]))
forcing = self.boundary[t : t + self.recurrent_passes]
labels = np.stack(
[
np.concatenate(
(self.prognostic[t + 2 * step + 1], self.prognostic[t + 2 * step + 2])
)
for step in range(self.recurrent_passes)
]
)
return torch.from_numpy(history), torch.from_numpy(forcing), torch.from_numpy(labels)
def set_seed(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
def load_config(path: str) -> dict:
with open(path, encoding="utf-8") as handle:
return yaml.safe_load(handle)
def distributed_setup(device: torch.device) -> tuple[int, int, bool]:
world_size = int(os.environ.get("WORLD_SIZE", "1"))
rank = int(os.environ.get("RANK", "0"))
distributed = world_size > 1
if distributed:
backend = "nccl" if device.type == "cuda" else "gloo"
torch.distributed.init_process_group(backend=backend)
return rank, world_size, distributed
def save_checkpoint(path: Path, model: nn.Module, optimizer: torch.optim.Optimizer, scheduler, epoch: int, loss: float) -> None:
state = model.module.state_dict() if isinstance(model, DistributedDataParallel) else model.state_dict()
torch.save({"epoch": epoch, "loss": loss, "model": state, "optimizer": optimizer.state_dict(), "scheduler": scheduler.state_dict()}, path)
def freeze_batch_norm_stats(module: nn.Module) -> None:
"""Prevent recurrent forwards from mutating BatchNorm buffers in one graph."""
for child in module.modules():
if isinstance(child, nn.modules.batchnorm._BatchNorm):
child.eval()
def train(
config_path: str,
data_path: str,
device_name: str | None = None,
epochs_override: int | None = None,
output_dir_override: str | None = None,
) -> None:
config = load_config(config_path)
rank = int(os.environ.get("RANK", "0"))
local_rank = int(os.environ.get("LOCAL_RANK", rank))
if device_name:
device = torch.device(device_name)
elif torch.cuda.is_available():
device = torch.device(f"cuda:{local_rank}")
else:
device = torch.device("cpu")
rank, world_size, distributed = distributed_setup(device)
set_seed(int(config["project"].get("seed", 1)) + rank)
if device.type == "cuda":
torch.cuda.set_device(device)
recurrent_passes = int(config["data"].get("recurrent_passes", 1))
dataset = SamudraDataset(data_path, recurrent_passes)
sampler = DistributedSampler(dataset, num_replicas=world_size, rank=rank, shuffle=True) if distributed else None
loader = DataLoader(dataset, batch_size=config["training"]["batch_size"], shuffle=sampler is None, sampler=sampler, num_workers=config["training"].get("num_workers", 0), pin_memory=device.type == "cuda")
model = build_model(config).to(device)
if distributed:
model = DistributedDataParallel(
model,
device_ids=[device.index] if device.type == "cuda" else None,
broadcast_buffers=False,
)
optimizer = torch.optim.Adam(model.parameters(), lr=config["training"]["learning_rate"], weight_decay=config["training"].get("weight_decay", 0.0))
epochs = epochs_override if epochs_override is not None else int(config["training"]["epochs"])
if epochs < 1:
raise ValueError("epochs must be positive")
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)
resume = config["training"].get("resume_checkpoint")
start_epoch = 0
if resume:
try:
checkpoint = torch.load(resume, map_location=device, weights_only=True)
except TypeError:
checkpoint = torch.load(resume, map_location=device)
target = model.module if isinstance(model, DistributedDataParallel) else model
state = {key: value for key, value in checkpoint["model"].items() if not key.endswith(".cap")}
target.load_state_dict(state)
optimizer.load_state_dict(checkpoint["optimizer"])
scheduler.load_state_dict(checkpoint["scheduler"])
start_epoch = int(checkpoint["epoch"]) + 1
output_dir = Path(output_dir_override or config["training"]["output_dir"])
output_dir.mkdir(parents=True, exist_ok=True)
for epoch in range(start_epoch, epochs):
if sampler is not None:
sampler.set_epoch(epoch)
model.train()
freeze_batch_norm_stats(model)
total_loss = 0.0
for batch in loader:
optimizer.zero_grad(set_to_none=True)
history, forcing, labels = (item.to(device, non_blocking=True) for item in batch)
previous, current = history[:, 0], history[:, 1]
losses = []
for step in range(recurrent_passes):
prediction = model(torch.cat((previous, current, forcing[:, step]), dim=1))
losses.append(nn.functional.mse_loss(prediction, labels[:, step]))
previous, current = prediction[:, :STATE_CHANNELS], prediction[:, STATE_CHANNELS:]
loss = torch.stack(losses).mean()
loss.backward()
optimizer.step()
total_loss += float(loss.detach())
scheduler.step()
mean_loss = total_loss / max(1, len(loader))
if rank == 0:
print(f"epoch={epoch + 1} loss={mean_loss:.6e} lr={scheduler.get_last_lr()[0]:.6e}")
frequency = int(config["training"].get("save_frequency", 5))
if (epoch + 1) % frequency == 0 or epoch + 1 == epochs:
save_checkpoint(output_dir / f"epoch_{epoch + 1:04d}.pt", model, optimizer, scheduler, epoch, mean_loss)
save_checkpoint(output_dir / "model_bak.pth", model, optimizer, scheduler, epoch, mean_loss)
if distributed:
torch.distributed.destroy_process_group()
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--config", default="./conf/config.yaml")
parser.add_argument("--data", default="./data/train.npz")
parser.add_argument("--device", default=None)
parser.add_argument("--epochs", type=int, default=None)
parser.add_argument("--output-dir", default=None)
args = parser.parse_args()
train(args.config, args.data, args.device, args.epochs, args.output_dir)
if __name__ == "__main__":
main()