Download scripts/train.py from OneScience-Group/Surya: direct link, hf CLI and curl.
- Browser
- Download file 6 kB
-
https://huggingface.co/OneScience-Group/Surya/resolve/main/scripts/train.py
- Command line
-
hf download hf://OneScience-Group/Surya/scripts/train.py
-
curl -L -o train.py https://huggingface.co/OneScience-Group/Surya/resolve/main/scripts/train.py
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() | |