"""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 "
".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 = "
".join(
f"{k}: {v:.3f}" for k, v in rec["breakdown"].items() if v is not None
)
hover = f"{rec['label']}
Overall: {rec['overall']:.3f}"
if breakdown_lines:
hover += "
" + 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 + "",
))
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"%{{text}}
{metric_label}: %{{x:.3f}}",
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"{col_name}
{metric_label}: {v_str}
{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"Overall
{metric_label}: {ov_str}
{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}",
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"{label}
%{{x}}: %{{y:.3f}}",
))
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"{label}
%{{x}}: %{{y:.3f}}",
))
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"{_ABLATION_DIMS[vary_by]}: {val} (fixed {fixed_desc})
"
"%{x}: %{y:.3f} — mean across %{customdata} task(s)"
),
))
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 = (
"%{text}
"
+ f"{speed_lbl} (mean): %{{x}}
"
+ f"{metric_label}: %{{y:.3f}}"
)
# 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=(
"%{customdata[0]}
"
+ f"{speed_lbl} (mean): %{{x}}
"
+ f"{metric_label}: %{{y:.3f}}"
),
))
# 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.
"
"Run python tools/compute_seq_lengths.py 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"{model}
{vlabel}
Max context: %{{x}} bp
"
f"{metric_label}: %{{y:.3f}} ({row['n']} tasks)
"
"Mean task seq length: %{customdata:.0f} bp"
),
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"{row['model_alias']} — {cat_label}
{vlabel}
Max context: %{{x}} bp
"
f"{metric_label}: %{{y:.3f}}
Mean task seq length: %{{customdata:.0f}} bp"
),
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