Spaces:
Sleeping
Sleeping
Download src/plots.py from lokahq/dna-benchmark: direct link, hf CLI and curl.
- Browser
- Download file 42.9 kB
-
https://huggingface.co/spaces/lokahq/dna-benchmark/resolve/main/src/plots.py
- Command line
-
hf download hf://spaces/lokahq/dna-benchmark/src/plots.py
-
curl -L -o plots.py https://huggingface.co/spaces/lokahq/dna-benchmark/resolve/main/src/plots.py
42.9 kB
| """Plotly figure factories for the DNA benchmark leaderboard.""" | |
| import json | |
| import numpy as np | |
| import pandas as pd | |
| import plotly.graph_objects as go | |
| import plotly.io as pio | |
| from plotly.subplots import make_subplots | |
| from src.constants import ( | |
| CATEGORY_COLORS, | |
| CATEGORY_DISPLAY_NAMES, | |
| CLASSIFICATION_METRICS, | |
| TASK_SEQ_LENGTHS, | |
| ) | |
| # Model colour palette — Loka-brand-anchored 8-hue categorical set (same family as | |
| # CATEGORY_COLORS, fixed order), used for model/variant series so the two never collide. | |
| _MODEL_PALETTE = [ | |
| "#1877F2", "#EB6834", "#0AA88F", "#D9A61F", "#E87BA4", "#008300", "#4A3AA7", "#E34948", | |
| ] | |
| def _bar_marker(color: str) -> dict: | |
| """Softened bar fill (opacity over white background) so grouped-bar charts read as | |
| pastel like the rest of the dashboard's charts, rather than solid/saturated blocks.""" | |
| return dict(color=color, opacity=0.78, line=dict(width=1, color="white")) | |
| # Brand chart chrome — registered as an additive Plotly template so every figure | |
| # picks up Loka typography/ink without touching each figure's own layout options. | |
| pio.templates["loka"] = go.layout.Template( | |
| layout=go.Layout( | |
| font=dict(family="Inter, -apple-system, 'Segoe UI', Roboto, sans-serif", color="#050517"), | |
| title=dict(font=dict(size=16, color="#050517")), | |
| colorway=_MODEL_PALETTE, | |
| paper_bgcolor="#ffffff", | |
| plot_bgcolor="#ffffff", | |
| ) | |
| ) | |
| _TEMPLATE = "plotly_white+loka" | |
| # Shared horizontal legend layout used across multiple figures. | |
| _LEGEND_H = dict( | |
| orientation="h", yanchor="bottom", y=1.02, xanchor="right", x=1, | |
| itemclick="toggle", itemdoubleclick="toggleothers", | |
| ) | |
| # Speed/resource column metadata: col -> (display label, higher_x_is_better) | |
| _SPEED_META: dict[str, tuple[str, bool]] = { | |
| "test_throughput_seq_s": ("Throughput (seq/s)", True), | |
| "test_embedding_time_s": ("Embedding Time (s)", False), | |
| "peak_vram_extraction_mb": ("Peak VRAM (MB)", False), | |
| "vram_model_mb": ("VRAM (MB)", False), | |
| } | |
| def _model_color_map(models: list[str]) -> dict[str, str]: | |
| return {m: _MODEL_PALETTE[i % len(_MODEL_PALETTE)] for i, m in enumerate(sorted(models))} | |
| def _pareto_mask(x: np.ndarray, y: np.ndarray, higher_x_better: bool) -> np.ndarray: | |
| """Return boolean mask: True = point is not dominated on (x, y). | |
| Dominated means some other point is at least as good on both axes and strictly | |
| better on at least one. y is always "higher is better" (MCC/metric). | |
| """ | |
| n = len(x) | |
| dominated = np.zeros(n, dtype=bool) | |
| for i in range(n): | |
| for j in range(n): | |
| if i == j: | |
| continue | |
| x_ge = (x[j] >= x[i]) if higher_x_better else (x[j] <= x[i]) | |
| y_ge = y[j] >= y[i] | |
| x_gt = (x[j] > x[i]) if higher_x_better else (x[j] < x[i]) | |
| y_gt = y[j] > y[i] | |
| if x_ge and y_ge and (x_gt or y_gt): | |
| dominated[i] = True | |
| break | |
| return ~dominated | |
| def _norm(val) -> str: | |
| """Normalise a possibly-null field to empty string.""" | |
| if val is None or (isinstance(val, float) and pd.isna(val)): | |
| return "" | |
| return str(val) | |
| def _fmt_layer(params_str) -> str: | |
| """Short label for layer selection from layer_selection_params JSON string.""" | |
| if params_str is None or (isinstance(params_str, float) and pd.isna(params_str)): | |
| return "default" | |
| try: | |
| params = json.loads(params_str) if isinstance(params_str, str) else params_str | |
| if not params: | |
| return "default" | |
| idx = params.get("layer_index") | |
| strategy = params.get("strategy") or params.get("layer_selection_strategy") | |
| parts = [] | |
| if strategy: | |
| parts.append(str(strategy)) | |
| if idx is not None: | |
| parts.append(f"({int(idx)})") | |
| return " ".join(parts) if parts else "default" | |
| except (json.JSONDecodeError, TypeError): | |
| return "default" | |
| def _variant_label(row: pd.Series) -> str: | |
| """Human-readable config label: 'short-hf-id · pooling · layer · head'.""" | |
| parts = [] | |
| hf = row.get("model_hf_id") | |
| if hf is not None and not (isinstance(hf, float) and pd.isna(hf)): | |
| parts.append(str(hf).rsplit("/", 1)[-1]) | |
| pool = row.get("pooling_strategy") | |
| if pool is not None and not (isinstance(pool, float) and pd.isna(pool)): | |
| parts.append(str(pool)) | |
| parts.append(_fmt_layer(row.get("layer_selection_params"))) | |
| head = row.get("head_type") | |
| if head is not None and not (isinstance(head, float) and pd.isna(head)): | |
| parts.append(str(head)) | |
| return " · ".join(parts) | |
| def _variant_key_col(df: pd.DataFrame) -> pd.Series: | |
| """Return a Series of stable variant key strings (model_hf_id|pooling|layer|head).""" | |
| def _col(name: str) -> pd.Series: | |
| return df[name].map(_norm) if name in df.columns else pd.Series("", index=df.index) | |
| layer = ( | |
| df["layer_selection_params"].apply(_fmt_layer) | |
| if "layer_selection_params" in df.columns | |
| else pd.Series("default", index=df.index) | |
| ) | |
| return _col("model_hf_id") + "|" + _col("pooling_strategy") + "|" + layer + "|" + _col("head_type") | |
| def _variant_hover_lines(vdf: pd.DataFrame) -> str: | |
| """Multi-line config block for heatmap hover, sourced from first row of variant group.""" | |
| row = vdf.iloc[0] | |
| lines = [] | |
| hf = row.get("model_hf_id") | |
| if hf is not None and not (isinstance(hf, float) and pd.isna(hf)): | |
| lines.append(f"HF model: {hf}") | |
| pool = row.get("pooling_strategy") | |
| if pool is not None and not (isinstance(pool, float) and pd.isna(pool)): | |
| lines.append(f"Pooling: {pool}") | |
| lines.append(f"Layer: {_fmt_layer(row.get('layer_selection_params'))}") | |
| head = row.get("head_type") | |
| if head is not None and not (isinstance(head, float) and pd.isna(head)): | |
| lines.append(f"Head: {head}") | |
| return "<br>".join(lines) | |
| # --------------------------------------------------------------------------- | |
| # 0. Top-N / full-coverage filter helpers | |
| # --------------------------------------------------------------------------- | |
| def filter_full_coverage(df: pd.DataFrame, metric: str) -> pd.DataFrame: | |
| """Keep only variants that report `metric` for every task anyone reports it for. | |
| Group means over a partial task set can look stronger than a fully-evaluated | |
| variant's mean simply because the missing tasks happen to be hard ones — this | |
| filter removes that confound before any downstream aggregation. | |
| """ | |
| if df.empty or metric not in df.columns or "task_name" not in df.columns: | |
| return df | |
| universe = { | |
| t for t in df["task_name"].dropna().unique() | |
| if df[df["task_name"] == t][metric].notna().any() | |
| } | |
| if not universe: | |
| return df | |
| tmp = df.copy() | |
| tmp["_vkey"] = _variant_key_col(tmp) | |
| reported = tmp[tmp[metric].notna()].groupby("_vkey")["task_name"].apply(set) | |
| full_keys = {vkey for vkey, tasks in reported.items() if universe.issubset(tasks)} | |
| return df[tmp["_vkey"].isin(full_keys)].reset_index(drop=True) | |
| def top_n_variants(df: pd.DataFrame, metric: str, n: int) -> pd.DataFrame: | |
| """Return df filtered to the top-N variants by mean metric (all tasks of those variants kept).""" | |
| if not n or df.empty or metric not in df.columns: | |
| return df | |
| tmp = df.copy() | |
| tmp["_vkey"] = _variant_key_col(tmp) | |
| # dropna first: nlargest on a nullable Float64 Series pads with NA once real values run | |
| # out (e.g. a metric only a few variants report, like AUPRC), pulling in variants that | |
| # never reported this metric at all instead of just returning fewer real ones. | |
| means = tmp.groupby("_vkey")[metric].mean().dropna() | |
| top_keys = set(means.nlargest(n).index) | |
| return df[tmp["_vkey"].isin(top_keys)].reset_index(drop=True) | |
| # --------------------------------------------------------------------------- | |
| # 1. Ranking bar — variant leaderboard as horizontal bar chart | |
| # --------------------------------------------------------------------------- | |
| def ranking_bar(df: pd.DataFrame, metric: str) -> go.Figure: | |
| """Horizontal bar chart: one bar per variant sorted by Overall mean, legend by model family.""" | |
| if df.empty or metric not in df.columns: | |
| return go.Figure() | |
| df = df.copy() | |
| df["_vkey"] = _variant_key_col(df) | |
| models = sorted(df["model_alias"].unique()) if "model_alias" in df.columns else [] | |
| model_color_map = _model_color_map(models) | |
| metric_label = CLASSIFICATION_METRICS.get(metric, metric) | |
| records = [] | |
| for vkey, group in df.groupby("_vkey"): | |
| first = group.iloc[0] | |
| model = _norm(first.get("model_alias", "")) | |
| label = _variant_label(first) | |
| overall_vals = group[metric].dropna() | |
| overall = round(float(overall_vals.mean()), 3) if not overall_vals.empty else None | |
| breakdown = {} | |
| for cat_key, cat_label in CATEGORY_DISPLAY_NAMES.items(): | |
| if cat_key == "other": | |
| continue | |
| cat_data = group[group["task_category"] == cat_key][metric].dropna() | |
| breakdown[cat_label] = round(float(cat_data.mean()), 3) if not cat_data.empty else None | |
| records.append({"label": label, "model": model, "overall": overall, "breakdown": breakdown}) | |
| records.sort(key=lambda r: r["overall"]) # ascending → best at top (plotly y goes bottom-to-top) | |
| y_order = [r["label"] for r in records] | |
| fig = go.Figure() | |
| legend_shown: set[str] = set() | |
| for rec in records: | |
| m = rec["model"] | |
| color = model_color_map.get(m, _MODEL_PALETTE[0]) | |
| show_legend = m not in legend_shown | |
| legend_shown.add(m) | |
| breakdown_lines = "<br>".join( | |
| f"{k}: {v:.3f}" for k, v in rec["breakdown"].items() if v is not None | |
| ) | |
| hover = f"<b>{rec['label']}</b><br>Overall: {rec['overall']:.3f}" | |
| if breakdown_lines: | |
| hover += "<br>" + breakdown_lines | |
| fig.add_trace(go.Bar( | |
| y=[rec["label"]], | |
| x=[rec["overall"]], | |
| orientation="h", | |
| name=m, | |
| marker_color=color, | |
| legendgroup=m, | |
| showlegend=show_legend, | |
| hovertemplate=hover + "<extra></extra>", | |
| )) | |
| n = len(records) | |
| height = max(400, n * 26 + 120) | |
| left_margin = min(max((len(lbl) for lbl in y_order), default=20) * 7, 380) | |
| right_margin = min(max((len(m) for m in models), default=8) * 8 + 40, 200) | |
| fig.update_layout( | |
| title=f"Variant ranking — {metric_label}", | |
| xaxis_title=metric_label, | |
| height=height, | |
| template=_TEMPLATE, | |
| # Vertical, outside-right legend — model count is unbounded (unlike fixed | |
| # category counts elsewhere), so a horizontal legend can wrap onto the plot. | |
| legend=dict(orientation="v", yanchor="top", y=1, xanchor="left", x=1.02), | |
| margin=dict(l=left_margin, r=right_margin, t=60, b=40), | |
| ) | |
| fig.update_yaxes( | |
| categoryorder="array", | |
| categoryarray=y_order, | |
| automargin=True, | |
| ) | |
| return fig | |
| # --------------------------------------------------------------------------- | |
| # 2. Strip plot — per-task scores per variant, coloured by task category | |
| # --------------------------------------------------------------------------- | |
| def strip_plot(df: pd.DataFrame, metric: str) -> go.Figure: | |
| """Strip/dot plot: y = variant (rank 1 at top), x = metric per task, dots by task category.""" | |
| if df.empty or metric not in df.columns: | |
| return go.Figure() | |
| df = df.copy() | |
| df["_vkey"] = _variant_key_col(df) | |
| metric_label = CLASSIFICATION_METRICS.get(metric, metric) | |
| # Order variants globally by mean metric descending (rank 1 = i=0 = top) | |
| vkey_means = df.groupby("_vkey")[metric].mean().sort_values(ascending=False) | |
| ordered_vkeys = vkey_means.index.tolist() | |
| vkey_to_i = {vk: i for i, vk in enumerate(ordered_vkeys)} | |
| fig = go.Figure() | |
| shown_cats: set[str] = set() | |
| rng = np.random.default_rng(seed=42) | |
| for cat, cat_df in df.groupby("task_category"): | |
| color = CATEGORY_COLORS.get(cat, CATEGORY_COLORS["other"]) | |
| cat_label = CATEGORY_DISPLAY_NAMES.get(cat, cat) | |
| show_legend = cat not in shown_cats | |
| shown_cats.add(cat) | |
| y_base = cat_df["_vkey"].map(vkey_to_i).astype(float) | |
| y_jitter = y_base + rng.uniform(-0.3, 0.3, len(cat_df)) | |
| fig.add_trace(go.Scatter( | |
| x=cat_df[metric], | |
| y=y_jitter, | |
| mode="markers", | |
| name=cat_label, | |
| marker=dict(color=color, size=6, opacity=0.6, line=dict(width=1, color="white")), | |
| text=cat_df["task_name"], | |
| hovertemplate=f"<b>%{{text}}</b><br>{metric_label}: %{{x:.3f}}<extra></extra>", | |
| legendgroup="cat_" + cat, | |
| showlegend=show_legend, | |
| )) | |
| # Short vertical tick at mean per variant | |
| for vk, i in vkey_to_i.items(): | |
| mean_val = float(vkey_means[vk]) | |
| fig.add_shape( | |
| type="line", | |
| x0=mean_val, x1=mean_val, | |
| y0=i - 0.35, y1=i + 0.35, | |
| line=dict(color="rgba(60,60,60,0.6)", width=2), | |
| ) | |
| tick_labels = [_variant_label(df[df["_vkey"] == vk].iloc[0]) for vk in ordered_vkeys] | |
| n = len(ordered_vkeys) | |
| height = max(500, n * 22 + 120) | |
| fig.update_layout( | |
| title=f"{metric_label} per task by variant — rank 1 at top", | |
| xaxis_title=metric_label, | |
| height=height, | |
| template=_TEMPLATE, | |
| legend=_LEGEND_H, | |
| ) | |
| fig.update_yaxes( | |
| tickmode="array", | |
| tickvals=list(range(n)), | |
| ticktext=tick_labels, | |
| autorange="reversed", # i=0 (rank 1) at top | |
| automargin=True, | |
| ) | |
| return fig | |
| # --------------------------------------------------------------------------- | |
| # 3. Heatmap — single model: all variants as rows; multi-model: best variant per family | |
| # --------------------------------------------------------------------------- | |
| def heatmap_variants( | |
| df: pd.DataFrame, | |
| metric: str, | |
| model_alias: str | list[str], | |
| aggregate_by_dataset: bool = False, | |
| show_all_variants: bool = False, | |
| ) -> go.Figure: | |
| """Heatmap: variant rows vs task columns. | |
| Single model: rows = all config variants. | |
| Multiple models: rows = best variant per family (highest mean metric for that family), | |
| unless show_all_variants is True, in which case every variant of every selected model is a row. | |
| """ | |
| metric_label = CLASSIFICATION_METRICS.get(metric, metric) | |
| if isinstance(model_alias, str): | |
| model_aliases = [model_alias] if model_alias else [] | |
| else: | |
| model_aliases = [m for m in (model_alias or []) if m] | |
| if not model_aliases or df.empty or metric not in df.columns: | |
| return go.Figure() | |
| multi_model = len(model_aliases) > 1 | |
| mdf = ( | |
| df[df["model_alias"].isin(model_aliases)] | |
| if multi_model | |
| else df[df["model_alias"] == model_aliases[0]] | |
| ).copy() | |
| if mdf.empty: | |
| return go.Figure() | |
| mdf["_vkey"] = _variant_key_col(mdf) | |
| if multi_model and not show_all_variants: | |
| rows_iter = [] | |
| for m in model_aliases: | |
| family_df = mdf[mdf["model_alias"] == m] | |
| if family_df.empty: | |
| continue | |
| # Pick the variant with highest mean metric for this family. A variant/head | |
| # (e.g. ridge, regression-only) that never reports this metric has an all-NaN | |
| # group mean — idxmax would raise, and it'd render as an all-blank row anyway. | |
| vkey_means = family_df.groupby("_vkey")[metric].mean().dropna() | |
| if vkey_means.empty: | |
| continue | |
| best_vkey = vkey_means.idxmax() | |
| best_df = family_df[family_df["_vkey"] == best_vkey] | |
| rows_iter.append((_variant_label(best_df.iloc[0]), best_df)) | |
| else: | |
| configs = sorted( | |
| c for c in mdf["_vkey"].unique() if mdf[mdf["_vkey"] == c][metric].notna().any() | |
| ) | |
| rows_iter = [ | |
| (_variant_label(mdf[mdf["_vkey"] == c].iloc[0]), mdf[mdf["_vkey"] == c]) | |
| for c in configs | |
| ] | |
| # A task/dataset only gets a column if the selected metric is ever reported for it — | |
| # e.g. bulk_rna_expression is regression-only (r2_test), so it never has mcc_test/ | |
| # accuracy_test/etc. and would otherwise render as an all-blank column. | |
| if aggregate_by_dataset: | |
| col_keys = ( | |
| sorted( | |
| d for d in mdf["dataset_name"].dropna().unique() | |
| if df[df["dataset_name"] == d][metric].notna().any() | |
| ) | |
| if "dataset_name" in mdf.columns | |
| else [] | |
| ) | |
| col_names = col_keys | |
| def _cell_vals(vdf: pd.DataFrame, col_key: str) -> pd.Series: | |
| return vdf[vdf["dataset_name"] == col_key][metric].dropna() | |
| else: | |
| col_keys, col_names = [], [] | |
| for cat in CATEGORY_DISPLAY_NAMES: | |
| for task in sorted(mdf[mdf["task_category"] == cat]["task_name"].unique()): | |
| if not df[df["task_name"] == task][metric].notna().any(): | |
| continue | |
| col_keys.append(task) | |
| col_names.append(task) | |
| def _cell_vals(vdf: pd.DataFrame, col_key: str) -> pd.Series: | |
| return vdf[vdf["task_name"] == col_key][metric].dropna() | |
| if not rows_iter: | |
| return go.Figure() | |
| col_names = col_names + ["Overall"] | |
| rows_z, rows_text, rows_hover, row_labels = [], [], [], [] | |
| for label, vdf in rows_iter: | |
| z_row, t_row, h_row = [], [], [] | |
| config_lines = _variant_hover_lines(vdf) | |
| for col_key, col_name in zip(col_keys, col_names[:-1]): | |
| vals = _cell_vals(vdf, col_key) | |
| v = round(float(vals.mean()), 3) if not vals.empty else None | |
| z_row.append(v) | |
| t_row.append(f"{v:.3f}" if v is not None else "—") | |
| v_str = f"{v:.3f}" if v is not None else "N/A" | |
| h_row.append(f"<b>{col_name}</b><br>{metric_label}: {v_str}<br>{config_lines}") | |
| overall_vals = vdf[metric].dropna() | |
| overall = round(float(overall_vals.mean()), 3) if not overall_vals.empty else None | |
| z_row.append(overall) | |
| t_row.append(f"{overall:.3f}" if overall is not None else "—") | |
| ov_str = f"{overall:.3f}" if overall is not None else "N/A" | |
| h_row.append(f"<b>Overall</b><br>{metric_label}: {ov_str}<br>{config_lines}") | |
| rows_z.append(z_row) | |
| rows_text.append(t_row) | |
| rows_hover.append(h_row) | |
| row_labels.append(label) | |
| if len(rows_z) > 1: | |
| sort_idx = sorted(range(len(rows_z)), key=lambda i: rows_z[i][-1] or -1, reverse=True) | |
| rows_z = [rows_z[i] for i in sort_idx][::-1] | |
| rows_text = [rows_text[i] for i in sort_idx][::-1] | |
| rows_hover = [rows_hover[i] for i in sort_idx][::-1] | |
| row_labels = [row_labels[i] for i in sort_idx][::-1] | |
| zmin = -1 if metric == "mcc_test" else 0 | |
| n_data_cols = len(col_keys) | |
| cell_font_size = 9 if n_data_cols > 10 else 12 | |
| tick_angle = -45 if not aggregate_by_dataset else 0 | |
| fig = go.Figure( | |
| go.Heatmap( | |
| z=rows_z, | |
| x=col_names, | |
| y=row_labels, | |
| text=rows_text, | |
| texttemplate="%{text}", | |
| textfont=dict(size=cell_font_size, color="black"), | |
| hovertext=rows_hover, | |
| hovertemplate="%{hovertext}<extra></extra>", | |
| colorscale="RdYlGn", | |
| zmin=zmin, | |
| zmax=1, | |
| colorbar=dict(title=metric_label, thickness=14, nticks=5, tickformat=".2f", lenmode="pixels", len=220), | |
| hoverongaps=False, | |
| ) | |
| ) | |
| fig.add_shape( | |
| type="rect", | |
| x0=n_data_cols - 0.5, x1=n_data_cols + 0.5, | |
| y0=-0.5, y1=len(rows_z) - 0.5, | |
| line=dict(color="#333", width=2), | |
| fillcolor="rgba(0,0,0,0)", | |
| ) | |
| n_rows = len(rows_z) | |
| left_margin = min(max(len(lbl) for lbl in row_labels) * 7, 320) | |
| bottom_margin = 120 if not aggregate_by_dataset else 20 | |
| t_margin = 120 if not aggregate_by_dataset else 80 | |
| v_overhead = t_margin + bottom_margin | |
| height = max(v_overhead + 80, n_rows * 48 + v_overhead + 20) | |
| title = f"{metric_label} — {', '.join(model_aliases)}" | |
| if multi_model: | |
| title += " (best variant per family)" | |
| fig.update_layout( | |
| title=title, | |
| height=height, | |
| template=_TEMPLATE, | |
| xaxis=dict(side="top", tickfont=dict(size=9 if not aggregate_by_dataset else 11), tickangle=tick_angle), | |
| yaxis=dict(tickfont=dict(size=10), automargin=True), | |
| margin=dict(l=left_margin, r=20, t=t_margin, b=bottom_margin), | |
| ) | |
| return fig | |
| # --------------------------------------------------------------------------- | |
| # 4 & 5. Grouped bar chart — shared helper + two public wrappers | |
| # --------------------------------------------------------------------------- | |
| def _grouped_bar( | |
| df: pd.DataFrame, | |
| metric: str, | |
| x_items: list, | |
| x_labels: list[str], | |
| get_vals, # Callable[[variant_df, item], pd.Series] | |
| title: str, | |
| legend_by_model: bool = False, | |
| ) -> go.Figure: | |
| """Grouped bar: one bar series per variant, coloured by model family. Sorted by overall mean. | |
| legend_by_model=True collapses legend entries to model alias (one entry per family). | |
| """ | |
| df = df.copy() | |
| df["_vkey"] = _variant_key_col(df) | |
| models = sorted(df["model_alias"].unique()) if "model_alias" in df.columns else [] | |
| model_color_map = _model_color_map(models) | |
| metric_label = CLASSIFICATION_METRICS.get(metric, metric) | |
| vkey_order = df.groupby("_vkey")[metric].mean().sort_values(ascending=False).index.tolist() | |
| fig = go.Figure() | |
| legend_shown: set[str] = set() | |
| for i, vkey in enumerate(vkey_order): | |
| vdf = df[df["_vkey"] == vkey] | |
| first = vdf.iloc[0] | |
| model = _norm(first.get("model_alias", "")) | |
| label = _variant_label(first) | |
| y_vals = [ | |
| round(float(vals.mean()), 3) if not (vals := get_vals(vdf, item).dropna()).empty else None | |
| for item in x_items | |
| ] | |
| if legend_by_model: | |
| color = model_color_map.get(model, _MODEL_PALETTE[0]) | |
| show_legend = model not in legend_shown | |
| legend_shown.add(model) | |
| fig.add_trace(go.Bar( | |
| name=model, x=x_labels, y=y_vals, marker=_bar_marker(color), | |
| legendgroup=model, showlegend=show_legend, | |
| hovertemplate=f"<b>{label}</b><br>%{{x}}: %{{y:.3f}}<extra></extra>", | |
| )) | |
| else: | |
| # One series per variant (e.g. pooling/layer ablation of a single model) — | |
| # colour by variant, not by model, or every bar would share one colour. | |
| color = _MODEL_PALETTE[i % len(_MODEL_PALETTE)] | |
| fig.add_trace(go.Bar( | |
| name=label, x=x_labels, y=y_vals, marker=_bar_marker(color), | |
| hovertemplate=f"<b>{label}</b><br>%{{x}}: %{{y:.3f}}<extra></extra>", | |
| )) | |
| fig.update_layout( | |
| title=title, yaxis_title=metric_label, barmode="group", | |
| height=440, template=_TEMPLATE, legend=_LEGEND_H, | |
| ) | |
| return fig | |
| def bar_plot_per_category(df: pd.DataFrame, metric: str, model_alias: str | None = None) -> go.Figure: | |
| """Grouped bar: X = task category, bar series = variants coloured by model family. | |
| model_alias: restrict to one model family and give each pooling/layer variant its | |
| own legend entry (pooling/layer ablation view) instead of collapsing legend by model. | |
| """ | |
| if df.empty or metric not in df.columns: | |
| return go.Figure() | |
| if model_alias: | |
| df = df[df["model_alias"] == model_alias] | |
| if df.empty: | |
| return go.Figure() | |
| active_cats = { | |
| c for c in df["task_category"].dropna().unique() | |
| if df[df["task_category"] == c][metric].notna().any() | |
| } | |
| categories = [k for k in CATEGORY_DISPLAY_NAMES if k != "other" and k in active_cats] | |
| metric_label = CLASSIFICATION_METRICS.get(metric, metric) | |
| title = f"{metric_label} per task category" | |
| if model_alias: | |
| title += f" — {model_alias} (by pooling/layer variant)" | |
| return _grouped_bar( | |
| df, metric, | |
| x_items=categories, | |
| x_labels=[CATEGORY_DISPLAY_NAMES[c] for c in categories], | |
| get_vals=lambda mdf, cat: mdf[mdf["task_category"] == cat][metric], | |
| title=title, | |
| legend_by_model=not model_alias, | |
| ) | |
| _ABLATION_DIMS = { | |
| "pooling_strategy": "Pooling", | |
| "_layer": "Layer", | |
| "head_type": "Head", | |
| } | |
| def bar_plot_pooling_ablation( | |
| df: pd.DataFrame, metric: str, model_alias: str, vary_by: str | |
| ) -> go.Figure: | |
| """Grouped bar isolating ONE config dimension for a single model, holding the other | |
| two dimensions fixed at whichever combo covers the most values of vary_by (falling | |
| back to the highest-scoring combo). E.g. vary_by='pooling_strategy' fixes layer and | |
| head_type, so the bars only differ by pooling — mean vs max, apples to apples. | |
| """ | |
| if df.empty or metric not in df.columns or not model_alias or vary_by not in _ABLATION_DIMS: | |
| return go.Figure() | |
| mdf = df[df["model_alias"] == model_alias].copy() | |
| if mdf.empty: | |
| return go.Figure() | |
| mdf["_layer"] = ( | |
| mdf["layer_selection_params"].apply(_fmt_layer) | |
| if "layer_selection_params" in mdf.columns | |
| else "" | |
| ) | |
| other_dims = [d for d in _ABLATION_DIMS if d != vary_by] | |
| mdf["_other_key"] = mdf[other_dims].astype(str).agg(" · ".join, axis=1) | |
| coverage = mdf.groupby("_other_key")[vary_by].nunique() | |
| combo_means = mdf.groupby("_other_key")[metric].mean() | |
| qualifying = coverage[coverage > 1].index | |
| candidates = combo_means[combo_means.index.isin(qualifying)].dropna() | |
| if candidates.empty: | |
| candidates = combo_means.dropna() | |
| if candidates.empty: | |
| return go.Figure() | |
| best_other = candidates.idxmax() | |
| fdf = mdf[mdf["_other_key"] == best_other] | |
| active_cats = { | |
| c for c in fdf["task_category"].dropna().unique() | |
| if fdf[fdf["task_category"] == c][metric].notna().any() | |
| } | |
| categories = [k for k in CATEGORY_DISPLAY_NAMES if k != "other" and k in active_cats] | |
| if not categories: | |
| return go.Figure() | |
| metric_label = CLASSIFICATION_METRICS.get(metric, metric) | |
| vary_vals = sorted( | |
| v for v in fdf[vary_by].dropna().unique().tolist() | |
| if fdf[fdf[vary_by] == v][metric].notna().any() | |
| ) | |
| if not vary_vals: | |
| return go.Figure() | |
| fixed_desc = " · ".join(f"{_ABLATION_DIMS[d]}: {v}" for d, v in zip(other_dims, best_other.split(" · "))) | |
| fig = go.Figure() | |
| for i, val in enumerate(vary_vals): | |
| vdf = fdf[fdf[vary_by] == val] | |
| y_vals, n_tasks = [], [] | |
| for cat in categories: | |
| vals = vdf[vdf["task_category"] == cat][metric].dropna() | |
| y_vals.append(round(float(vals.mean()), 3) if not vals.empty else None) | |
| n_tasks.append(vdf[vdf["task_category"] == cat]["task_name"].nunique()) | |
| color = _MODEL_PALETTE[i % len(_MODEL_PALETTE)] | |
| fig.add_trace(go.Bar( | |
| name=str(val), | |
| x=[CATEGORY_DISPLAY_NAMES[c] for c in categories], | |
| y=y_vals, | |
| customdata=n_tasks, | |
| marker=_bar_marker(color), | |
| hovertemplate=( | |
| f"<b>{_ABLATION_DIMS[vary_by]}: {val}</b> (fixed {fixed_desc})<br>" | |
| "%{x}: %{y:.3f} — mean across %{customdata} task(s)<extra></extra>" | |
| ), | |
| )) | |
| fig.update_layout( | |
| title=f"{metric_label} per task category — {model_alias}, varying {_ABLATION_DIMS[vary_by].lower()} " | |
| f"(fixed {fixed_desc})", | |
| yaxis_title=f"{metric_label} (mean across tasks in category)", | |
| barmode="group", | |
| height=440, template=_TEMPLATE, legend=_LEGEND_H, | |
| ) | |
| return fig | |
| def bar_plot_per_task(df: pd.DataFrame, metric: str, category: str) -> go.Figure: | |
| """Grouped bar: X = individual tasks within category, bar series = variants.""" | |
| if df.empty or metric not in df.columns or not category: | |
| return go.Figure() | |
| task_df = df[df["task_category"] == category] | |
| if task_df.empty: | |
| fig = go.Figure() | |
| fig.add_annotation(text=f"No data for category '{category}'", showarrow=False, | |
| xref="paper", yref="paper", x=0.5, y=0.5) | |
| return fig | |
| tasks = sorted(task_df["task_name"].unique()) | |
| metric_label = CLASSIFICATION_METRICS.get(metric, metric) | |
| cat_label = CATEGORY_DISPLAY_NAMES.get(category, category) | |
| return _grouped_bar( | |
| task_df, metric, | |
| x_items=tasks, | |
| x_labels=tasks, | |
| get_vals=lambda mdf, task: mdf[mdf["task_name"] == task][metric], | |
| title=f"{metric_label} per task — {cat_label}", | |
| legend_by_model=True, | |
| ) | |
| # --------------------------------------------------------------------------- | |
| # 6. Speed vs Performance scatter — one point per variant, Pareto highlighted | |
| # --------------------------------------------------------------------------- | |
| def scatter_speed( | |
| df: pd.DataFrame, | |
| metric: str, | |
| speed_col: str | None = None, | |
| log_scale: bool = True, | |
| ) -> go.Figure: | |
| """Scatter: mean metric vs mean speed/VRAM, one point per variant. | |
| Pareto-frontier variants (not dominated on the two axes) are highlighted | |
| with a larger marker, black outline, and short text label. Dominated | |
| variants are shown faded. Speed values are averaged across tasks because | |
| throughput and embedding time vary with sequence length. | |
| """ | |
| if df.empty or metric not in df.columns: | |
| return go.Figure() | |
| available = [c for c in _SPEED_META if c in df.columns] | |
| if not available: | |
| fig = go.Figure() | |
| fig.add_annotation( | |
| text="Speed metrics not available in data", | |
| showarrow=False, xref="paper", yref="paper", x=0.5, y=0.5, | |
| ) | |
| return fig | |
| if speed_col is None or speed_col not in df.columns: | |
| speed_col = available[0] | |
| speed_lbl, higher_x_better = _SPEED_META[speed_col] | |
| metric_label = CLASSIFICATION_METRICS.get(metric, metric) | |
| df = df.copy() | |
| df["_vkey"] = _variant_key_col(df) | |
| # Aggregate to one row per variant; speed cols averaged across tasks | |
| agg_cols: dict = {metric: "mean", "model_alias": "first"} | |
| for c in available: | |
| agg_cols[c] = "mean" | |
| vdf = df.groupby("_vkey").agg(agg_cols).reset_index() | |
| # Full hover label (variant_label) and short display label for Pareto text | |
| first_rows = {vk: grp.iloc[0] for vk, grp in df.groupby("_vkey")} | |
| vdf["_label"] = vdf["_vkey"].map(lambda vk: _variant_label(first_rows[vk])) | |
| vdf["_short"] = vdf["_vkey"].map( | |
| lambda vk: ( | |
| _norm(first_rows[vk].get("model_alias", "")) | |
| + (" · " + _norm(first_rows[vk].get("pooling_strategy", "")) | |
| if _norm(first_rows[vk].get("pooling_strategy", "")) else "") | |
| ) | |
| ) | |
| valid = vdf[[speed_col, metric]].notna().all(axis=1) | |
| vdf = vdf[valid].reset_index(drop=True) | |
| if vdf.empty: | |
| fig = go.Figure() | |
| fig.add_annotation( | |
| text=f"No valid data for {speed_lbl}", | |
| showarrow=False, xref="paper", yref="paper", x=0.5, y=0.5, | |
| ) | |
| return fig | |
| pareto = _pareto_mask( | |
| vdf[speed_col].to_numpy(dtype=float), | |
| vdf[metric].to_numpy(dtype=float), | |
| higher_x_better=higher_x_better, | |
| ) | |
| vdf["_pareto"] = pareto | |
| models = sorted(vdf["model_alias"].dropna().unique()) if "model_alias" in vdf.columns else [] | |
| model_color_map = _model_color_map(models) | |
| legend_shown: set[str] = set() | |
| fig = go.Figure() | |
| hover_tmpl = ( | |
| "<b>%{text}</b><br>" | |
| + f"{speed_lbl} (mean): %{{x}}<br>" | |
| + f"{metric_label}: %{{y:.3f}}<extra></extra>" | |
| ) | |
| # Dominated points — faded, small | |
| for model in models: | |
| mdf = vdf[(vdf["model_alias"] == model) & ~vdf["_pareto"]] | |
| if mdf.empty: | |
| continue | |
| color = model_color_map.get(model, _MODEL_PALETTE[0]) | |
| show_legend = model not in legend_shown | |
| legend_shown.add(model) | |
| fig.add_trace(go.Scatter( | |
| x=mdf[speed_col], y=mdf[metric], | |
| mode="markers", | |
| name=model, legendgroup=model, showlegend=show_legend, | |
| marker=dict(color=color, size=8, opacity=0.25, line=dict(width=0)), | |
| text=mdf["_label"], | |
| hovertemplate=hover_tmpl, | |
| )) | |
| # Pareto points — solid, labeled, black outline | |
| for model in models: | |
| mdf = vdf[(vdf["model_alias"] == model) & vdf["_pareto"]] | |
| if mdf.empty: | |
| continue | |
| color = model_color_map.get(model, _MODEL_PALETTE[0]) | |
| show_legend = model not in legend_shown | |
| legend_shown.add(model) | |
| fig.add_trace(go.Scatter( | |
| x=mdf[speed_col], y=mdf[metric], | |
| mode="markers+text", | |
| name=model, legendgroup=model, showlegend=show_legend, | |
| marker=dict(color=color, size=14, opacity=0.9, line=dict(width=2, color="black")), | |
| text=mdf["_short"], | |
| textposition="top center", | |
| textfont=dict(size=9, color="#222"), | |
| customdata=mdf[["_label"]].values, | |
| hovertemplate=( | |
| "<b>%{customdata[0]}</b><br>" | |
| + f"{speed_lbl} (mean): %{{x}}<br>" | |
| + f"{metric_label}: %{{y:.3f}}<extra></extra>" | |
| ), | |
| )) | |
| # Pareto legend symbol | |
| fig.add_trace(go.Scatter( | |
| x=[None], y=[None], mode="markers", | |
| name="Pareto frontier", | |
| marker=dict( | |
| color="rgba(0,0,0,0)", size=14, | |
| line=dict(width=2, color="black"), | |
| symbol="circle", | |
| ), | |
| showlegend=True, legendgroup="__pareto__", | |
| )) | |
| direction = "higher = faster" if higher_x_better else "lower = cheaper" | |
| fig.update_layout( | |
| title=f"{metric_label} vs {speed_lbl} ({direction})", | |
| xaxis=dict( | |
| title_text=f"{speed_lbl} — mean across tasks", | |
| type="log" if log_scale else "linear", | |
| ), | |
| yaxis_title=metric_label, | |
| height=540, | |
| template=_TEMPLATE, | |
| legend=_LEGEND_H, | |
| ) | |
| return fig | |
| def efficiency_table(df: pd.DataFrame, metric: str) -> pd.DataFrame: | |
| """Per-variant efficiency summary: MCC, speed, VRAM, and derived MCC/sec, MCC/MB. | |
| Sorted by metric descending. Speed values are means across tasks. | |
| """ | |
| if df.empty or metric not in df.columns: | |
| return pd.DataFrame() | |
| df = df.copy() | |
| df["_vkey"] = _variant_key_col(df) | |
| available_speed = [c for c in _SPEED_META if c in df.columns] | |
| agg_cols: dict = {metric: "mean", "model_alias": "first"} | |
| for c in available_speed: | |
| agg_cols[c] = "mean" | |
| vdf = df.groupby("_vkey").agg(agg_cols).reset_index() | |
| first_rows = {vk: grp.iloc[0] for vk, grp in df.groupby("_vkey")} | |
| vdf["_label"] = vdf["_vkey"].map(lambda vk: _variant_label(first_rows[vk])) | |
| vdf = vdf.sort_values(metric, ascending=False).reset_index(drop=True) | |
| metric_label = CLASSIFICATION_METRICS.get(metric, metric) | |
| out = pd.DataFrame() | |
| out["Variant"] = vdf["_label"] | |
| out["Family"] = vdf.get("model_alias", pd.Series("", index=vdf.index)) | |
| out[metric_label] = vdf[metric].round(3) | |
| thr_col = next((c for c in ["test_throughput_seq_s"] if c in vdf.columns), None) | |
| emb_col = next((c for c in ["test_embedding_time_s"] if c in vdf.columns), None) | |
| vram_col = next((c for c in ["peak_vram_extraction_mb", "vram_model_mb"] if c in vdf.columns), None) | |
| if thr_col: | |
| out["Throughput (seq/s)"] = vdf[thr_col].round(2) | |
| if emb_col: | |
| out["Emb. Time (s)"] = vdf[emb_col].round(4) | |
| out["MCC/sec"] = (vdf[metric] / vdf[emb_col].replace(0, float("nan"))).round(4) | |
| if vram_col: | |
| vram_lbl = _SPEED_META[vram_col][0] | |
| out[vram_lbl] = vdf[vram_col].round(1) | |
| out["MCC/MB"] = (vdf[metric] / vdf[vram_col].replace(0, float("nan"))).round(6) | |
| return out.reset_index(drop=True) | |
| # --------------------------------------------------------------------------- | |
| # 7. Context length overview (overall + by-category, X = model max context size) | |
| # --------------------------------------------------------------------------- | |
| # Marker symbol per task category, used in the by-category panel legend. | |
| _CATEGORY_SYMBOLS: dict[str, str] = { | |
| "histone_marks": "circle", | |
| "promoter": "square", | |
| "enhancer": "triangle-up", | |
| "splice_site": "diamond", | |
| "variant_effect": "triangle-down", | |
| "rna_expression": "star", | |
| "other": "cross", | |
| } | |
| def context_length_overview(df: pd.DataFrame, metric: str) -> go.Figure: | |
| """Two-panel figure: model max context size (X) vs performance (Y). | |
| One point per variant combination (model + pooling + layer + head) — not | |
| collapsed to a single "best" variant per model, since that's the granularity | |
| actually deployed. Left panel: mean metric across that variant's tasks. | |
| Right panel: same variant, broken out by task category (color=model, | |
| symbol=category). Bubble size in both panels encodes mean task sequence | |
| length, so that dimension stays visible even though it's no longer on the axis. | |
| """ | |
| if not TASK_SEQ_LENGTHS: | |
| fig = go.Figure() | |
| fig.add_annotation( | |
| text=( | |
| "Sequence length data not available.<br>" | |
| "Run <b>python tools/compute_seq_lengths.py</b> to populate it." | |
| ), | |
| showarrow=False, font=dict(size=14), xref="paper", yref="paper", x=0.5, y=0.5, | |
| ) | |
| fig.update_layout(height=480, template=_TEMPLATE) | |
| return fig | |
| required = ["model_alias", "task_name", "task_category", "task_seq_length", "max_context_size", metric] | |
| if df.empty or any(c not in df.columns for c in required): | |
| return go.Figure() | |
| _CONFIG_COLS = ["model_hf_id", "pooling_strategy", "layer_selection_params", "head_type"] | |
| extra_cols = [c for c in _CONFIG_COLS if c in df.columns] | |
| d = df.dropna(subset=[metric, "task_seq_length", "max_context_size"])[required + extra_cols].copy() | |
| if d.empty: | |
| return go.Figure() | |
| d["_vkey"] = _variant_key_col(d) | |
| metric_label = CLASSIFICATION_METRICS.get(metric, metric) | |
| models = sorted(d["model_alias"].unique()) | |
| model_color_map = _model_color_map(models) | |
| n_tasks = d.groupby("model_alias")["task_name"].nunique().max() | |
| seq_all = d["task_seq_length"] | |
| seq_lo, seq_hi = seq_all.min(), seq_all.max() | |
| def _size(seq_len: float) -> float: | |
| if seq_hi == seq_lo: | |
| return 22.0 | |
| return 14 + (seq_len - seq_lo) / (seq_hi - seq_lo) * 30 | |
| vkey_label = {vkey: _variant_label(g.iloc[0]) for vkey, g in d.groupby("_vkey")} | |
| fig = make_subplots( | |
| rows=1, cols=2, | |
| subplot_titles=( | |
| "Context length vs. overall performance", | |
| "Context length vs. performance, by task category", | |
| ), | |
| horizontal_spacing=0.1, | |
| ) | |
| # ---- Left panel: one point per variant combination --------------------- | |
| overall = ( | |
| d.groupby(["model_alias", "_vkey"]) | |
| .agg(y=(metric, "mean"), n=(metric, "count"), max_ctx=("max_context_size", "first"), | |
| seq_mean=("task_seq_length", "mean")) | |
| .reset_index() | |
| ) | |
| legend_shown: set[str] = set() | |
| for _, row in overall.iterrows(): | |
| model = row["model_alias"] | |
| vlabel = vkey_label.get(row["_vkey"], row["_vkey"]) | |
| color = model_color_map.get(model, _MODEL_PALETTE[0]) | |
| show_legend = model not in legend_shown | |
| legend_shown.add(model) | |
| fig.add_trace( | |
| go.Scatter( | |
| x=[row["max_ctx"]], y=[row["y"]], mode="markers", | |
| name=model, legendgroup=model, showlegend=show_legend, | |
| marker=dict(size=_size(row["seq_mean"]), color=color, opacity=0.75, line=dict(width=1, color="white")), | |
| hovertemplate=( | |
| f"<b>{model}</b><br>{vlabel}<br>Max context: %{{x}} bp<br>" | |
| f"{metric_label}: %{{y:.3f}} ({row['n']} tasks)<br>" | |
| "Mean task seq length: %{customdata:.0f} bp<extra></extra>" | |
| ), | |
| customdata=[row["seq_mean"]], | |
| ), | |
| row=1, col=1, | |
| ) | |
| # ---- Right panel: per (variant, task category) -------------------------- | |
| by_cat = ( | |
| d.groupby(["model_alias", "_vkey", "task_category"]) | |
| .agg(y=(metric, "mean"), max_ctx=("max_context_size", "first"), seq_mean=("task_seq_length", "mean")) | |
| .reset_index() | |
| ) | |
| for _, row in by_cat.iterrows(): | |
| cat = row["task_category"] | |
| vlabel = vkey_label.get(row["_vkey"], row["_vkey"]) | |
| color = model_color_map.get(row["model_alias"], _MODEL_PALETTE[0]) | |
| symbol = _CATEGORY_SYMBOLS.get(cat, "cross") | |
| cat_label = CATEGORY_DISPLAY_NAMES.get(cat, cat) | |
| fig.add_trace( | |
| go.Scatter( | |
| x=[row["max_ctx"]], y=[row["y"]], mode="markers", | |
| name=row["model_alias"], legendgroup=row["model_alias"], showlegend=False, | |
| marker=dict(size=_size(row["seq_mean"]), color=color, symbol=symbol, opacity=0.75, | |
| line=dict(width=1, color="white")), | |
| hovertemplate=( | |
| f"<b>{row['model_alias']}</b> — {cat_label}<br>{vlabel}<br>Max context: %{{x}} bp<br>" | |
| f"{metric_label}: %{{y:.3f}}<br>Mean task seq length: %{{customdata:.0f}} bp<extra></extra>" | |
| ), | |
| customdata=[row["seq_mean"]], | |
| ), | |
| row=1, col=2, | |
| ) | |
| # Legend-only dummy traces so task-category shapes have a key. | |
| for cat in sorted(by_cat["task_category"].unique()): | |
| fig.add_trace( | |
| go.Scatter( | |
| x=[None], y=[None], mode="markers", | |
| marker=dict(size=10, color="#888", symbol=_CATEGORY_SYMBOLS.get(cat, "cross")), | |
| name=CATEGORY_DISPLAY_NAMES.get(cat, cat), | |
| legendgroup="category", legendgrouptitle_text="Task category", | |
| showlegend=True, | |
| ), | |
| row=1, col=2, | |
| ) | |
| fig.update_xaxes(title_text="Max context size (bp)", type="log", row=1, col=1) | |
| fig.update_xaxes(title_text="Max context size (bp)", type="log", row=1, col=2) | |
| fig.update_yaxes(title_text=f"Mean {metric_label} per variant ({n_tasks} tasks)", row=1, col=1) | |
| fig.update_yaxes(title_text=f"Mean {metric_label}", row=1, col=2) | |
| fig.update_layout( | |
| height=560, | |
| template=_TEMPLATE, | |
| legend=dict(orientation="v", yanchor="top", y=1, xanchor="left", x=1.02), | |
| margin=dict(r=160), | |
| ) | |
| return fig | |