"""Train Surya in one-step pretraining and rollout-tuning phases.""" import argparse, importlib.util, json, math, os, random from contextlib import nullcontext from pathlib import Path import numpy as np, torch, yaml from torch import distributed as dist from torch.nn.parallel import DistributedDataParallel from torch.utils.data import DataLoader, Dataset, DistributedSampler ROOT = Path(__file__).resolve().parents[1] class SolarDataset(Dataset): def __init__(self, path, cfg): d = np.load(path); self.inputs, self.targets = d["inputs"], d["targets"] if self.inputs.ndim != 5 or self.targets.ndim != 5 or self.inputs.shape[2] != 13 or self.inputs.shape[1] != 2: raise ValueError("Expected inputs [N,2,13,H,W] and targets [N,S,13,H,W]") self.mean = np.asarray(cfg["data"]["channel_mean"], dtype=np.float32)[None, :, None, None] self.std = np.asarray(cfg["data"]["channel_std"], dtype=np.float32)[None, :, None, None] def __len__(self): return len(self.inputs) def __getitem__(self, i): transform = lambda x: (np.sign(x) * np.log1p(np.abs(x)) - self.mean) / self.std return torch.from_numpy(transform(self.inputs[i]).astype(np.float32)), torch.from_numpy(transform(self.targets[i]).astype(np.float32)) def load_model(): spec = importlib.util.spec_from_file_location("surya_model", ROOT / "model/surya.py"); mod = importlib.util.module_from_spec(spec); spec.loader.exec_module(mod); return mod.Surya def args(): p = argparse.ArgumentParser(); p.add_argument("--config", type=Path, default=ROOT / "conf/config.yaml"); p.add_argument("--data", type=Path); p.add_argument("--output", type=Path); p.add_argument("--epochs", type=int); p.add_argument("--batch-size", type=int); p.add_argument("--device", choices=["auto", "cpu", "cuda"]); return p.parse_args() def lr_at(progress, total, warmup, peak, floor): if warmup and progress < warmup: return peak * progress / warmup phase = min(max((progress - warmup) / max(total - warmup, 1), 0), 1) return floor + (peak - floor) * (1 + math.cos(math.pi * phase)) / 2 def main(): a = args(); cfg = yaml.safe_load(a.config.read_text()); tc = cfg["training"] total_epochs = a.epochs or tc["epochs"]; world = int(os.getenv("WORLD_SIZE", "1")); rank = int(os.getenv("RANK", "0")); local = int(os.getenv("LOCAL_RANK", "0")); distributed = world > 1 requested = a.device or cfg["runtime"]["device"]; cuda = torch.cuda.is_available() and requested != "cpu" if requested == "cuda" and not cuda: raise RuntimeError("CUDA requested but unavailable") if distributed: dist.init_process_group("nccl" if cuda else "gloo") device = torch.device(f"cuda:{local}" if cuda else "cpu"); torch.cuda.set_device(local) if cuda else None seed = cfg["seed"] + rank; random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) data_path = a.data or ROOT / cfg["data"]["root"] / "train.npz" dataset = SolarDataset(data_path, cfg); sampler = DistributedSampler(dataset) if distributed else None loader = DataLoader(dataset, batch_size=a.batch_size or tc["batch_size"], shuffle=sampler is None, sampler=sampler, num_workers=tc["num_workers"], pin_memory=cuda) model = load_model()(**cfg["model"]).to(device); bare = model if distributed: model = DistributedDataParallel(model, device_ids=[local] if cuda else None); bare = model.module decay, no_decay = [], [] for n, p in bare.named_parameters(): (no_decay if p.ndim == 1 or n.endswith("bias") else decay).append(p) opt = torch.optim.AdamW([{"params": decay, "weight_decay": tc["weight_decay"]}, {"params": no_decay, "weight_decay": 0}], lr=tc["learning_rate"]) amp = bool(tc.get("amp", True) and cuda); scaler = torch.amp.GradScaler("cuda", enabled=amp); history=[]; opt.zero_grad(set_to_none=True) one_step = tc.get("one_step_epochs", max(1, total_epochs // 2)) for epoch in range(total_epochs): if sampler: sampler.set_epoch(epoch) model.train(); total=0.0; phase = "one_step" if epoch < one_step else "rollout" for step, (x, y) in enumerate(loader): x, y = x.to(device, non_blocking=cuda), y.to(device, non_blocking=cuda); pred_steps = 1 if phase == "one_step" else y.shape[1] for group in opt.param_groups: group["lr"] = lr_at(epoch + step/max(len(loader),1), total_epochs, tc["warmup_epochs"], tc["learning_rate"], tc["min_learning_rate"]) context = torch.amp.autocast("cuda") if amp else nullcontext() with context: pred = model(x, steps=pred_steps); loss = (pred - y[:, :pred_steps]).square().mean() / tc["accum_iter"] if not torch.isfinite(loss): raise FloatingPointError("non-finite training loss") scaler.scale(loss).backward() if (step + 1) % tc["accum_iter"] == 0 or step + 1 == len(loader): scaler.unscale_(opt); torch.nn.utils.clip_grad_norm_(bare.parameters(), tc["grad_clip"]); scaler.step(opt); scaler.update(); opt.zero_grad(set_to_none=True) total += loss.item() * tc["accum_iter"] values = torch.tensor([total / max(len(loader),1)], device=device); dist.all_reduce(values) if distributed else None record={"epoch":epoch+1,"phase":phase,"loss":float(values.item()/world),"learning_rate":opt.param_groups[0]["lr"]}; history.append(record) if rank == 0: print(json.dumps(record)) if rank == 0: path=a.output or ROOT / cfg["paths"]["checkpoint"]; path.parent.mkdir(parents=True,exist_ok=True); torch.save({"model":bare.state_dict(),"optimizer":opt.state_dict(),"scaler":scaler.state_dict(),"epoch":total_epochs-1,"history":history,"config":cfg},path) out=ROOT / cfg["paths"]["training_metrics"]; out.parent.mkdir(parents=True,exist_ok=True); out.write_text(json.dumps({"history":history,"protocol":cfg["data"]["protocol"]},indent=2)+"\n"); print("checkpoint=",path) if distributed: dist.destroy_process_group() if __name__ == "__main__": main()