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