Files
hk-weather-mkt/weather/model_runner.py
T
ramseshk c93af97059 HK Weather Prediction Market Pipeline: WeatherNext + HKO + Polymarket
- Open-Meteo WeatherNext API client for HK forecasts
- HKO public data client (current conditions, 9-day forecast, typhoon warnings)
- HK-specific weather extraction and calibration
- Polymarket market scanning, price discovery, and market creation proposals
- Trading strategy engine: edge detection, Kelly criterion sizing, probability calibration
- End-to-end pipeline with dry-run mode and scheduled runner
- Interactive dashboard with live HK weather + forecasts + trading signals

Dependencies: Python 3.10+, openmeteo-requests, pandas
No API keys needed for dry-run mode.
Polymarket trading requires private key in .env.
2026-08-10 12:48:05 +08:00

217 lines
6.9 KiB
Python

"""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),
}