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