Spaces:
Sleeping
Sleeping
github-actions[bot] commited on
Commit ·
c833974
1
Parent(s): e8805a3
Rebuild Context Length tab as per-variant, model-context-size chart
Browse filesSwitch X-axis from task sequence length to model max context size so the
chart actually answers whether larger context windows improve performance;
split into an overall + by-task-category panel. Task sequence length is
kept via bubble size instead of being dropped. Plots every variant
combination (not aggregated per model) to match production deployment
granularity.
- app.py +3 -3
- src/plots.py +113 -50
app.py
CHANGED
|
@@ -8,7 +8,7 @@ from src.data import apply_filters, deduplicate, load_data
|
|
| 8 |
from src.plots import (
|
| 9 |
bar_plot_per_category,
|
| 10 |
bar_plot_per_task,
|
| 11 |
-
|
| 12 |
efficiency_table,
|
| 13 |
heatmap_variants,
|
| 14 |
ranking_bar,
|
|
@@ -109,7 +109,7 @@ def render_speed(
|
|
| 109 |
def render_context(df: pd.DataFrame, metric_label: str, top_n: int = 20):
|
| 110 |
key = _metric_key(metric_label)
|
| 111 |
df = top_n_variants(df, key, int(top_n))
|
| 112 |
-
return
|
| 113 |
|
| 114 |
|
| 115 |
def render_variant_heatmap(
|
|
@@ -369,7 +369,7 @@ with gr.Blocks(title="DNA Benchmark Leaderboard", theme=gr.themes.Soft()) as dem
|
|
| 369 |
|
| 370 |
# ── Tab 4: Context Length ─────────────────────────────────────────────
|
| 371 |
with gr.Tab("🔬 Context Length"):
|
| 372 |
-
context_fig = gr.Plot(label="Performance vs task sequence length")
|
| 373 |
|
| 374 |
# ── Tab 5: Scope / Filter ─────────────────────────────────────────────
|
| 375 |
with gr.Tab("🔬 Scope / Filter"):
|
|
|
|
| 8 |
from src.plots import (
|
| 9 |
bar_plot_per_category,
|
| 10 |
bar_plot_per_task,
|
| 11 |
+
context_length_overview,
|
| 12 |
efficiency_table,
|
| 13 |
heatmap_variants,
|
| 14 |
ranking_bar,
|
|
|
|
| 109 |
def render_context(df: pd.DataFrame, metric_label: str, top_n: int = 20):
|
| 110 |
key = _metric_key(metric_label)
|
| 111 |
df = top_n_variants(df, key, int(top_n))
|
| 112 |
+
return context_length_overview(df, key)
|
| 113 |
|
| 114 |
|
| 115 |
def render_variant_heatmap(
|
|
|
|
| 369 |
|
| 370 |
# ── Tab 4: Context Length ─────────────────────────────────────────────
|
| 371 |
with gr.Tab("🔬 Context Length"):
|
| 372 |
+
context_fig = gr.Plot(label="Performance vs model context size (bubble size = task sequence length)")
|
| 373 |
|
| 374 |
# ── Tab 5: Scope / Filter ─────────────────────────────────────────────
|
| 375 |
with gr.Tab("🔬 Scope / Filter"):
|
src/plots.py
CHANGED
|
@@ -5,6 +5,7 @@ import json
|
|
| 5 |
import numpy as np
|
| 6 |
import pandas as pd
|
| 7 |
import plotly.graph_objects as go
|
|
|
|
| 8 |
|
| 9 |
from src.constants import (
|
| 10 |
CATEGORY_COLORS,
|
|
@@ -772,13 +773,30 @@ def efficiency_table(df: pd.DataFrame, metric: str) -> pd.DataFrame:
|
|
| 772 |
|
| 773 |
|
| 774 |
# ---------------------------------------------------------------------------
|
| 775 |
-
# 7. Context length
|
| 776 |
# ---------------------------------------------------------------------------
|
| 777 |
|
| 778 |
-
|
| 779 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 780 |
|
| 781 |
-
One
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 782 |
"""
|
| 783 |
if not TASK_SEQ_LENGTHS:
|
| 784 |
fig = go.Figure()
|
|
@@ -792,73 +810,118 @@ def bubble_context_length(df: pd.DataFrame, metric: str) -> go.Figure:
|
|
| 792 |
fig.update_layout(height=480, template="plotly_white")
|
| 793 |
return fig
|
| 794 |
|
| 795 |
-
|
|
|
|
| 796 |
return go.Figure()
|
| 797 |
|
| 798 |
_CONFIG_COLS = ["model_hf_id", "pooling_strategy", "layer_selection_params", "head_type"]
|
| 799 |
extra_cols = [c for c in _CONFIG_COLS if c in df.columns]
|
| 800 |
-
|
| 801 |
-
|
|
|
|
|
|
|
| 802 |
|
| 803 |
-
|
| 804 |
-
|
| 805 |
-
|
|
|
|
| 806 |
|
| 807 |
-
|
| 808 |
-
|
| 809 |
-
for col, lbl in [("model_hf_id", "HF model"), ("pooling_strategy", "Pooling"),
|
| 810 |
-
("_layer_info", "Layer"), ("head_type", "Head")]:
|
| 811 |
-
if col not in row.index:
|
| 812 |
-
continue
|
| 813 |
-
v = row[col]
|
| 814 |
-
if v is not None and not (isinstance(v, float) and pd.isna(v)):
|
| 815 |
-
parts.append(f"{lbl}: {v}")
|
| 816 |
-
return "<br>".join(parts)
|
| 817 |
|
| 818 |
-
|
|
|
|
|
|
|
|
|
|
| 819 |
|
| 820 |
-
|
| 821 |
-
model_color_map = _model_color_map(models)
|
| 822 |
-
max_ctx = plot_df["max_context_size"].max() or 1
|
| 823 |
-
metric_label = CLASSIFICATION_METRICS.get(metric, metric)
|
| 824 |
-
fig = go.Figure()
|
| 825 |
|
| 826 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 827 |
|
| 828 |
-
|
| 829 |
-
|
| 830 |
-
|
| 831 |
-
|
| 832 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 833 |
color = model_color_map.get(model, _MODEL_PALETTE[0])
|
| 834 |
show_legend = model not in legend_shown
|
| 835 |
legend_shown.add(model)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 836 |
|
| 837 |
-
|
| 838 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 839 |
fig.add_trace(
|
| 840 |
go.Scatter(
|
| 841 |
-
x=
|
| 842 |
-
name=
|
| 843 |
-
marker=dict(size=
|
| 844 |
-
|
| 845 |
-
legendgroup=model, showlegend=show_legend,
|
| 846 |
hovertemplate=(
|
| 847 |
-
f"<b>
|
| 848 |
-
"
|
| 849 |
-
f"{metric_label}: %{{y:.3f}}<br>"
|
| 850 |
-
"Max context: %{customdata[0]} bp<br>"
|
| 851 |
-
"%{customdata[1]}<extra></extra>"
|
| 852 |
),
|
| 853 |
-
|
|
|
|
|
|
|
| 854 |
)
|
| 855 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 856 |
fig.update_layout(
|
| 857 |
-
|
| 858 |
-
xaxis_title="Task sequence length (bp)",
|
| 859 |
-
yaxis_title=metric_label,
|
| 860 |
-
height=520,
|
| 861 |
template="plotly_white",
|
| 862 |
-
legend=
|
|
|
|
| 863 |
)
|
| 864 |
return fig
|
|
|
|
| 5 |
import numpy as np
|
| 6 |
import pandas as pd
|
| 7 |
import plotly.graph_objects as go
|
| 8 |
+
from plotly.subplots import make_subplots
|
| 9 |
|
| 10 |
from src.constants import (
|
| 11 |
CATEGORY_COLORS,
|
|
|
|
| 773 |
|
| 774 |
|
| 775 |
# ---------------------------------------------------------------------------
|
| 776 |
+
# 7. Context length overview (overall + by-category, X = model max context size)
|
| 777 |
# ---------------------------------------------------------------------------
|
| 778 |
|
| 779 |
+
# Marker symbol per task category, used in the by-category panel legend.
|
| 780 |
+
_CATEGORY_SYMBOLS: dict[str, str] = {
|
| 781 |
+
"histone_marks": "circle",
|
| 782 |
+
"promoter": "square",
|
| 783 |
+
"enhancer": "triangle-up",
|
| 784 |
+
"splice_site": "diamond",
|
| 785 |
+
"variant_effect": "triangle-down",
|
| 786 |
+
"rna_expression": "star",
|
| 787 |
+
"other": "cross",
|
| 788 |
+
}
|
| 789 |
+
|
| 790 |
+
|
| 791 |
+
def context_length_overview(df: pd.DataFrame, metric: str) -> go.Figure:
|
| 792 |
+
"""Two-panel figure: model max context size (X) vs performance (Y).
|
| 793 |
|
| 794 |
+
One point per variant combination (model + pooling + layer + head) — not
|
| 795 |
+
collapsed to a single "best" variant per model, since that's the granularity
|
| 796 |
+
actually deployed. Left panel: mean metric across that variant's tasks.
|
| 797 |
+
Right panel: same variant, broken out by task category (color=model,
|
| 798 |
+
symbol=category). Bubble size in both panels encodes mean task sequence
|
| 799 |
+
length, so that dimension stays visible even though it's no longer on the axis.
|
| 800 |
"""
|
| 801 |
if not TASK_SEQ_LENGTHS:
|
| 802 |
fig = go.Figure()
|
|
|
|
| 810 |
fig.update_layout(height=480, template="plotly_white")
|
| 811 |
return fig
|
| 812 |
|
| 813 |
+
required = ["model_alias", "task_name", "task_category", "task_seq_length", "max_context_size", metric]
|
| 814 |
+
if df.empty or any(c not in df.columns for c in required):
|
| 815 |
return go.Figure()
|
| 816 |
|
| 817 |
_CONFIG_COLS = ["model_hf_id", "pooling_strategy", "layer_selection_params", "head_type"]
|
| 818 |
extra_cols = [c for c in _CONFIG_COLS if c in df.columns]
|
| 819 |
+
d = df.dropna(subset=[metric, "task_seq_length", "max_context_size"])[required + extra_cols].copy()
|
| 820 |
+
if d.empty:
|
| 821 |
+
return go.Figure()
|
| 822 |
+
d["_vkey"] = _variant_key_col(d)
|
| 823 |
|
| 824 |
+
metric_label = CLASSIFICATION_METRICS.get(metric, metric)
|
| 825 |
+
models = sorted(d["model_alias"].unique())
|
| 826 |
+
model_color_map = _model_color_map(models)
|
| 827 |
+
n_tasks = d.groupby("model_alias")["task_name"].nunique().max()
|
| 828 |
|
| 829 |
+
seq_all = d["task_seq_length"]
|
| 830 |
+
seq_lo, seq_hi = seq_all.min(), seq_all.max()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 831 |
|
| 832 |
+
def _size(seq_len: float) -> float:
|
| 833 |
+
if seq_hi == seq_lo:
|
| 834 |
+
return 22.0
|
| 835 |
+
return 14 + (seq_len - seq_lo) / (seq_hi - seq_lo) * 30
|
| 836 |
|
| 837 |
+
vkey_label = {vkey: _variant_label(g.iloc[0]) for vkey, g in d.groupby("_vkey")}
|
|
|
|
|
|
|
|
|
|
|
|
|
| 838 |
|
| 839 |
+
fig = make_subplots(
|
| 840 |
+
rows=1, cols=2,
|
| 841 |
+
subplot_titles=(
|
| 842 |
+
"Context length vs. overall performance",
|
| 843 |
+
"Context length vs. performance, by task category",
|
| 844 |
+
),
|
| 845 |
+
horizontal_spacing=0.1,
|
| 846 |
+
)
|
| 847 |
|
| 848 |
+
# ---- Left panel: one point per variant combination ---------------------
|
| 849 |
+
overall = (
|
| 850 |
+
d.groupby(["model_alias", "_vkey"])
|
| 851 |
+
.agg(y=(metric, "mean"), n=(metric, "count"), max_ctx=("max_context_size", "first"),
|
| 852 |
+
seq_mean=("task_seq_length", "mean"))
|
| 853 |
+
.reset_index()
|
| 854 |
+
)
|
| 855 |
+
legend_shown: set[str] = set()
|
| 856 |
+
for _, row in overall.iterrows():
|
| 857 |
+
model = row["model_alias"]
|
| 858 |
+
vlabel = vkey_label.get(row["_vkey"], row["_vkey"])
|
| 859 |
color = model_color_map.get(model, _MODEL_PALETTE[0])
|
| 860 |
show_legend = model not in legend_shown
|
| 861 |
legend_shown.add(model)
|
| 862 |
+
fig.add_trace(
|
| 863 |
+
go.Scatter(
|
| 864 |
+
x=[row["max_ctx"]], y=[row["y"]], mode="markers",
|
| 865 |
+
name=model, legendgroup=model, showlegend=show_legend,
|
| 866 |
+
marker=dict(size=_size(row["seq_mean"]), color=color, opacity=0.75, line=dict(width=1, color="white")),
|
| 867 |
+
hovertemplate=(
|
| 868 |
+
f"<b>{model}</b><br>{vlabel}<br>Max context: %{{x}} bp<br>"
|
| 869 |
+
f"{metric_label}: %{{y:.3f}} ({row['n']} tasks)<br>"
|
| 870 |
+
"Mean task seq length: %{customdata:.0f} bp<extra></extra>"
|
| 871 |
+
),
|
| 872 |
+
customdata=[row["seq_mean"]],
|
| 873 |
+
),
|
| 874 |
+
row=1, col=1,
|
| 875 |
+
)
|
| 876 |
|
| 877 |
+
# ---- Right panel: per (variant, task category) --------------------------
|
| 878 |
+
by_cat = (
|
| 879 |
+
d.groupby(["model_alias", "_vkey", "task_category"])
|
| 880 |
+
.agg(y=(metric, "mean"), max_ctx=("max_context_size", "first"), seq_mean=("task_seq_length", "mean"))
|
| 881 |
+
.reset_index()
|
| 882 |
+
)
|
| 883 |
+
for _, row in by_cat.iterrows():
|
| 884 |
+
cat = row["task_category"]
|
| 885 |
+
vlabel = vkey_label.get(row["_vkey"], row["_vkey"])
|
| 886 |
+
color = model_color_map.get(row["model_alias"], _MODEL_PALETTE[0])
|
| 887 |
+
symbol = _CATEGORY_SYMBOLS.get(cat, "cross")
|
| 888 |
+
cat_label = CATEGORY_DISPLAY_NAMES.get(cat, cat)
|
| 889 |
fig.add_trace(
|
| 890 |
go.Scatter(
|
| 891 |
+
x=[row["max_ctx"]], y=[row["y"]], mode="markers",
|
| 892 |
+
name=row["model_alias"], legendgroup=row["model_alias"], showlegend=False,
|
| 893 |
+
marker=dict(size=_size(row["seq_mean"]), color=color, symbol=symbol, opacity=0.75,
|
| 894 |
+
line=dict(width=1, color="white")),
|
|
|
|
| 895 |
hovertemplate=(
|
| 896 |
+
f"<b>{row['model_alias']}</b> — {cat_label}<br>{vlabel}<br>Max context: %{{x}} bp<br>"
|
| 897 |
+
f"{metric_label}: %{{y:.3f}}<br>Mean task seq length: %{{customdata:.0f}} bp<extra></extra>"
|
|
|
|
|
|
|
|
|
|
| 898 |
),
|
| 899 |
+
customdata=[row["seq_mean"]],
|
| 900 |
+
),
|
| 901 |
+
row=1, col=2,
|
| 902 |
)
|
| 903 |
|
| 904 |
+
# Legend-only dummy traces so task-category shapes have a key.
|
| 905 |
+
for cat in sorted(by_cat["task_category"].unique()):
|
| 906 |
+
fig.add_trace(
|
| 907 |
+
go.Scatter(
|
| 908 |
+
x=[None], y=[None], mode="markers",
|
| 909 |
+
marker=dict(size=10, color="#888", symbol=_CATEGORY_SYMBOLS.get(cat, "cross")),
|
| 910 |
+
name=CATEGORY_DISPLAY_NAMES.get(cat, cat),
|
| 911 |
+
legendgroup="category", legendgrouptitle_text="Task category",
|
| 912 |
+
showlegend=True,
|
| 913 |
+
),
|
| 914 |
+
row=1, col=2,
|
| 915 |
+
)
|
| 916 |
+
|
| 917 |
+
fig.update_xaxes(title_text="Max context size (bp)", type="log", row=1, col=1)
|
| 918 |
+
fig.update_xaxes(title_text="Max context size (bp)", type="log", row=1, col=2)
|
| 919 |
+
fig.update_yaxes(title_text=f"Mean {metric_label} per variant ({n_tasks} tasks)", row=1, col=1)
|
| 920 |
+
fig.update_yaxes(title_text=f"Mean {metric_label}", row=1, col=2)
|
| 921 |
fig.update_layout(
|
| 922 |
+
height=560,
|
|
|
|
|
|
|
|
|
|
| 923 |
template="plotly_white",
|
| 924 |
+
legend=dict(orientation="v", yanchor="top", y=1, xanchor="left", x=1.02),
|
| 925 |
+
margin=dict(r=160),
|
| 926 |
)
|
| 927 |
return fig
|