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