import json import os from pathlib import Path import sys import numpy as np import torch import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel from torch.utils.data import DataLoader, DistributedSampler, TensorDataset import yaml ROOT = Path(__file__).resolve().parents[1] if str(ROOT) not in sys.path: sys.path.insert(0, str(ROOT)) from model.stablenn_phys import StableNNPhys, rollout, rollout_loss CHECKPOINT_FORMAT_VERSION = "stablenn_phys_checkpoint_v1" def main(): cfg = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) torch.manual_seed(cfg["seed"]) distributed = int(os.environ.get("WORLD_SIZE", "1")) > 1 world_size = int(os.environ.get("WORLD_SIZE", "1")) use_accelerator = torch.cuda.is_available() and torch.cuda.device_count() >= world_size if distributed: dist.init_process_group("nccl" if use_accelerator else "gloo") rank = dist.get_rank() if distributed else 0 local_rank = int(os.environ.get("LOCAL_RANK", "0")) device = torch.device(f"cuda:{local_rank}" if use_accelerator else "cpu") if use_accelerator: torch.cuda.set_device(device) raw = np.load(ROOT / cfg["data"]["file"]) tensors = [torch.from_numpy(raw[name].astype(np.float32)) for name in ("initial", "target", "surface", "advection")] source = torch.from_numpy(raw["source"].astype(np.int64)) dataset = TensorDataset(*tensors, source) sampler = DistributedSampler(dataset, shuffle=True, seed=cfg["seed"]) if distributed else None loader = DataLoader(dataset, batch_size=cfg["train"]["batch_size"], sampler=sampler, shuffle=sampler is None) all_state = torch.from_numpy(raw["target"].astype(np.float32)) state_mean = all_state.mean((0, 1)).to(device) state_std = all_state.std((0, 1), unbiased=False).clamp_min(1e-5).to(device) tendency_mean = torch.zeros(68, device=device) tendency_std = torch.cat((torch.full((34,), 0.02), torch.full((34,), 2e-7))).to(device) layer_mass = torch.from_numpy(raw["layer_mass"].astype(np.float32)).to(device) model = StableNNPhys(cfg["model"]["hidden_size"]).to(device) if distributed: model = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None) optimizer = torch.optim.Adam(model.parameters(), lr=cfg["train"]["learning_rate"]) history = [] for epoch in range(cfg["train"]["epochs"]): if sampler: sampler.set_epoch(epoch) total = 0.0 for initial, target, surface, advection, _ in loader: initial, target, surface, advection = [x.to(device) for x in (initial, target, surface, advection)] prediction, _ = rollout(model, initial, surface, advection, state_mean, state_std, tendency_mean, tendency_std, float(raw["dt_seconds"])) loss = rollout_loss(prediction, target, layer_mass.expand(initial.shape[0], -1), cfg["train"]["loss_mode"]) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() total += loss.item() * initial.shape[0] value = total / len(dataset) history.append(value) if rank == 0: print(f"epoch={epoch + 1} loss={value:.6g}") if rank == 0: module = model.module if distributed else model checkpoint = ROOT / cfg["paths"]["checkpoint"] checkpoint.parent.mkdir(parents=True, exist_ok=True) payload = {"format_version": CHECKPOINT_FORMAT_VERSION, "model": module.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "epoch": cfg["train"]["epochs"], "model_config": dict(cfg["model"]), "training_config": dict(cfg["train"]), "normalization": {"state_mean": state_mean.cpu(), "state_std": state_std.cpu(), "tendency_mean": tendency_mean.cpu(), "tendency_std": tendency_std.cpu()}, "variables": {"input": ["sL[34]", "qT[34]", "SHF", "LHF", "SOLIN"], "output": ["dsL_dt[34]", "dqT_dt[34]"]}, "dt_seconds": float(raw["dt_seconds"])} torch.save(payload, checkpoint) metrics = ROOT / cfg["paths"]["training_metrics"] metrics.parent.mkdir(parents=True, exist_ok=True) metrics.write_text(json.dumps({"loss": history, "world_size": dist.get_world_size() if distributed else 1, "paper_model": cfg["train"]["paper_model"]}, indent=2)) print(f"saved {checkpoint}") if distributed: dist.destroy_process_group() if __name__ == "__main__": main()