Download scripts/train.py from OneScience-Group/Stormer: direct link, hf CLI and curl.
- Browser
- Download file 20.3 kB
-
https://huggingface.co/OneScience-Group/Stormer/resolve/main/scripts/train.py
- Command line
-
hf download hf://OneScience-Group/Stormer/scripts/train.py
-
curl -L -o train.py https://huggingface.co/OneScience-Group/Stormer/resolve/main/scripts/train.py
20.3 kB
| """ | |
| Stormer Training Script — Matching Official Implementation. | |
| Training follows the official Stormer approach exactly: | |
| 1. Input x(t) is normalized with inp_transform (N(0,1) per variable) | |
| 2. Model predicts Δx in DIFF-NORMALIZED space | |
| 3. Ground truth: Δx_gt = diff_transform[Δt](raw_out - raw_in) | |
| 4. Loss: L1 between pred_diff and gt_diff (both in diff-normalized space) | |
| 5. Autoregressive rollout with: norm_diff → raw_diff → original → re-normalize | |
| Supports: | |
| - Single GPU: python scripts/train.py | |
| - Multi-GPU: torchrun --nproc_per_node=N scripts/train.py | |
| - Cluster (slurm): sbatch work_slurm.sh | |
| """ | |
| import torch | |
| import torch.distributed as dist | |
| import torch.nn as nn | |
| import os | |
| import sys | |
| import warnings | |
| from pathlib import Path | |
| # Suppress warnings from external libraries (apex, etc.) | |
| warnings.filterwarnings("ignore", category=UserWarning, module="apex") | |
| warnings.filterwarnings("ignore", message=".*DtypeTensor constructors.*") | |
| root_path = Path(__file__).parent.parent | |
| sys.path.append(str(root_path)) | |
| import numpy as np | |
| import logging | |
| import time | |
| import random | |
| from torch.nn.parallel import DistributedDataParallel | |
| from model.stormer import Stormer | |
| from onescience.datapipes.climate import ERA5Datapipe | |
| from onescience.utils.YParams import YParams | |
| try: | |
| from apex import optimizers | |
| HAS_APEX = True | |
| except ImportError: | |
| HAS_APEX = False | |
| # ============================================================================ | |
| # Normalization utilities | |
| # ============================================================================ | |
| class Normalize: | |
| """Per-variable normalization: y = (x - mean) / std. | |
| Replaces torchvision.transforms.Normalize to avoid dependency. | |
| """ | |
| def __init__(self, mean, std): | |
| # mean, std: (V,) tensors | |
| self.mean = mean.view(1, -1, 1, 1) | |
| self.std = std.view(1, -1, 1, 1) | |
| def __call__(self, x): | |
| # x: (B, V, H, W) or (V, H, W) | |
| if x.dim() == 3: | |
| x = x.unsqueeze(0) | |
| return ((x - self.mean) / self.std).squeeze(0) | |
| return (x - self.mean) / self.std | |
| def get_reverse_transform(transform): | |
| """Return the inverse of a Normalize transform.""" | |
| mean = transform.mean.view(-1) | |
| std = transform.std.view(-1) | |
| std_rev = 1.0 / std | |
| mean_rev = -mean * std_rev | |
| return Normalize(mean_rev, std_rev) | |
| def load_normalization_stats(normalize_dir, variables): | |
| """Load official Stormer normalization constants. | |
| Returns: | |
| inp_transform: Normalize for input fields | |
| reverse_inp_transform: inverse of inp_transform | |
| diff_transform: dict {interval: Normalize} for diff fields | |
| reverse_diff_transform: dict {interval: Normalize} inverse | |
| """ | |
| # Input normalization | |
| mean_dict = dict(np.load(os.path.join(normalize_dir, "normalize_mean.npz"))) | |
| std_dict = dict(np.load(os.path.join(normalize_dir, "normalize_std.npz"))) | |
| inp_mean = np.concatenate([mean_dict[v] for v in variables], axis=0) | |
| inp_std = np.concatenate([std_dict[v] for v in variables], axis=0) | |
| inp_mean_t = torch.from_numpy(inp_mean).float() | |
| inp_std_t = torch.from_numpy(inp_std).float() | |
| inp_transform = Normalize(inp_mean_t, inp_std_t) | |
| reverse_inp_transform = get_reverse_transform(inp_transform) | |
| # Diff normalization for each interval | |
| diff_transform = {} | |
| reverse_diff_transform = {} | |
| for interval in [6, 12, 24]: | |
| dmean_dict = dict(np.load( | |
| os.path.join(normalize_dir, f"normalize_diff_mean_{interval}.npz"))) | |
| dstd_dict = dict(np.load( | |
| os.path.join(normalize_dir, f"normalize_diff_std_{interval}.npz"))) | |
| dmean = np.concatenate([dmean_dict[v] for v in variables], axis=0) | |
| dstd = np.concatenate([dstd_dict[v] for v in variables], axis=0) | |
| dmean_t = torch.from_numpy(dmean).float() | |
| dstd_t = torch.from_numpy(dstd).float() | |
| diff_transform[interval] = Normalize(dmean_t, dstd_t) | |
| reverse_diff_transform[interval] = get_reverse_transform( | |
| diff_transform[interval]) | |
| return (inp_transform, reverse_inp_transform, | |
| diff_transform, reverse_diff_transform) | |
| # ============================================================================ | |
| # Training | |
| # ============================================================================ | |
| def main(): | |
| logging.basicConfig( | |
| level=logging.INFO, | |
| format="%(asctime)s - %(levelname)s - %(message)s" | |
| ) | |
| logger = logging.getLogger() | |
| # ============================================================ | |
| # Config | |
| # ============================================================ | |
| config_file_path = os.path.join(current_path, "conf/config.yaml") | |
| cfg = YParams(config_file_path, "model") | |
| cfg_data = YParams(config_file_path, "datapipe") | |
| # ============================================================ | |
| # Distributed setup | |
| # ============================================================ | |
| cfg.world_size = 1 | |
| if "WORLD_SIZE" in os.environ: | |
| cfg.world_size = int(os.environ["WORLD_SIZE"]) | |
| world_rank = 0 | |
| local_rank = 0 | |
| if cfg.world_size > 1: | |
| dist.init_process_group(backend="nccl", init_method="env://") | |
| local_rank = int(os.environ["LOCAL_RANK"]) | |
| world_rank = dist.get_rank() | |
| # ============================================================ | |
| # Load normalization stats | |
| # ============================================================ | |
| normalize_dir = cfg.normalize_dir | |
| variables = cfg_data.dataset.channels | |
| (inp_transform, reverse_inp_transform, | |
| diff_transform, reverse_diff_transform) = load_normalization_stats( | |
| normalize_dir, variables) | |
| # Move transforms to device | |
| for key in diff_transform: | |
| diff_transform[key].mean = diff_transform[key].mean.to(local_rank) | |
| diff_transform[key].std = diff_transform[key].std.to(local_rank) | |
| reverse_diff_transform[key].mean = reverse_diff_transform[key].mean.to(local_rank) | |
| reverse_diff_transform[key].std = reverse_diff_transform[key].std.to(local_rank) | |
| inp_transform.mean = inp_transform.mean.to(local_rank) | |
| inp_transform.std = inp_transform.std.to(local_rank) | |
| reverse_inp_transform.mean = reverse_inp_transform.mean.to(local_rank) | |
| reverse_inp_transform.std = reverse_inp_transform.std.to(local_rank) | |
| # ============================================================ | |
| # DataLoader — get RAW data (normalize=False since we apply | |
| # official normalization manually) | |
| # ============================================================ | |
| max_output_steps = max(cfg.list_train_intervals) // cfg.data_freq # 24/6 = 4 | |
| datapipe = ERA5Datapipe( | |
| dataset_dir=cfg_data.dataset.data_dir, | |
| used_variables=variables, | |
| used_years=cfg_data.dataset.train_time, | |
| distributed=dist.is_initialized(), | |
| input_steps=1, | |
| output_steps=max_output_steps, | |
| normalize=False, # Raw data — apply official normalization manually | |
| batch_size=cfg_data.dataloader.batch_size, | |
| num_workers=cfg_data.dataloader.num_workers, | |
| ) | |
| train_dataloader, train_sampler = datapipe.get_dataloader("train") | |
| # Validation: use output_steps=1 (6h) to avoid negative samples_per_year | |
| # with T=10 fake data (10-1-12+1 < 0 for output_steps=12) | |
| val_datapipe = ERA5Datapipe( | |
| dataset_dir=cfg_data.dataset.data_dir, | |
| used_variables=variables, | |
| used_years=cfg_data.dataset.val_time, | |
| distributed=dist.is_initialized(), | |
| input_steps=1, | |
| output_steps=1, | |
| normalize=False, | |
| batch_size=cfg_data.dataloader.batch_size, | |
| num_workers=cfg_data.dataloader.num_workers, | |
| ) | |
| val_dataloader, val_sampler = val_datapipe.get_dataloader("valid") | |
| # ============================================================ | |
| # Model | |
| # ============================================================ | |
| model = Stormer( | |
| in_img_size=cfg.in_img_size, | |
| variables=variables, | |
| patch_size=cfg.patch_size, | |
| hidden_size=cfg.hidden_size, | |
| depth=cfg.depth, | |
| num_heads=cfg.num_heads, | |
| mlp_ratio=cfg.mlp_ratio, | |
| ).to(local_rank) | |
| # Mixed precision: FP16 autocast + GradScaler (matches official precision=16) | |
| use_amp = (local_rank >= 0) # enable AMP when GPU/DCU available | |
| scaler = torch.amp.GradScaler('cuda', enabled=use_amp) | |
| # Optimizer — matching official: separate weight decay for embedding params | |
| decay = [] | |
| no_decay = [] | |
| for name, m in model.named_parameters(): | |
| if "channel_embed" in name or "pos_embed" in name: | |
| no_decay.append(m) | |
| else: | |
| decay.append(m) | |
| if HAS_APEX: | |
| optimizer = optimizers.FusedAdam( | |
| [{"params": decay, "lr": cfg.lr, | |
| "betas": (cfg.beta_1, cfg.beta_2), | |
| "weight_decay": cfg.weight_decay}, | |
| {"params": no_decay, "lr": cfg.lr, | |
| "betas": (cfg.beta_1, cfg.beta_2), | |
| "weight_decay": 0}] | |
| ) | |
| else: | |
| optimizer = torch.optim.AdamW( | |
| [{"params": decay, "lr": cfg.lr, | |
| "betas": (cfg.beta_1, cfg.beta_2), | |
| "weight_decay": cfg.weight_decay}, | |
| {"params": no_decay, "lr": cfg.lr, | |
| "betas": (cfg.beta_1, cfg.beta_2), | |
| "weight_decay": 0}] | |
| ) | |
| scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau( | |
| optimizer, factor=0.2, patience=5, mode="min" | |
| ) | |
| # Latitude-weighted L1 loss | |
| lat_path = os.path.join(cfg_data.dataset.data_dir, "static", "lat.npy") | |
| if os.path.exists(lat_path): | |
| lat = np.load(lat_path) | |
| else: | |
| lat = np.linspace(90, -90, cfg.in_img_size[0]) | |
| loss_obj = _LatWeightedL1Loss(lat, local_rank) | |
| # ============================================================ | |
| # Training state | |
| # ============================================================ | |
| os.makedirs(cfg.checkpoint_dir, exist_ok=True) | |
| train_loss_file = f"{cfg.checkpoint_dir}/trloss.npy" | |
| valid_loss_file = f"{cfg.checkpoint_dir}/valoss.npy" | |
| best_valid_loss = 1.0e6 | |
| best_loss_epoch = 0 | |
| train_losses = np.empty((0,), dtype=np.float32) | |
| valid_losses = np.empty((0,), dtype=np.float32) | |
| if world_rank == 0: | |
| total_params = sum(p.numel() for p in model.parameters()) | |
| print("\n" + "-" * 50) | |
| print(f"📂 Stormer params: {total_params:,} " | |
| f"({total_params / 1e6:.2f}M, {total_params / 1e9:.2f}B)") | |
| print(f" Resolution: {cfg.in_img_size}") | |
| print(f" Patch size: {cfg.patch_size}, Hidden: {cfg.hidden_size}") | |
| print(f" Depth: {cfg.depth}, Heads: {cfg.num_heads}") | |
| print(f" Variables: {len(variables)}") | |
| print(f" Intervals: {cfg.list_train_intervals}") | |
| print(f" Norm dir: {normalize_dir}") | |
| print("-" * 50 + "\n") | |
| # Load checkpoint | |
| if os.path.exists(f"{cfg.checkpoint_dir}/model_bak.pth"): | |
| if world_rank == 0: | |
| print(f"✅ Found checkpoint, resuming training...") | |
| ckpt = torch.load( | |
| f"{cfg.checkpoint_dir}/model_bak.pth", | |
| map_location=f'cuda:{local_rank}', weights_only=False, | |
| ) | |
| model.load_state_dict(ckpt["model_state_dict"]) | |
| optimizer.load_state_dict(ckpt["optimizer_state_dict"]) | |
| scheduler.load_state_dict(ckpt["scheduler_state_dict"]) | |
| best_valid_loss = ckpt["best_valid_loss"] | |
| best_loss_epoch = ckpt["best_loss_epoch"] | |
| if os.path.exists(train_loss_file): | |
| train_losses = np.load(train_loss_file) | |
| if os.path.exists(valid_loss_file): | |
| valid_losses = np.load(valid_loss_file) | |
| # DDP | |
| if cfg.world_size > 1: | |
| model = DistributedDataParallel( | |
| model, device_ids=[local_rank], output_device=local_rank, | |
| find_unused_parameters=True, | |
| ) | |
| if world_rank == 0: | |
| logger.info("Starting Stormer training...") | |
| # ============================================================ | |
| # Training loop | |
| # ============================================================ | |
| for epoch in range(cfg.start_epoch, cfg.max_epoch): | |
| if dist.is_initialized(): | |
| train_sampler.set_epoch(epoch) | |
| val_sampler.set_epoch(epoch) | |
| # ---- Train ---- | |
| model.train() | |
| train_loss = 0 | |
| start_time = time.time() | |
| for j, data in enumerate(train_dataloader): | |
| # ERA5Dataset (normalize=False): | |
| # DataLoader adds batch dim → need to squeeze(batch_size=1) | |
| # invar: (B, C, H, W) → squeeze → (C, H, W) raw at time t | |
| # outvar: (B, T_out, C, H, W) → squeeze → (T_out, C, H, W) raw future | |
| invar = data[0].to(local_rank, dtype=torch.float32).squeeze(0) | |
| outvar = data[1].to(local_rank, dtype=torch.float32).squeeze(0) | |
| # Randomly select training interval | |
| chosen_interval = random.choice(cfg.list_train_intervals) | |
| step_idx = chosen_interval // cfg.data_freq - 1 # 6→0, 12→1, 24→3 | |
| # ---- Forward pass following official forward_train logic ---- | |
| # Step 1: Normalize input: (C,H,W) → (V,H,W) → unsqueeze → (1,V,H,W) | |
| x = inp_transform(invar).unsqueeze(0) # (1, V, H, W) | |
| # Step 2: Compute ground truth diff in DIFF-NORMALIZED space | |
| # raw_diff = outvar[step_idx] - invar (both raw, step_idx on time dim) | |
| raw_diff = outvar[step_idx] - invar # (V, H, W) | |
| gt_norm_diff = diff_transform[chosen_interval](raw_diff) # (V, H, W) | |
| # Step 3: Model forward with AMP + gradient checkpointing | |
| interval_tensor = torch.tensor( | |
| [chosen_interval], device=local_rank, dtype=torch.float32 | |
| ) | |
| with torch.amp.autocast('cuda', enabled=use_amp, dtype=torch.float16): | |
| pred_norm_diff = model( | |
| x, variables, interval_tensor, use_checkpoint=True, | |
| ) # (1, V, H, W) | |
| pred_norm_diff = _replace_constant(pred_norm_diff.float(), variables) | |
| # Step 4: Loss in diff-normalized space | |
| gt_norm_diff = gt_norm_diff.unsqueeze(0) # (1, V, H, W) | |
| loss = loss_obj(pred_norm_diff, gt_norm_diff) | |
| optimizer.zero_grad() | |
| scaler.scale(loss).backward() | |
| scaler.step(optimizer) | |
| scaler.update() | |
| train_loss += loss.item() | |
| if world_rank == 0: | |
| elapsed = time.time() - start_time | |
| logger.info( | |
| f'Train: Epoch {epoch}-{j+1}/{len(train_dataloader)} ' | |
| f'[cost {int(elapsed // 60):02}:{int(elapsed % 60):02}] ' | |
| f'[{elapsed/(j+1):.02f}s/batch] ' | |
| f'interval={chosen_interval}h ' | |
| f'loss:{train_loss / (j+1):.04f}' | |
| ) | |
| train_loss /= len(train_dataloader) | |
| # ---- Validation ---- | |
| model.eval() | |
| valid_loss = 0 | |
| val_start = time.time() | |
| with torch.no_grad(): | |
| for j, data in enumerate(val_dataloader): | |
| # Squeeze batch dim (batch_size=1) | |
| invar = data[0].to(local_rank, dtype=torch.float32).squeeze(0) | |
| outvar = data[1].to(local_rank, dtype=torch.float32).squeeze(0) | |
| # Use 6h interval for validation (simple next-step) | |
| val_interval = 6 | |
| target_frame = outvar[0] # t+6h | |
| # Same forward logic as training | |
| x = inp_transform(invar).unsqueeze(0) | |
| raw_diff = target_frame - invar | |
| gt_norm_diff = diff_transform[val_interval](raw_diff) | |
| interval_tensor = torch.tensor( | |
| [val_interval], device=local_rank, dtype=torch.float32 | |
| ) | |
| with torch.amp.autocast('cuda', enabled=use_amp, dtype=torch.float16): | |
| pred_norm_diff = model( | |
| x, variables, interval_tensor, use_checkpoint=False, | |
| ) | |
| pred_norm_diff = _replace_constant(pred_norm_diff.float(), variables) | |
| gt_norm_diff = gt_norm_diff.unsqueeze(0) | |
| loss = loss_obj(pred_norm_diff, gt_norm_diff) | |
| if cfg.world_size > 1: | |
| loss_tensor = loss.detach().to(local_rank) | |
| dist.all_reduce(loss_tensor) | |
| loss = loss_tensor.item() / cfg.world_size | |
| valid_loss += loss | |
| else: | |
| valid_loss += loss.item() | |
| if world_rank == 0: | |
| logger.info( | |
| f'Valid: Epoch {epoch}-{j+1}/{len(val_dataloader)} ' | |
| f'[{(time.time()-val_start)/(j+1):.02f}s/batch] ' | |
| f'loss:{valid_loss / (j+1):.04f}' | |
| ) | |
| valid_loss /= len(val_dataloader) | |
| # ---- Checkpoint & Early stopping ---- | |
| is_save_ckp = False | |
| if valid_loss < best_valid_loss: | |
| best_valid_loss = valid_loss | |
| best_loss_epoch = epoch | |
| if world_rank == 0: | |
| save_checkpoint( | |
| model, optimizer, scheduler, | |
| best_valid_loss, best_loss_epoch, cfg.checkpoint_dir | |
| ) | |
| is_save_ckp = True | |
| scheduler.step(valid_loss) | |
| if world_rank == 0: | |
| logger.info( | |
| f"Epoch [{epoch + 1}/{cfg.max_epoch}], " | |
| f"Train Loss: {train_loss:.4f}, " | |
| f"Valid Loss: {valid_loss:.4f}, " | |
| f"Best loss at Epoch: {best_loss_epoch + 1}" | |
| + (", saving checkpoint" if is_save_ckp else "") | |
| ) | |
| train_losses = np.append(train_losses, train_loss) | |
| valid_losses = np.append(valid_losses, valid_loss) | |
| np.save(train_loss_file, train_losses) | |
| np.save(valid_loss_file, valid_losses) | |
| if epoch - best_loss_epoch > cfg.patience: | |
| if world_rank == 0: | |
| print(f"Loss has not decreased in {cfg.patience} epochs, stopping.") | |
| break | |
| # ============================================================================ | |
| # Helper: zero out constant variable predictions | |
| # ============================================================================ | |
| from model.stormer import CONSTANTS | |
| def _replace_constant(yhat, out_variables): | |
| """Zero out diffs for constant/invariant variables.""" | |
| for i in range(yhat.shape[1]): | |
| if out_variables[i] in CONSTANTS: | |
| yhat[:, i] = 0.0 | |
| return yhat | |
| # ============================================================================ | |
| # Latitude-weighted L1 Loss | |
| # ============================================================================ | |
| class _LatWeightedL1Loss(nn.Module): | |
| """Latitude-weighted L1 loss.""" | |
| def __init__(self, lat, device): | |
| super().__init__() | |
| w_lat = np.cos(np.deg2rad(lat)) | |
| w_lat = w_lat / w_lat.mean() | |
| self.w_lat = torch.from_numpy(w_lat).float().to(device).unsqueeze(0).unsqueeze(-1) | |
| def forward(self, pred, target): | |
| error = torch.abs(pred - target) # (B, V, H, W) | |
| weighted = error * self.w_lat.unsqueeze(1) # (B, V, H, W) | |
| return weighted.mean() | |
| # ============================================================================ | |
| # Checkpoint utilities | |
| # ============================================================================ | |
| def save_checkpoint(model, optimizer, scheduler, best_valid_loss, | |
| best_loss_epoch, model_path): | |
| model_to_save = model.module if hasattr(model, "module") else model | |
| state = { | |
| "model_state_dict": model_to_save.state_dict(), | |
| "optimizer_state_dict": optimizer.state_dict(), | |
| "scheduler_state_dict": scheduler.state_dict(), | |
| "best_valid_loss": best_valid_loss, | |
| "best_loss_epoch": best_loss_epoch, | |
| } | |
| torch.save(state, f"{model_path}/model.pth") | |
| os.system(f"cp {model_path}/model.pth {model_path}/model_bak.pth") | |
| if __name__ == "__main__": | |
| current_path = os.getcwd() | |
| sys.path.append(current_path) | |
| main() | |