c93af97059
- 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.
217 lines
6.9 KiB
Python
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),
|
|
}
|