github-actions[bot] commited on
Commit
c833974
·
1 Parent(s): e8805a3

Rebuild Context Length tab as per-variant, model-context-size chart

Browse files

Switch 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.

Files changed (2) hide show
  1. app.py +3 -3
  2. 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
- bubble_context_length,
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 bubble_context_length(df, key)
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 bubble chart
776
  # ---------------------------------------------------------------------------
777
 
778
- def bubble_context_length(df: pd.DataFrame, metric: str) -> go.Figure:
779
- """Bubble chart: X=task_seq_length, Y=metric, color=model family, size=max_context_size.
 
 
 
 
 
 
 
 
 
 
 
 
780
 
781
- One bubble per (variant, task). Legend grouped by model family.
 
 
 
 
 
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
- if df.empty or metric not in df.columns or "task_seq_length" not in df.columns:
 
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
- base_cols = ["model_alias", "task_name", "task_seq_length", metric, "max_context_size"]
801
- plot_df = df[base_cols + extra_cols].dropna(subset=["task_seq_length", metric]).copy()
 
 
802
 
803
- plot_df["_vkey"] = _variant_key_col(plot_df)
804
- if "layer_selection_params" in plot_df.columns:
805
- plot_df["_layer_info"] = plot_df["layer_selection_params"].apply(_fmt_layer)
 
806
 
807
- def _config_str(row: pd.Series) -> str:
808
- parts = []
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
- plot_df["_config_str"] = plot_df.apply(_config_str, axis=1)
 
 
 
819
 
820
- models = sorted(plot_df["model_alias"].unique()) if "model_alias" in plot_df.columns else []
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
- legend_shown: set[str] = set()
 
 
 
 
 
 
 
827
 
828
- for vkey in plot_df["_vkey"].unique():
829
- vdf = plot_df[plot_df["_vkey"] == vkey]
830
- first = vdf.iloc[0]
831
- model = _norm(first.get("model_alias", ""))
832
- label = _variant_label(first)
 
 
 
 
 
 
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
- sizes = (vdf["max_context_size"].fillna(max_ctx) / max_ctx * 40 + 8).clip(8, 50)
838
- cd = vdf[["max_context_size", "_config_str"]].to_numpy(dtype=object)
 
 
 
 
 
 
 
 
 
 
839
  fig.add_trace(
840
  go.Scatter(
841
- x=vdf["task_seq_length"], y=vdf[metric], mode="markers",
842
- name=model, text=vdf["task_name"],
843
- marker=dict(size=sizes, color=color, opacity=0.7, line=dict(width=1, color="white")),
844
- customdata=cd,
845
- legendgroup=model, showlegend=show_legend,
846
  hovertemplate=(
847
- f"<b>%{{text}}</b> — {label}<br>"
848
- "Seq length: %{x} bp<br>"
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
- title=f"{metric_label} vs task sequence length (bubble size = model max context)",
858
- xaxis_title="Task sequence length (bp)",
859
- yaxis_title=metric_label,
860
- height=520,
861
  template="plotly_white",
862
- legend=_LEGEND_H,
 
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