TiRex-2-demo / src /plotting.py
Christian Ganhör
Initial commit
c9ef990
Raw
History Blame Contribute Delete
9.35 kB
"""Plotly figures for the TiRex-2 demo, with a consistent brand palette."""
from __future__ import annotations
import numpy as np
import pandas as pd
import plotly.express as px
import plotly.graph_objects as go
from plotly.subplots import make_subplots
import plotly.colors as pc
from tirex2.plotting import (
COVARIATE_COLORS,
plot_covariate,
plot_forecast,
)
from core import ForecastResult
COLOR_PALETTE = px.colors.qualitative.G10
# Brand-ish, colour-blind-friendly palette (shared with the covariate example).
C_HISTORY = COLOR_PALETTE[0] # navy - observed history
C_TRUTH = COLOR_PALETTE[0] # black - ground truth (dashed)
C_MEDIAN = COLOR_PALETTE[1] # orange - forecast median
C_MULTI = COLOR_PALETTE[2] # green - multivariate forecast (covariate lab)
PLOTLY_TEMPLATE = "plotly_white"
# colour-blind-friendly, skipping first three which are similar to history and forecast
C_COVARIATES = COLOR_PALETTE[3:]
# Shared typography / chrome so charts match the app's Inter-based styling.
FONT = dict(family="Inter, -apple-system, BlinkMacSystemFont, sans-serif",
color="#16202e", size=11)
C_GRID = "#eef2f7"
def increase_color_brightness(hex_color: str, factor: float) -> str:
"""Increase brightness of a hex color by a given factor."""
rgb = pc.hex_to_rgb(hex_color)
brighter_rgb = pc.find_intermediate_color(rgb, pc.hex_to_rgb("#FFFFFF"), factor)
return pc.label_rgb(brighter_rgb)
def hex_to_rgba(hex_color: str, alpha: float) -> str:
r, g, b = pc.hex_to_rgb(hex_color)
return f"rgba({r}, {g}, {b}, {alpha})"
def _polish_axes(fig) -> None:
"""Lighten gridlines and drop the heavy axis chrome for a cleaner look."""
fig.update_xaxes(showgrid=True, gridcolor=C_GRID, zeroline=False,
showline=False, ticks="", title_font=dict(size=12, color="#8a93a3"))
fig.update_yaxes(showgrid=True, gridcolor=C_GRID, zeroline=False,
showline=False, ticks="", title_font=dict(size=12, color="#8a93a3"))
def _iter_covariate_rows(result: ForecastResult, x):
"""Yield ``(label, x, y)`` per covariate, aligned to the context origin.
The full covariate history is returned (never trimmed) so no data is dropped; the
view is narrowed later purely via the shared x-axis range. Past covariates carry
only history, so their x-values stop at the forecast start; future covariates run
through the horizon.
"""
ts = result.timeseries
if result.cov_mode == "past":
arrays = ts.past_covariates if ts is not None else None
elif result.cov_mode == "future":
arrays = ts.future_covariates if ts is not None else None
else:
arrays = None
if arrays is None:
return
labels = result.cov_names or []
x = np.asarray(x)
for i, cov in enumerate(np.asarray(arrays, dtype=np.float32)):
xc = x[:len(cov)]
yc = cov[:len(xc)]
label = labels[i] if i < len(labels) else f"Covariate {i + 1}"
yield label, xc, yc
def _as_positions(x):
"""Return integer plot positions plus optional date labels.
The tirex2 plotting primitives crash on a datetime x-axis when a series is absent
(they compare a ``Timestamp`` against ``np.inf``). Feeding them plain positions and
relabelling the axis with dates sidesteps that while keeping a readable time axis.
"""
x = np.asarray(x)
is_datetime = x.dtype.kind == "M" or (
x.dtype == object and len(x) and isinstance(x[0], pd.Timestamp)
)
if is_datetime:
return np.arange(len(x)), pd.to_datetime(x)
return x, None
def _apply_date_ticks(fig, positions, labels) -> None:
"""Relabel the (numeric) x-axis with ~8 formatted date ticks."""
n = min(8, len(labels))
if n < 2:
return
idx = np.unique(np.linspace(0, len(labels) - 1, n).astype(int))
span = labels[-1] - labels[0]
if span <= pd.Timedelta(days=3):
fmt = "%Y-%m-%d %H:%M"
elif span <= pd.Timedelta(days=1200):
fmt = "%Y-%m-%d"
else:
fmt = "%Y-%m"
fig.update_xaxes(
tickmode="array",
tickvals=[positions[i] for i in idx],
ticktext=[labels[i].strftime(fmt) for i in idx],
)
def build_forecast_figure(
result: ForecastResult,
baseline_result: ForecastResult | None,
x,
*,
max_context_to_show: int,
ground_truth=None,
):
"""Assemble the forecast figure and return ``(fig, n_rows)``.
With covariates, the target is drawn as two stacked, directly comparable panels -
a univariate TiRex baseline and the covariate-informed forecast - followed by one
panel per covariate (mirrors ``tirex2.demo.plot_demo_forecast``). Without covariates
it draws a single target panel.
In every case the *full* context and covariate history is plotted; ``max_context_to_show``
only narrows the initial visible window by setting a shared x-axis range (zoom), so no
data is cut off - the viewer can pan/zoom out to reveal the entire history.
"""
quantile_levels = tuple(result.quantile_levels)
context = np.asarray(result.context[0], dtype=np.float32)
context_len = len(context)
positions, date_labels = _as_positions(x)
if baseline_result is None:
fig = make_subplots(rows=1, cols=1)
plot_forecast(
context=context,
forecasts=result.quantiles[0],
ground_truth=ground_truth,
x=positions,
quantile_levels=quantile_levels,
max_context_to_show=max_context_to_show,
engine="plotly",
fig=fig,
row=1,
col=1,
)
n_rows = 1
else:
cov_rows = list(_iter_covariate_rows(result, positions))
n_cov = len(cov_rows)
cov_heights = [0.32 / n_cov] * n_cov if n_cov else []
fig = make_subplots(
rows=2 + n_cov,
cols=1,
shared_xaxes=True,
vertical_spacing=0.06,
row_heights=[0.34, 0.34, *cov_heights],
row_titles=["Univariate", "Multivariate",
*(lbl for lbl, _, _ in cov_rows)],
)
for row, forecast in ((1, baseline_result.quantiles[0]), (2, result.quantiles[0])):
plot_forecast(
context=context,
forecasts=forecast,
ground_truth=ground_truth,
x=positions,
quantile_levels=quantile_levels,
max_context_to_show=max_context_to_show,
engine="plotly",
fig=fig,
row=row,
col=1,
)
for i, (label, xc, yc) in enumerate(cov_rows):
plot_covariate(
yc, x=xc, label=label, engine="plotly", fig=fig, row=i + 3, col=1,
color=COVARIATE_COLORS[i % len(COVARIATE_COLORS)],
)
n_rows = 2 + n_cov
# Enforce the zoom window as a shared axis range on *every* row (target and covariate
# panels alike) without dropping any data. This is what keeps context and covariates
# from being cut off: all points remain plotted, only the initial view is narrowed.
start = max(0, context_len - max_context_to_show) if max_context_to_show else 0
fig.update_xaxes(range=[positions[start], positions[-1]], autorange=False)
if date_labels is not None:
_apply_date_ticks(fig, positions, date_labels)
return fig, n_rows
def build_dataset_figure(x, y, *, label: str):
"""Plot a single selected series over time - a dataset preview with no forecast.
Shown as soon as a dataset/target is chosen, before (and regardless of) any run,
so users can eyeball the raw series. Uses the same datetime-axis handling and brand
chrome as the forecast figure for a consistent look.
"""
positions, date_labels = _as_positions(x)
y = np.asarray(y, dtype=np.float32)
positions = positions[: len(y)]
fig = go.Figure()
fig.add_trace(go.Scatter(
x=positions, y=y[: len(positions)], mode="lines", name=label,
line=dict(color=C_HISTORY, width=1.6),
))
if date_labels is not None:
_apply_date_ticks(fig, positions, date_labels)
fig.update_layout(
template=PLOTLY_TEMPLATE, title="", hovermode="x unified", font=FONT,
height=360, paper_bgcolor="rgba(0,0,0,0)", plot_bgcolor="rgba(0,0,0,0)",
margin=dict(t=48, b=40, l=54, r=20),
legend=dict(orientation="h", yanchor="bottom", y=1.03, xanchor="left", x=0),
)
_polish_axes(fig)
fig.update_xaxes(title_text="time")
return fig
def style_forecast_figure(fig, n_rows: int) -> None:
"""Apply the shared brand template, legend, and per-row axis chrome in place."""
fig.update_layout(
template=PLOTLY_TEMPLATE,
title="",
hovermode="x unified",
font=FONT,
height=300 + 150 * (n_rows - 1),
paper_bgcolor="rgba(0,0,0,0)",
plot_bgcolor="rgba(0,0,0,0)",
margin=dict(t=72, b=40, l=54, r=20),
legend=dict(orientation="h", yanchor="bottom", y=1.03, xanchor="left", x=0),
)
_polish_axes(fig)
for row in range(1, n_rows):
fig.update_xaxes(showticklabels=False, title_text="", row=row, col=1)
fig.update_xaxes(showticklabels=True, title_text="time", row=n_rows, col=1)