"""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()