Download src/plotting.py from fchavelli/tsseg: direct link, hf CLI and curl.
- Browser
- Download file 13.9 kB
-
https://huggingface.co/spaces/fchavelli/tsseg/resolve/main/src/plotting.py
- Command line
-
hf download hf://spaces/fchavelli/tsseg/src/plotting.py
-
curl -L -o plotting.py https://huggingface.co/spaces/fchavelli/tsseg/resolve/main/src/plotting.py
13.9 kB
| """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 | |