gausubash's picture
Enable ZeroGPU for Chronos forecasts
8396be4 verified
Raw
History Blame Contribute Delete
8.04 kB
from __future__ import annotations
from functools import lru_cache
from pathlib import Path
import joblib
import lightgbm as lgb
import pandas as pd
import plotly.express as px
import plotly.graph_objects as go
from dashboard.runtime import configure_runtime
configure_runtime()
from huggingface_hub import hf_hub_download
REGIONS = ["NSW1", "VIC1", "QLD1", "SA1", "TAS1"]
HORIZONS = {
"5 minutes (1 step)": 1,
"30 minutes (6 steps)": 6,
"1 hour (12 steps)": 12,
}
TIMESTAMP_COLUMN = "SETTLEMENTDATE"
REGION_COLUMN = "REGIONID"
DEMAND_COLUMN = "TOTALDEMAND"
DATASET_REPO = "gausubash/nem-demand-5min"
FALLBACK_DATASET_REPO = "PowerZooJax/PowerZooDataset"
FALLBACK_DATASET_FILE = "parquet/AEMO_5min_Demand_2025_2026.parquet"
ASSETS_DIR = Path(__file__).resolve().parent.parent / "assets"
LOCAL_DEMAND_PATH = ASSETS_DIR / "demand_all_regions.parquet"
@lru_cache(maxsize=1)
def load_demand() -> pd.DataFrame:
if LOCAL_DEMAND_PATH.exists():
frame = pd.read_parquet(LOCAL_DEMAND_PATH)
else:
download_dir = ASSETS_DIR / "downloads"
download_dir.mkdir(parents=True, exist_ok=True)
try:
path = hf_hub_download(
repo_id=DATASET_REPO,
repo_type="dataset",
filename="demand_all_regions.parquet",
local_dir=str(download_dir),
local_dir_use_symlinks=False,
)
frame = pd.read_parquet(path)
except Exception:
path = hf_hub_download(
repo_id=FALLBACK_DATASET_REPO,
repo_type="dataset",
filename=FALLBACK_DATASET_FILE,
local_dir=str(download_dir),
local_dir_use_symlinks=False,
)
raw = pd.read_parquet(path)
frame = raw.rename(
columns={
"INTERVAL_DATETIME": TIMESTAMP_COLUMN,
"OPERATIONAL_DEMAND": DEMAND_COLUMN,
}
)
frame[TIMESTAMP_COLUMN] = pd.to_datetime(frame[TIMESTAMP_COLUMN], utc=True).dt.tz_localize(None)
frame[TIMESTAMP_COLUMN] = pd.to_datetime(frame[TIMESTAMP_COLUMN])
frame[DEMAND_COLUMN] = frame[DEMAND_COLUMN].astype(float)
return frame.sort_values([REGION_COLUMN, TIMESTAMP_COLUMN]).reset_index(drop=True)
@lru_cache(maxsize=1)
def load_benchmark_metrics() -> pd.DataFrame:
path = ASSETS_DIR / "benchmark_metrics.csv"
if path.exists():
return pd.read_csv(path)
return pd.DataFrame()
def load_lightgbm_model(region: str, horizon: int) -> lgb.LGBMRegressor | None:
path = ASSETS_DIR / "models" / f"{region}_h{horizon}.pkl"
if not path.exists():
return None
return joblib.load(path)
def horizon_label(steps: int) -> str:
for label, value in HORIZONS.items():
if value == steps:
return label
return f"{steps} steps"
def demand_summary_table(demand: pd.DataFrame) -> pd.DataFrame:
rows = []
for region in REGIONS:
region_frame = demand.loc[demand[REGION_COLUMN] == region]
if region_frame.empty:
continue
rows.append(
{
"Region": region,
"Rows": len(region_frame),
"Start": region_frame[TIMESTAMP_COLUMN].min(),
"End": region_frame[TIMESTAMP_COLUMN].max(),
"Min MW": round(region_frame[DEMAND_COLUMN].min(), 1),
"Max MW": round(region_frame[DEMAND_COLUMN].max(), 1),
"Mean MW": round(region_frame[DEMAND_COLUMN].mean(), 1),
}
)
return pd.DataFrame(rows)
def region_timeseries_figure(demand: pd.DataFrame, region: str, days: int = 7) -> go.Figure:
region_frame = demand.loc[demand[REGION_COLUMN] == region].copy()
cutoff = region_frame[TIMESTAMP_COLUMN].max() - pd.Timedelta(days=days)
region_frame = region_frame.loc[region_frame[TIMESTAMP_COLUMN] >= cutoff]
fig = px.line(
region_frame,
x=TIMESTAMP_COLUMN,
y=DEMAND_COLUMN,
title=f"{region} demand — last {days} days",
labels={DEMAND_COLUMN: "Demand (MW)", TIMESTAMP_COLUMN: "Time"},
)
fig.update_layout(template="plotly_dark", height=420, margin=dict(l=40, r=20, t=60, b=40))
return fig
def benchmark_bar_figure(metrics: pd.DataFrame, region: str, horizon_steps: int) -> go.Figure:
subset = metrics.loc[
(metrics["region"] == region) & (metrics["horizon_steps"] == horizon_steps)
].copy()
if subset.empty:
fig = go.Figure()
fig.update_layout(title="Benchmark metrics not available yet", template="plotly_dark")
return fig
subset = subset.sort_values("mae")
fig = px.bar(
subset,
x="model",
y="mae",
title=f"Model comparison — {region}, {horizon_label(horizon_steps)}",
labels={"mae": "MAE (MW)", "model": "Model"},
color="model",
)
fig.update_layout(template="plotly_dark", height=420, showlegend=False)
return fig
def benchmark_heatmap_figure(metrics: pd.DataFrame, model: str) -> go.Figure:
subset = metrics.loc[metrics["model"] == model].copy()
if subset.empty:
fig = go.Figure()
fig.update_layout(title=f"No metrics for {model}", template="plotly_dark")
return fig
pivot = subset.pivot(index="region", columns="horizon_steps", values="mae")
fig = px.imshow(
pivot,
text_auto=".1f",
aspect="auto",
title=f"MAE heatmap — {model}",
labels={"color": "MAE (MW)"},
)
fig.update_layout(template="plotly_dark", height=420)
return fig
def forecast_with_chronos(
demand: pd.DataFrame,
region: str,
anchor_time: pd.Timestamp,
horizon: int,
) -> tuple[go.Figure, pd.DataFrame]:
from dashboard.chronos_model import predict_demand
region_frame = demand.loc[demand[REGION_COLUMN] == region].sort_values(TIMESTAMP_COLUMN)
history = region_frame.loc[region_frame[TIMESTAMP_COLUMN] <= anchor_time].tail(512)
if len(history) < 24:
raise ValueError("Not enough history before the selected timestamp.")
context_df = pd.DataFrame(
{
"item_id": region,
"timestamp": history[TIMESTAMP_COLUMN],
"target": history[DEMAND_COLUMN],
}
)
pred_df = predict_demand(context_df, horizon)
history_plot = history.tail(288)
fig = go.Figure()
fig.add_trace(
go.Scatter(
x=history_plot[TIMESTAMP_COLUMN],
y=history_plot[DEMAND_COLUMN],
mode="lines",
name="History",
line=dict(color="#60a5fa"),
)
)
fig.add_trace(
go.Scatter(
x=pred_df["timestamp"],
y=pred_df["0.5"],
mode="lines+markers",
name="Forecast (p50)",
line=dict(color="#34d399"),
)
)
fig.add_trace(
go.Scatter(
x=pred_df["timestamp"],
y=pred_df["0.9"],
mode="lines",
line=dict(width=0),
showlegend=False,
)
)
fig.add_trace(
go.Scatter(
x=pred_df["timestamp"],
y=pred_df["0.1"],
mode="lines",
fill="tonexty",
name="80% interval",
line=dict(width=0),
fillcolor="rgba(52, 211, 153, 0.2)",
)
)
fig.update_layout(
title=f"Chronos GPU forecast — {region}, {horizon_label(horizon)}",
template="plotly_dark",
height=460,
xaxis_title="Time",
yaxis_title="Demand (MW)",
)
display_df = pred_df[["timestamp", "0.1", "0.5", "0.9"]].rename(
columns={"0.1": "p10", "0.5": "p50", "0.9": "p90"}
)
return fig, display_df