Surya / scripts /train.py
zhangrenchao's picture
Upload Surya model package
a13f4b9 verified
Raw
History Blame Contribute Delete
6 kB
"""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()