WeatherPred / weatherpred /modeling.py
yyyc's picture
fist push the app
5194558
Raw
History Blame Contribute Delete
5.01 kB
"""ONNX weather inference and forecast comparison logic."""
from __future__ import annotations
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any
import numpy as np
import pandas as pd
from weatherpred.config import (
API_DISPLAY_COLUMNS,
FORECAST_DAYS,
LOOKBACK_DAYS,
MODEL_PATH,
)
@dataclass
class WeatherBundle:
history: pd.DataFrame
api_forecast: pd.DataFrame
model_forecast: pd.DataFrame
comparison: pd.DataFrame
summary: dict[str, Any]
def run_onnx_temperature_model(history: pd.DataFrame) -> np.ndarray:
recent_highs = (
history.sort_values("date")["high_c"].tail(LOOKBACK_DAYS).to_numpy(dtype=np.float32)
)
if len(recent_highs) < LOOKBACK_DAYS:
raise ValueError(f"Need at least {LOOKBACK_DAYS} historical days for ONNX inference.")
try:
import onnxruntime as ort
except ImportError as exc:
raise RuntimeError(
"onnxruntime is required to run models/mock_model/model.onnx. "
"Install requirements.txt or deploy on HuggingFace Spaces."
) from exc
session = ort.InferenceSession(str(MODEL_PATH), providers=["CPUExecutionProvider"])
input_name = session.get_inputs()[0].name
output_name = session.get_outputs()[0].name
prediction = session.run([output_name], {input_name: recent_highs.reshape(1, -1)})[0]
return np.asarray(prediction, dtype=float).reshape(-1)[:FORECAST_DAYS]
def seasonal_reference(history: pd.DataFrame, forecast_dates: list[pd.Timestamp]) -> pd.DataFrame:
frame = history.copy()
frame["day_of_year"] = frame["date"].dt.dayofyear
rows = []
for forecast_date in forecast_dates:
day_of_year = forecast_date.dayofyear
window = frame[
(frame["day_of_year"] >= day_of_year - 14)
& (frame["day_of_year"] <= day_of_year + 14)
]
if window.empty:
window = frame.tail(365)
rows.append(
{
"date": forecast_date,
"seasonal_high_c": float(window["high_c"].median()),
"seasonal_low_gap": float((window["high_c"] - window["low_c"]).median()),
"rain_mm": float(window["rain_mm"].median()),
"wind_kmh": float(window["wind_kmh"].median()),
"high_std": float(window["high_c"].std(ddof=0) or 3.0),
}
)
return pd.DataFrame(rows)
def build_model_forecast(history: pd.DataFrame, api_forecast: pd.DataFrame) -> pd.DataFrame:
forecast_dates = list(api_forecast["date"])
seasonal = seasonal_reference(history, forecast_dates)
predicted_highs = run_onnx_temperature_model(history)
horizon = len(forecast_dates)
if horizon <= len(predicted_highs):
model_highs = predicted_highs[:horizon]
else:
seasonal_extension = seasonal["seasonal_high_c"].to_numpy(dtype=float)[len(predicted_highs) : horizon]
model_highs = np.concatenate([predicted_highs, seasonal_extension])
result = pd.DataFrame({"date": forecast_dates, "model_high_c": model_highs})
result = result.merge(seasonal, on="date", how="left")
result["model_low_c"] = result["model_high_c"] - result["seasonal_low_gap"]
result["model_rain_mm"] = result["rain_mm"].clip(lower=0)
result["model_wind_kmh"] = result["wind_kmh"].clip(lower=0)
result["model_confidence"] = (1 - (result["high_std"] / 12)).clip(lower=0.45, upper=0.88)
return result[
[
"date",
"model_high_c",
"model_low_c",
"model_rain_mm",
"model_wind_kmh",
"model_confidence",
]
]
def build_weather_bundle(force_refresh_history: bool = False, horizon_days: int = FORECAST_DAYS) -> WeatherBundle:
from weatherpred.data import fetch_api_forecast, load_or_fetch_history
history = load_or_fetch_history(force_refresh=force_refresh_history)
api_forecast = fetch_api_forecast(days=horizon_days)
model_forecast = build_model_forecast(history, api_forecast)
comparison = api_forecast.rename(columns=API_DISPLAY_COLUMNS).merge(
model_forecast, on="date", how="inner"
)
comparison["date"] = comparison["date"].dt.strftime("%Y-%m-%d")
comparison["high_delta_c"] = comparison["model_high_c"] - comparison["api_high_c"]
comparison["rain_delta_mm"] = comparison["model_rain_mm"] - comparison["api_rain_mm"]
summary = {
"history_rows": int(len(history)),
"history_start": history["date"].min().strftime("%Y-%m-%d"),
"history_end": history["date"].max().strftime("%Y-%m-%d"),
"fetched_at": datetime.now(timezone.utc).isoformat(),
"avg_model_confidence": float(model_forecast["model_confidence"].mean()),
"max_api_high_c": float(api_forecast["high_c"].max()),
"max_model_high_c": float(model_forecast["model_high_c"].max()),
}
return WeatherBundle(history, api_forecast, model_forecast, comparison, summary)