"""Plotting utilities using Plotly for the Gradio app.""" from __future__ import annotations from typing import Iterable, Optional, Sequence import numpy as np import plotly.graph_objects as go from plotly.subplots import make_subplots _BASE_COLORS = [ "#1f77b4", "#ff7f0e", "#2ca02c", "#d62728", "#9467bd", "#8c564b", "#e377c2", "#7f7f7f", "#bcbd22", "#17becf", "#aec7e8", "#ffbb78", "#98df8a", "#ff9896", "#c5b0d5", ] # Neutral alternating grays for CPD predictions (no state semantics) _CPD_GRAYS = ["#555555", "#aaaaaa"] def _adjust_color_hex(hex_color: str, factor: float) -> str: hex_color = hex_color.lstrip("#") r, g, b = (int(hex_color[i : i + 2], 16) for i in (0, 2, 4)) r = min(255, max(0, int(r + (255 - r) * factor))) g = min(255, max(0, int(g + (255 - g) * factor))) b = min(255, max(0, int(b + (255 - b) * factor))) return f"#{r:02x}{g:02x}{b:02x}" def build_segment_color_map(n_segments: int, n_dims: int) -> list[list[str]]: colors = [] for idx in range(n_segments): base = _BASE_COLORS[idx % len(_BASE_COLORS)] segment_colors = [_adjust_color_hex(base, min(0.15 * dim, 0.85)) for dim in range(n_dims)] colors.append(segment_colors) return colors def build_cpd_color_map(n_segments: int, n_dims: int) -> list[list[str]]: """Build an alternating neutral-gray color map for CPD predictions. Uses two distinguishable gray shades that alternate across segments, avoiding any implication of state identity. """ colors = [] for idx in range(n_segments): base = _CPD_GRAYS[idx % len(_CPD_GRAYS)] segment_colors = [_adjust_color_hex(base, min(0.10 * dim, 0.85)) for dim in range(n_dims)] colors.append(segment_colors) return colors def _iter_segments(states: Sequence[int]) -> Iterable[tuple[int, int, int]]: start = 0 current = states[0] for idx in range(1, len(states)): if states[idx] != current: yield start, idx, current start = idx current = states[idx] yield start, len(states), current def make_annotation_figure( signal: np.ndarray, change_points: Sequence[int], states: Sequence[int], title: str, color_map: list[list[str]], ) -> go.Figure: """Build a Plotly figure highlighting segments and change points.""" n_samples, n_dims = signal.shape fig = go.Figure() for start, stop, state in _iter_segments(states): for dim in range(n_dims): color = color_map[state % len(color_map)][dim % len(color_map[state % len(color_map)])] fig.add_trace( go.Scatter( x=np.arange(start, stop), y=signal[start:stop, dim], mode="lines", line=dict(color=color, width=2), name=f"State {state} · Dim {dim}", showlegend=False, ) ) for cp in change_points: fig.add_vline(x=cp, line=dict(color="#333333", width=1, dash="dash")) fig.update_layout( title=title, xaxis_title="Time", yaxis_title="Value", template="plotly_white", height=300 + 30 * n_dims, margin=dict(t=40, b=40, l=50, r=20), ) return fig def make_signal_figure(signal: np.ndarray, title: str = "Signal") -> go.Figure: """Simple signal visualization without state coloring.""" n_samples, n_dims = signal.shape fig = go.Figure() for dim in range(n_dims): fig.add_trace( go.Scatter( x=np.arange(n_samples), y=signal[:, dim], mode="lines", name=f"Dim {dim}", line=dict(color=_BASE_COLORS[dim % len(_BASE_COLORS)], width=1.5), ) ) fig.update_layout( title=title, xaxis_title="Time", yaxis_title="Value", template="plotly_white", height=300, margin=dict(t=40, b=40, l=50, r=20), ) return fig def _make_cpd_prediction_figure( signal: np.ndarray, pred_cps: Sequence[int], ) -> go.Figure: """Build a prediction subplot with alternating neutral grays for CPD. Does NOT imply state identity — only visually marks segment boundaries. """ n_samples, n_dims = signal.shape # Build alternating segments from CPs boundaries = [0] + sorted(pred_cps) + [n_samples] n_segs = len(boundaries) - 1 cpd_cmap = build_cpd_color_map(n_segs, n_dims) fig = go.Figure() for seg_idx, (start, stop) in enumerate(zip(boundaries[:-1], boundaries[1:])): for dim in range(n_dims): color = cpd_cmap[seg_idx % len(cpd_cmap)][dim % len(cpd_cmap[seg_idx % len(cpd_cmap)])] fig.add_trace( go.Scatter( x=np.arange(start, stop), y=signal[start:stop, dim], mode="lines", line=dict(color=color, width=2), showlegend=False, ) ) for cp in pred_cps: fig.add_vline(x=cp, line=dict(color="#333333", width=1, dash="dash")) return fig def make_comparison_figure( signal: np.ndarray, truth_states: Optional[Sequence[int]], truth_cps: Optional[Sequence[int]], pred_states: Optional[Sequence[int]], pred_cps: Optional[Sequence[int]], color_map: list[list[str]], *, fallback_states: Optional[Sequence[int]] = None, task: str = "change_point_detection", ) -> go.Figure: """Two-row figure contrasting ground-truth and predictions. For CPD-only algorithms (task='change_point_detection' and no native pred_states), the prediction row uses neutral alternating grays to avoid implying state identity. """ has_truth = truth_states is not None n_rows = 2 if has_truth else 1 titles = ["Ground Truth", "Prediction"] if has_truth else ["Prediction"] fig = make_subplots( rows=n_rows, cols=1, shared_xaxes=True, vertical_spacing=0.08, subplot_titles=titles, ) if has_truth: truth_fig = make_annotation_figure(signal, truth_cps or [], truth_states, "", color_map) for trace in truth_fig.data: fig.add_trace(trace, row=1, col=1) for shape in truth_fig.layout.shapes or []: fig.add_shape(shape, row=1, col=1) pred_row = 2 if has_truth else 1 is_cpd = task == "change_point_detection" if pred_states is not None and not is_cpd: # State Detection: use full state colors pred_fig = make_annotation_figure(signal, pred_cps or [], pred_states, "", color_map) for trace in pred_fig.data: fig.add_trace(trace, row=pred_row, col=1) for shape in pred_fig.layout.shapes or []: fig.add_shape(shape, row=pred_row, col=1) elif pred_cps: # CPD or CPD-only: neutral alternating grays pred_fig = _make_cpd_prediction_figure(signal, pred_cps) for trace in pred_fig.data: fig.add_trace(trace, row=pred_row, col=1) for shape in pred_fig.layout.shapes or []: fig.add_shape(shape, row=pred_row, col=1) elif fallback_states is not None: pred_fig = make_annotation_figure(signal, pred_cps or [], fallback_states, "", color_map) for trace in pred_fig.data: fig.add_trace(trace, row=pred_row, col=1) for shape in pred_fig.layout.shapes or []: fig.add_shape(shape, row=pred_row, col=1) else: # No predicted change points and no fallback states: render the signal # as a single neutral segment so the dimensions don't get distinct # state-like colors (which would falsely imply "1 state per channel"). n_dims = signal.shape[1] neutral_cmap = build_cpd_color_map(1, n_dims) for dim in range(n_dims): fig.add_trace( go.Scatter( x=np.arange(signal.shape[0]), y=signal[:, dim], mode="lines", name=f"Dim {dim}", line=dict(color=neutral_cmap[0][dim], width=2), showlegend=False, ), row=pred_row, col=1, ) height = 350 * n_rows fig.update_layout(height=height, template="plotly_white", margin=dict(t=50, b=40, l=50, r=20)) fig.update_yaxes(title_text="Value", row=1, col=1) if n_rows > 1: fig.update_yaxes(title_text="Value", row=2, col=1) fig.update_xaxes(title_text="Time", row=n_rows, col=1) return fig def make_multi_comparison_figure( signal: np.ndarray, runs: list[dict], truth_states: Optional[Sequence[int]] = None, truth_cps: Optional[Sequence[int]] = None, ) -> go.Figure: """Multi-row figure comparing several algorithm runs. Parameters ---------- signal : np.ndarray The input signal, shape (n_timestamps, n_dims). runs : list of dict Each dict has keys: 'name', 'pred_states', 'pred_cps'. truth_states, truth_cps : optional ground truth. """ has_truth = truth_states is not None n_rows = len(runs) + (1 if has_truth else 0) if n_rows == 0: return go.Figure() n_states_max = 1 if has_truth: n_states_max = max(n_states_max, int(np.max(truth_states)) + 1) for run in runs: if run.get("pred_states") is not None: n_states_max = max(n_states_max, int(np.max(run["pred_states"])) + 1) color_map = build_segment_color_map(n_states_max, signal.shape[1]) titles = [] if has_truth: titles.append("Ground Truth") for run in runs: titles.append(run["name"]) fig = make_subplots( rows=n_rows, cols=1, shared_xaxes=True, vertical_spacing=0.04, subplot_titles=titles, ) row = 1 if has_truth: truth_fig = make_annotation_figure(signal, truth_cps or [], truth_states, "", color_map) for trace in truth_fig.data: fig.add_trace(trace, row=row, col=1) for shape in truth_fig.layout.shapes or []: fig.add_shape(shape, row=row, col=1) row += 1 for run in runs: ps = run.get("pred_states") pc = run.get("pred_cps", []) run_task = run.get("task", "change_point_detection") is_sd = run_task == "state_detection" if ps is not None and is_sd: # State Detection run: use full state colors run_fig = make_annotation_figure(signal, pc, ps, "", color_map) for trace in run_fig.data: fig.add_trace(trace, row=row, col=1) for shape in run_fig.layout.shapes or []: fig.add_shape(shape, row=row, col=1) elif pc: # CPD run: neutral alternating grays run_fig = _make_cpd_prediction_figure(signal, pc) for trace in run_fig.data: fig.add_trace(trace, row=row, col=1) for shape in run_fig.layout.shapes or []: fig.add_shape(shape, row=row, col=1) else: # No predicted change points: render the signal as a single # neutral segment to avoid implying that each dimension is a # distinct state. n_dims = signal.shape[1] neutral_cmap = build_cpd_color_map(1, n_dims) for dim in range(n_dims): fig.add_trace( go.Scatter( x=np.arange(signal.shape[0]), y=signal[:, dim], mode="lines", line=dict(color=neutral_cmap[0][dim], width=2), showlegend=False, ), row=row, col=1, ) row += 1 fig.update_layout( height=max(250 * n_rows, 400), template="plotly_white", margin=dict(t=50, b=40, l=50, r=20), showlegend=False, ) fig.update_xaxes(title_text="Time", row=n_rows, col=1) return fig def make_metrics_bar_chart(scores_df) -> go.Figure: """Bar chart comparing metric scores across algorithms. Colors are grouped semantically: - F1-related metrics → blue palette - Covering metrics → green palette - State detection metrics → purple/orange palette """ import pandas as pd if not isinstance(scores_df, pd.DataFrame) or scores_df.empty: return go.Figure() # Semantic color mapping for metrics _METRIC_COLORS = { # F1-related → blues "F1 Score": "#1f77b4", "F1 Score (precision)": "#4a9ecf", "F1 Score (recall)": "#7ec4e8", # Covering → greens "Covering": "#2ca02c", "Bidirectional Covering": "#5cc85c", # State metrics → warm tones "Adjusted Rand Index": "#d62728", "Normalized Mutual Information": "#e87d2e", "Adjusted Mutual Information": "#f0a050", "Weighted ARI": "#9467bd", "Weighted NMI": "#b694d4", "State Matching Score": "#8c564b", } fig = go.Figure() metric_cols = [c for c in scores_df.columns if c != "Algorithm"] for i, metric in enumerate(metric_cols): color = _METRIC_COLORS.get(metric, _BASE_COLORS[i % len(_BASE_COLORS)]) fig.add_trace( go.Bar( x=scores_df["Algorithm"], y=scores_df[metric], name=metric, marker_color=color, ) ) fig.update_layout( barmode="group", title="Metric Scores Comparison", xaxis_title="Algorithm", yaxis_title="Score", template="plotly_white", height=400, margin=dict(t=50, b=40, l=50, r=20), ) return fig