"""Local WeatherNext model runner using JAX/GPU. Runs WeatherNext Cyclones Mini on RTX 4070 SUPER (12GB VRAM). """ import os from datetime import datetime from pathlib import Path from typing import Optional, Dict import numpy as np import pandas as pd from config import WEIGHTS_DIR, HK_BBOX, HK_COORDS class WeatherNextRunner: """Run WeatherNext model locally for Hong Kong region forecasts.""" def __init__(self, model_name: str = "WeatherNextCyclones_Mini_<2024"): self.model_name = model_name self.weights_path = WEIGHTS_DIR / f"{model_name}.npz" self.model = None self.params = None self.state = None self.initialized = False def check_ready(self) -> bool: """Check if model weights exist and JAX is working.""" if not self.weights_path.exists(): print(f"Weights not found: {self.weights_path}") print("Run: python scripts/download_weights.py") return False try: import jax devices = jax.devices() if not devices: print("No JAX devices found") return False print(f"JAX devices: {devices}") return True except ImportError: print("JAX not installed") return False def load_model(self) -> bool: """Load WeatherNext model and weights.""" if not self.check_ready(): return False try: import jax import jax.numpy as jnp from weathernext.models.fgn import FGN, FGNConfig print(f"Loading model weights from {self.weights_path}...") params = dict(np.load(self.weights_path, allow_pickle=True)) config = FGNConfig( resolution=60, # 1° for Mini num_layers=12, model_dim=384, num_heads=6, mlp_ratio=4.0, patch_size=4, max_path_length=240, ) model = FGN(config) self.model = model self.params = params self.initialized = True print("Model loaded successfully") return True except Exception as e: print(f"Failed to load model: {e}") print("Using Open-Meteo WeatherNext API as fallback.") return False def load_initial_state(self, source: str = "era5") -> Optional[Dict]: """Load initial atmospheric state for model input. source: 'era5' (reanalysis), 'hres' (operational), or 'gfs' """ import xarray as xr try: if source == "era5": ds = self._load_era5_latest() elif source == "gfs": ds = self._load_gfs_latest() else: print(f"Unknown source: {source}") return None return self._preprocess_for_model(ds) except Exception as e: print(f"Failed to load initial state: {e}") return None def run_forecast(self, lead_hours: int = 120, steps: int = 20) -> Optional[pd.DataFrame]: """Run an autoregressive forecast for Hong Kong region. lead_hours: Total forecast hours steps: Number of model steps (at 6h per step for Mini) """ if not self.initialized: if not self.load_model(): return None import jax import jax.numpy as jnp initial_state = self.load_initial_state("era5") if initial_state is None: return None try: input_tensor = jnp.array(initial_state["fields"]) @jax.jit def step_fn(params, state): return self.model.apply({"params": params}, state) results = [] current = input_tensor for i in range(steps): current = step_fn(self.params, current) if (i + 1) * 6 <= lead_hours: results.append(self._extract_hk_region(np.array(current), i * 6 + 6)) return self._format_forecast_df(results) except Exception as e: print(f"Model inference failed: {e}") print("Falling back to API-based forecasts.") return None def _extract_hk_region(self, field: np.ndarray, lead_hour: int) -> Dict: """Extract Hong Kong region data from global field. With 1° resolution, HK is roughly a single grid cell. """ lat_idx = slice( int((HK_BBOX["lat_min"] + 90) / 1.0), int((HK_BBOX["lat_max"] + 90) / 1.0) + 1, ) lon_idx = slice( int((HK_BBOX["lon_min"] + 180) / 1.0), int((HK_BBOX["lon_max"] + 180) / 1.0) + 1, ) # Placeholder - actual variable mapping depends on WeatherNext output channels hk_slice = field[..., lat_idx, lon_idx] return { "lead_hour": lead_hour, "temperature_2m_mean": float(np.mean(hk_slice[0]) if hk_slice.size > 0 else np.nan), "precipitation_mean": float(np.mean(hk_slice[-2]) if hk_slice.size > 1 else np.nan), } def _format_forecast_df(self, results: list) -> pd.DataFrame: """Format forecast results as DataFrame.""" if not results: return pd.DataFrame() df = pd.DataFrame(results) now = datetime.now() df["valid_time"] = [now + pd.Timedelta(hours=r["lead_hour"]) for r in results] return df.set_index("valid_time") def _load_era5_latest(self): """Load latest ERA5 data. Requires CDS API setup.""" import xarray as xr from datetime import datetime, timedelta cds_ds = xr.open_dataset( "https://storage.googleapis.com/dm_graphcast/dataset/dataset_test.nc", engine="h5netcdf", ) return cds_ds def _load_gfs_latest(self): """Load latest GFS analysis.""" import xarray as xr from datetime import datetime now = datetime.utcnow() url = f"https://nomads.ncep.noaa.gov/dods/gfs_0p25/gfs{now:%Y%m%d}/gfs_0p25_00z" try: return xr.open_dataset(url) except Exception: return None def _preprocess_for_model(self, ds) -> Dict: """Convert raw dataset to model input format.""" required_vars = [ "2m_temperature", "10m_u_component_of_wind", "10m_v_component_of_wind", "mean_sea_level_pressure", "geopotential", "specific_humidity", "temperature", "u_component_of_wind", "v_component_of_wind", ] fields = [] for var in required_vars: if var in ds: fields.append(np.array(ds[var].isel(time=-1))) else: fields.append(np.zeros((721, 1440))) return { "fields": np.stack(fields), "timestamp": str(ds.time.isel(time=-1).values), }