""" Run NVIDIA StormScope (regional nowcasting) on RTX 5090. StormScope predicts CONUS satellite (GOES) and radar (MRMS) imagery at 6km or 3km resolution with 10-min or 60-min timesteps. Operational modes: nowcast - 10-min steps, pure-obs (no GFS), uses latest GOES/MRMS forecast - 60-min steps, GFS-conditioned, longer range Usage: # Operational: use latest observations python run_stormscope.py --mode nowcast --steps 6 python run_stormscope.py --mode forecast --steps 6 # Historical: specify init time python run_stormscope.py --mode nowcast --steps 6 --date 2023-12-05T12:00:00 python run_stormscope.py --mode nowcast --res 3km --steps 3 # 3km satellite only Output: Each step -> stormscope_{YYYYMMDD}_{HH}z_f{NN}.npz Metadata JSON -> stormscope_{YYYYMMDD}_{HH}z_meta.json """ import argparse import json import os import sys import time from datetime import datetime, timedelta, timezone import numpy as np import torch # When True, GOES/MRMS inputs come from the rust fast path (rustwx_stormscope.rust_datasources) instead of # earth2studio's fetch_data. Set by rustwx_stormscope.datasources.install_backend. RUST_BACKEND = False def _env(name: str, legacy_name: str | None = None, default: str = "") -> str: if name in os.environ: return os.environ[name] if legacy_name and legacy_name in os.environ: return os.environ[legacy_name] return default # Map earth2studio variable names -> wxsection output field names GOES_CHANNEL_MAP = { "abi01c": "ch_blue_047", # GOES Ch1 blue 0.47um "abi02c": "ch_red_064", # GOES Ch2 red 0.64um "abi03c": "ch_veggie_086", # GOES Ch3 veggie 0.86um "abi07c": "ch_swir_390", # GOES Ch7 shortwave IR 3.90um "abi08c": "ch_wv_upper_619", # GOES Ch8 upper WV 6.19um "abi09c": "ch_wv_mid_695", # GOES Ch9 mid WV 6.95um "abi10c": "ch_wv_lower_734", # GOES Ch10 lower WV 7.34um "abi13c": "ch_ir_1035", # GOES Ch13 clean IR 10.35um } # Channel metadata for wxsection colormaps CHANNEL_META = { "ch_blue_047": {"units": "reflectance", "range": [0.0, 1.0], "type": "visible"}, "ch_red_064": {"units": "reflectance", "range": [0.0, 1.0], "type": "visible"}, "ch_veggie_086": {"units": "reflectance", "range": [0.0, 1.0], "type": "visible"}, "ch_swir_390": {"units": "K", "range": [180.0, 310.0], "type": "ir"}, "ch_wv_upper_619": {"units": "K", "range": [190.0, 260.0], "type": "wv"}, "ch_wv_mid_695": {"units": "K", "range": [190.0, 270.0], "type": "wv"}, "ch_wv_lower_734": {"units": "K", "range": [190.0, 275.0], "type": "wv"}, "ch_ir_1035": {"units": "K", "range": [180.0, 310.0], "type": "ir"}, "refc": {"units": "dBZ", "range": [-20.0, 75.0], "type": "radar"}, } # Model registry # Nowcast 6km uses "6steps" sliding-window models (both GOES and MRMS need # 6 past observations at 10-min intervals = 50 min of history). # Nowcast 3km uses "eoe" single-step model (GOES only, no MRMS available). # Forecast uses GFS-conditioned "eoe" models (single-step with GFS context). MODEL_CONFIGS = { ("nowcast", "6km"): { "goes": "6km_10min_natten_pure_obs_zenith_6steps", "mrms": "6km_10min_natten_pure_obs_mrms_obs_6steps", "step_minutes": 10, "needs_gfs": False, }, ("nowcast", "3km"): { "goes": "3km_10min_natten_pure_obs_cos_zenith_input_eoe", "mrms": None, # No 3km MRMS model available "step_minutes": 10, "needs_gfs": False, }, ("forecast", "6km"): { "goes": "6km_60min_natten_cos_zenith_input_eoe_v2", "mrms": "6km_60min_natten_cos_zenith_input_mrms_eoe", "step_minutes": 60, "needs_gfs": True, "conditioning": "gfs", }, } STORMSCOPE_VERSION = "0.12.1" # earth2studio version def forecast_conditioning_kind(cfg: dict | None = None) -> str: """Return the forecast conditioner kind. Default is GFS, matching Earth2Studio's StormScopeGOES.load_model default. RUSTWX_STORMSCOPE_FORECAST_CONDITIONING=hrrr is kept as an explicit experiment switch. """ default = (cfg or {}).get("conditioning", "gfs") kind = _env("RUSTWX_STORMSCOPE_FORECAST_CONDITIONING", "SSFAST_FORECAST_CONDITIONING", default).strip().lower() if kind not in {"gfs", "hrrr"}: raise ValueError("RUSTWX_STORMSCOPE_FORECAST_CONDITIONING must be 'gfs' or 'hrrr'") return kind def get_latest_init_time(step_minutes: int, needs_history: bool, gfs_aligned: bool = False) -> str: """Get the best init time for operational runs using latest observations. Rounds down to the nearest step_minutes boundary and subtracts a small buffer to ensure GOES/MRMS data is available on NOAA S3. If gfs_aligned=True, rounds to the nearest past 6-hour GFS cycle that has had time to post on AWS (~4 hours after init). """ now = datetime.now(timezone.utc) if gfs_aligned: # GFS_FX runs at 00/06/12/18Z. The hourly leads are available from # those cycle roots, but the cycle itself needs time to publish. gfs_buffer_hours = int(_env("RUSTWX_STORMSCOPE_GFS_BUFFER_HOURS", "SSFAST_GFS_BUFFER_HOURS", "4")) cutoff = now - timedelta(hours=gfs_buffer_hours) cutoff = cutoff.replace(minute=0, second=0, microsecond=0) cutoff = cutoff - timedelta(hours=cutoff.hour % 6) return cutoff.strftime("%Y-%m-%dT%H:%M:%S") # GOES CONUS scans every 5 min, MRMS every 2 min. Default to a modest # 10-minute operational buffer; override with RUSTWX_STORMSCOPE_OBS_BUFFER_MIN for # very fresh or very conservative cycling. buffer_min = int(_env("RUSTWX_STORMSCOPE_OBS_BUFFER_MIN", "SSFAST_OBS_BUFFER_MIN", "10")) init = now - timedelta(minutes=buffer_min) # Round down to nearest step_minutes init = init.replace(second=0, microsecond=0) init = init - timedelta(minutes=init.minute % step_minutes) return init.strftime("%Y-%m-%dT%H:%M:%S") def save_step(outdir, init_time_str, step_idx, step_minutes, mode, res, goes_pred, goes_vars, mrms_pred, lats, lons, valid_mask, mrms_valid_mask, dtype_name, goes_model_name, mrms_model_name): """Save one forecast step as a self-describing .npz file.""" init_dt = datetime.fromisoformat(init_time_str) valid_dt = init_dt + timedelta(minutes=step_minutes * (step_idx + 1)) fhr = step_minutes * (step_idx + 1) / 60.0 cycle_str = init_dt.strftime("%Y%m%d_%H%Mz") fname = f"stormscope_{mode}_{res}_{cycle_str}_f{step_idx + 1:02d}.npz" fields = {} # Extract individual GOES channels for i, var_name in enumerate(goes_vars): field = goes_pred[0, 0, 0, i].detach().cpu().to(torch.float32).numpy() field = np.where(valid_mask[0, 0, 0, i].cpu().numpy(), field, np.nan) out_name = GOES_CHANNEL_MAP.get(var_name, var_name) fields[out_name] = field # Extract MRMS reflectivity (if available) if mrms_pred is not None: refc = mrms_pred[0, 0, 0, 0].detach().cpu().to(torch.float32).numpy() refc = np.where(mrms_valid_mask[0, 0, 0, 0].cpu().numpy(), refc, np.nan) fields["refc"] = refc # Grid metadata - convert lons from 0-360 to -180/+180 for wxsection lons_fixed = lons.copy().astype(np.float32) lons_fixed[lons_fixed > 180] -= 360 fields["lats"] = lats.astype(np.float32) fields["lons"] = lons_fixed # Timing metadata fields["init_time"] = init_dt.isoformat() fields["valid_time"] = valid_dt.isoformat() fields["step_minutes"] = step_minutes fields["forecast_hour"] = fhr fields["mode"] = mode fields["resolution"] = res fields["dtype"] = dtype_name fields["goes_model"] = goes_model_name fields["mrms_model"] = mrms_model_name or "none" fields["stormscope_version"] = STORMSCOPE_VERSION outpath = os.path.join(outdir, fname) np.savez_compressed(outpath, **fields) # Print stats label = f"T+{int(fhr)}h" if fhr == int(fhr) else f"T+{fhr:.1f}h" print(f" Saved {fname} ({label}, valid {valid_dt.strftime('%Y-%m-%d %H:%M')}Z)") for name, arr in fields.items(): if isinstance(arr, np.ndarray) and arr.ndim == 2: valid = arr[~np.isnan(arr)] if len(valid) > 0: print(f" {name}: [{valid.min():.1f}, {valid.max():.1f}] mean={valid.mean():.1f}") return outpath def save_cycle_metadata(outdir, init_time_str, mode, res, cfg, saved_files, goes_model_name, mrms_model_name, grid_shape, dtype_name, total_elapsed, complete=True): """Write a JSON manifest for the entire cycle (for unified_dashboard ingestion).""" init_dt = datetime.fromisoformat(init_time_str) cycle_str = init_dt.strftime("%Y%m%d_%H%Mz") step_minutes = cfg["step_minutes"] fields_available = list(GOES_CHANNEL_MAP.values()) if cfg.get("mrms"): fields_available.append("refc") meta = { "model": "stormscope", "cycle": cycle_str, "init_time": init_dt.isoformat() + "Z", "mode": mode, "complete": complete, "resolution": res, "step_minutes": step_minutes, "n_steps": len(saved_files), "valid_times": [ (init_dt + timedelta(minutes=step_minutes * (i + 1))).isoformat() + "Z" for i in range(len(saved_files)) ], "fields": fields_available, "field_metadata": {k: v for k, v in CHANNEL_META.items() if k in fields_available}, "grid": { "shape": list(grid_shape), "projection": "hrrr_lcc", }, "models": { "goes": goes_model_name, "mrms": mrms_model_name or None, }, "inference": { "dtype": dtype_name, "device": torch.cuda.get_device_name() if torch.cuda.is_available() else "cpu", "stormscope_version": STORMSCOPE_VERSION, "conditioning_source": forecast_conditioning_kind(cfg) if cfg.get("needs_gfs") else None, "total_seconds": round(total_elapsed, 1), }, "files": [os.path.basename(f) for f in saved_files], } meta_path = os.path.join(outdir, f"stormscope_{mode}_{res}_{cycle_str}_meta.json") with open(meta_path, "w") as f: json.dump(meta, f, indent=2) if complete: print(f"\nMetadata: {os.path.basename(meta_path)}") return meta_path def run_stormscope(date: str, nsteps: int, outdir: str, mode: str, res: str, use_bf16: bool = True): """Run StormScope in the specified mode.""" from earth2studio.data import GFS_FX, HRRR_FX, GOES, MRMS, fetch_data from earth2studio.models.px.stormscope import ( StormScopeBase, StormScopeGOES, StormScopeMRMS, ) config_key = (mode, res) if config_key not in MODEL_CONFIGS: raise ValueError(f"No model config for mode={mode}, res={res}. " f"Available: {list(MODEL_CONFIGS.keys())}") cfg = MODEL_CONFIGS[config_key] step_minutes = cfg["step_minutes"] conditioning_kind = forecast_conditioning_kind(cfg) if cfg["needs_gfs"] else None # Determine dtype if use_bf16 and torch.cuda.is_bf16_supported(): compute_dtype = torch.bfloat16 dtype_name = "bfloat16" else: compute_dtype = torch.float32 dtype_name = "float32" if use_bf16: print("WARNING: bf16 requested but not supported, falling back to fp32") device = torch.device("cuda" if torch.cuda.is_available() else "cpu") os.makedirs(outdir, exist_ok=True) print(f"Mode: {mode} ({res}, {step_minutes}-min steps)") print(f"Precision: {dtype_name}") if conditioning_kind: print(f"Forecast conditioning: {conditioning_kind.upper()}_FX") else: print("Forecast conditioning: no (pure obs)") print(f"MRMS model: {'yes' if cfg.get('mrms') else 'no'}") run_start = time.time() # --- Load models --- print("\nLoading StormScope models...") t0 = time.time() package = StormScopeBase.load_default_package() # GOES model - load to device first, cast to bf16 AFTER building interpolators goes_model_name = cfg["goes"] if cfg["needs_gfs"]: if conditioning_kind == "gfs": conditioning_source = GFS_FX() elif conditioning_kind == "hrrr": conditioning_source = HRRR_FX() else: raise ValueError(f"Unsupported forecast conditioning source: {conditioning_kind}") model = StormScopeGOES.load_model( package=package, conditioning_data_source=conditioning_source, model_name=goes_model_name, ) else: model = StormScopeGOES.load_model( package=package, model_name=goes_model_name, ) model = model.to(device=device) model.eval() print(f" GOES model ({goes_model_name}) loaded in {time.time() - t0:.1f}s") # Determine GOES satellite: GOES-16 retired Apr 7 2025, GOES-19 took over init_dt = datetime.fromisoformat(date) goes_cutover = datetime(2025, 4, 7) goes_satellite = "goes19" if init_dt >= goes_cutover else "goes16" scan_mode = "C" print(f" Satellite: {goes_satellite.upper()} ({scan_mode})") # MRMS model (optional) model_mrms = None mrms_model_name = cfg.get("mrms") if mrms_model_name: t0 = time.time() model_mrms = StormScopeMRMS.load_model( package=package, conditioning_data_source=GOES(satellite=goes_satellite, scan_mode=scan_mode), model_name=mrms_model_name, ) model_mrms = model_mrms.to(device=device) model_mrms.eval() print(f" MRMS model ({mrms_model_name}) loaded in {time.time() - t0:.1f}s") start_date = [np.datetime64(datetime.fromisoformat(date))] print(f"\nInit time: {start_date[0]}") goes_vars = list(model.input_coords()["variable"]) lat_out = model.latitudes.detach().cpu().to(torch.float32).numpy() lon_out = model.longitudes.detach().cpu().to(torch.float32).numpy() grid_shape = lat_out.shape print(f"Output grid: {grid_shape}, lat [{lat_out.min():.1f}, {lat_out.max():.1f}], " f"lon [{lon_out.min():.1f}, {lon_out.max():.1f}]") goes = GOES(satellite=goes_satellite, scan_mode=scan_mode) goes_lat, goes_lon = GOES.grid(satellite=goes_satellite, scan_mode=scan_mode) # Build interpolators print("Building interpolators...") model.build_input_interpolator(goes_lat, goes_lon) if cfg["needs_gfs"]: if conditioning_kind == "gfs": gfs_lon, gfs_lat = np.meshgrid(GFS_FX.GFS_LON, GFS_FX.GFS_LAT) model.build_conditioning_interpolator( gfs_lat.astype(np.float32), gfs_lon.astype(np.float32), ) elif conditioning_kind == "hrrr": hrrr_lat, hrrr_lon = HRRR_FX.grid() model.build_conditioning_interpolator(hrrr_lat, hrrr_lon) in_coords = model.input_coords() # --- Fetch initial data --- print("Fetching GOES data...") t0 = time.time() if RUST_BACKEND: from rustwx_stormscope import rust_datasources as _rds x, x_coords = _rds.goes_input(goes_satellite, scan_mode, start_date, goes_vars, in_coords["lead_time"], device, goes_lat.shape) else: x, x_coords = fetch_data(goes, time=start_date, variable=np.array(goes_vars), lead_time=in_coords["lead_time"], device=device) print(f" GOES data fetched in {time.time() - t0:.1f}s, shape={x.shape}") # MRMS data (if model available) x_mrms = None x_coords_mrms = None if model_mrms: print("Fetching MRMS data...") t0 = time.time() mrms_in_coords = model_mrms.input_coords() if RUST_BACKEND: from rustwx_stormscope import rust_datasources as _rds x_mrms, x_coords_mrms, _mlat, _mlon = _rds.mrms_input( start_date, mrms_in_coords["lead_time"], device) else: mrms = MRMS() x_mrms, x_coords_mrms = fetch_data(mrms, time=start_date, variable=np.array(["refc"]), lead_time=mrms_in_coords["lead_time"], device=device) print(f" MRMS data fetched in {time.time() - t0:.1f}s, shape={x_mrms.shape}") model_mrms.build_input_interpolator(x_coords_mrms["lat"], x_coords_mrms["lon"]) model_mrms.build_conditioning_interpolator(goes_lat, goes_lon) # --- Add batch dimension --- batch_size = 1 if x.dim() == 5: x = x.unsqueeze(0).repeat(batch_size, 1, 1, 1, 1, 1) x_coords["batch"] = np.arange(batch_size) x_coords.move_to_end("batch", last=False) if x_mrms is not None and x_mrms.dim() == 5: x_mrms = x_mrms.unsqueeze(0).repeat(batch_size, 1, 1, 1, 1, 1) x_coords_mrms["batch"] = np.arange(batch_size) x_coords_mrms.move_to_end("batch", last=False) x = x.to(dtype=torch.float32) if x_mrms is not None: x_mrms = x_mrms.to(dtype=torch.float32) # Use autocast for bf16 - PyTorch handles mixed precision automatically use_autocast = (compute_dtype == torch.bfloat16) if use_autocast: print(f"Using torch.autocast(bf16) for inference") # --- Run coupled forecast, saving each step --- y, y_coords = x, x_coords y_mrms, y_coords_mrms = x_mrms, x_coords_mrms saved_files = [] total_minutes = nsteps * step_minutes print(f"\nRunning {nsteps}-step {mode} ({total_minutes} min = {total_minutes/60:.1f}h)...") for step_idx in range(nsteps): t0 = time.time() # GOES prediction with torch.autocast("cuda", dtype=compute_dtype, enabled=use_autocast): y_pred, y_pred_coords = model(y, y_coords) # MRMS prediction conditioned on GOES (if available) y_mrms_pred = None y_coords_mrms_pred = None if model_mrms and y_mrms is not None: with torch.autocast("cuda", dtype=compute_dtype, enabled=use_autocast): y_mrms_pred, y_coords_mrms_pred = model_mrms.call_with_conditioning( y_mrms, y_coords_mrms, conditioning=y, conditioning_coords=y_coords ) elapsed = time.time() - t0 print(f" Step {step_idx + 1}/{nsteps}: inference {elapsed:.1f}s") # Save this step mrms_valid_mask = None if model_mrms and y_mrms_pred is not None: mrms_valid_mask = model_mrms.valid_mask.expand_as(y_mrms_pred) outpath = save_step( outdir, date, step_idx, step_minutes, mode, res, y_pred, goes_vars, y_mrms_pred, lat_out, lon_out, model.valid_mask.expand_as(y_pred), mrms_valid_mask, dtype_name, goes_model_name, mrms_model_name, ) saved_files.append(outpath) # Keep the operational renderer fed as soon as f01 exists. The service # only treats complete=true manifests as completed cycles. save_cycle_metadata( outdir, date, mode, res, cfg, saved_files, goes_model_name, mrms_model_name, grid_shape, dtype_name, time.time() - run_start, complete=False, ) # Update sliding window for next step y_pred, y_pred_coords = model.next_input(y_pred, y_pred_coords, y, y_coords) if model_mrms and y_mrms_pred is not None: y_mrms_pred, y_coords_mrms_pred = model_mrms.next_input( y_mrms_pred, y_coords_mrms_pred, y_mrms, y_coords_mrms ) y = y_pred y_coords = y_pred_coords if y_mrms_pred is not None: y_mrms = y_mrms_pred y_coords_mrms = y_coords_mrms_pred total_elapsed = time.time() - run_start # Write cycle metadata JSON for unified_dashboard meta_path = save_cycle_metadata( outdir, date, mode, res, cfg, saved_files, goes_model_name, mrms_model_name, grid_shape, dtype_name, total_elapsed, complete=True, ) print(f"\n=== StormScope {mode} complete ===") print(f" {len(saved_files)} steps, {total_elapsed:.0f}s total") print(f" Output: {outdir}/") for f in saved_files: print(f" {os.path.basename(f)}") return saved_files if __name__ == "__main__": parser = argparse.ArgumentParser(description="Run NVIDIA StormScope") parser.add_argument("--mode", type=str, default="nowcast", choices=["nowcast", "forecast"], help="nowcast=10min pure-obs, forecast=60min GFS-conditioned") parser.add_argument("--res", type=str, default="6km", choices=["6km", "3km"], help="Grid resolution (3km only for nowcast mode)") parser.add_argument("--steps", type=int, default=6, help="Forecast steps (default: 6)") parser.add_argument("--date", type=str, default=None, help="Init datetime ISO format (default: latest available)") parser.add_argument("--outdir", type=str, default="outputs/stormscope", help="Output directory") parser.add_argument("--fp32", action="store_true", help="Force fp32 instead of bf16") args = parser.parse_args() print(f"StormScope on {torch.cuda.get_device_name()}") print(f"VRAM: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB") print() # Resolve init time cfg = MODEL_CONFIGS[(args.mode, args.res)] if args.date is None: needs_history = "6steps" in cfg["goes"] date = get_latest_init_time( cfg["step_minutes"], needs_history, gfs_aligned=cfg.get("needs_gfs", False) and forecast_conditioning_kind(cfg) == "gfs", ) print(f"Using latest init time: {date}Z") else: date = args.date run_stormscope(date, args.steps, args.outdir, args.mode, args.res, use_bf16=not args.fp32)