pefidias commited on
Commit
717e0e5
·
1 Parent(s): 9f99eeb

Major refactor - V2 (#7)

Browse files

- refactored the whole webapp (a486179cf332b50a0585aff1550b344815676d91)
- added raw input inspector (a011d958b616b93695e594aefaa0f9ae6bc7e336)

Makefile CHANGED
@@ -1,8 +1,11 @@
1
- .PHONY: style run
2
 
3
  style:
4
  python -m black --line-length 119 .
5
  ruff check --fix .
6
 
7
- run:
8
  gradio app.py
 
 
 
 
1
+ .PHONY: style app seq-lengths
2
 
3
  style:
4
  python -m black --line-length 119 .
5
  ruff check --fix .
6
 
7
+ app:
8
  gradio app.py
9
+
10
+ seq-lengths:
11
+ python tools/compute_seq_lengths.py
README.md CHANGED
@@ -13,31 +13,8 @@ sdk_version: 5.19.0
13
 
14
  # DNA Benchmark Leaderboard
15
 
16
- A simple leaderboard for evaluating DNA foundational models on genomics tasks.
17
-
18
- ## Configuration
19
-
20
- - `src/envs.py` - Configure repository paths and environment variables
21
- - `src/about.py` - Define tasks and leaderboard text
22
- - `src/display/utils.py` - Configure table columns and metrics
23
-
24
- ## Results Format
25
-
26
- Results files should be stored as JSON with the following structure:
27
- ```json
28
- {
29
- "config": {
30
- "model_dtype": "torch.float16",
31
- "model_name": "org/model",
32
- "model_sha": "revision"
33
- },
34
- "results": {
35
- "task_name": {
36
- "metric_name": score
37
- }
38
- }
39
- }
40
- ```
41
 
42
  ## Development
43
 
@@ -46,8 +23,29 @@ Results files should be stored as JSON with the following structure:
46
  pip install -r requirements.txt
47
 
48
  # Run locally
49
- gradio app.py
50
 
51
  # Format code
52
  make style
53
  ```
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
13
 
14
  # DNA Benchmark Leaderboard
15
 
16
+ Interactive leaderboard comparing DNA foundation models across genomics classification tasks.
17
+ Data is loaded from [`lokahq/genomic-benchmark-metrics`](https://huggingface.co/datasets/lokahq/genomic-benchmark-metrics) on startup.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
18
 
19
  ## Development
20
 
 
23
  pip install -r requirements.txt
24
 
25
  # Run locally
26
+ make app
27
 
28
  # Format code
29
  make style
30
  ```
31
+
32
+ ## Context length bubble chart
33
+
34
+ The bubble chart requires mean sequence lengths per task to be pre-computed once:
35
+
36
+ ```bash
37
+ make seq-lengths # computes lengths and writes tools/task_seq_lengths.json
38
+ git add tools/task_seq_lengths.json && git commit -m "populate task seq lengths"
39
+ ```
40
+
41
+ The chart is a no-op until `tools/task_seq_lengths.json` is populated and committed.
42
+
43
+ ## Configuration
44
+
45
+ | File | Purpose |
46
+ |---|---|
47
+ | `src/constants.py` | Task categories, metric names, seq length lookup |
48
+ | `src/data.py` | HF Hub loading, deduplication, filtering |
49
+ | `src/plots.py` | All Plotly figure factories |
50
+ | `tools/task_seq_lengths.json` | Mean sequence lengths per task (committed) |
51
+ | `.env` | `HF_TOKEN` and `LEADERBOARD_DATASET` |
app.py CHANGED
@@ -1,163 +1,414 @@
 
 
1
  import gradio as gr
2
- from apscheduler.schedulers.background import BackgroundScheduler
3
- from gradio_leaderboard import ColumnFilter, Leaderboard, SelectColumns
4
- from huggingface_hub import snapshot_download
5
 
6
- from src.about import INTRODUCTION_TEXT, LLM_BENCHMARKS_TEXT, TITLE
7
- from src.display.css_html_js import custom_css
8
- from src.display.utils import BenchRawColumn, fields
9
- from src.envs import API, EVAL_RESULTS_PATH, REPO_ID, RESULTS_REPO, TOKEN
10
- from src.populate import get_leaderboard_df_from_hf_dataset, summarize_model_task_type_performance
11
- from src.display.plotting import METRICS_FOR_PLOTS, extract_mean_std, prepare_leaderboard_df, make_plot_wrapper, make_model_all_datasets_wrapper, plot_metric_bar
 
 
 
 
 
12
 
13
- def restart_space():
14
- API.restart_space(repo_id=REPO_ID)
 
15
 
 
16
  try:
17
- print(EVAL_RESULTS_PATH)
18
- snapshot_download(
19
- repo_id=RESULTS_REPO,
20
- local_dir=EVAL_RESULTS_PATH,
21
- repo_type="dataset",
22
- tqdm_class=None,
23
- etag_timeout=30,
24
- token=TOKEN,
25
- )
26
 
27
- LEADERBOARD_DF = get_leaderboard_df_from_hf_dataset(EVAL_RESULTS_PATH)
28
- LEADERBOARD_DF = summarize_model_task_type_performance(LEADERBOARD_DF)
29
- LEADERBOARD_DF = prepare_leaderboard_df(LEADERBOARD_DF)
30
-
31
- wrapped_task_plot = make_plot_wrapper(leaderboard_df=LEADERBOARD_DF, group_by="Model", filter_col="Task", orientation="h")
32
- wrapped_model_all_datasets_plot = make_model_all_datasets_wrapper(leaderboard_df=LEADERBOARD_DF)
33
-
34
- except Exception as e:
35
- print(e)
36
- restart_space()
37
-
38
- def init_leaderboard(dataframe):
39
- if dataframe is None or dataframe.empty:
40
- raise ValueError("Leaderboard DataFrame is empty or None.")
41
-
42
- return Leaderboard(
43
- value=dataframe,
44
- datatype=[c.type for c in fields(BenchRawColumn)],
45
- select_columns=SelectColumns(
46
- default_selection=[c.name for c in fields(BenchRawColumn) if c.displayed_by_default],
47
- cant_deselect=[c.name for c in fields(BenchRawColumn) if c.never_hidden],
48
- label="Select Columns to Display:",
49
- ),
50
- filter_columns=[
51
- ColumnFilter(BenchRawColumn.task.name, type="checkboxgroup", label=BenchRawColumn.task.name),
52
- # ColumnFilter(BenchRawColumn.precision.name, type="checkboxgroup", label="Precision"),
53
- ColumnFilter(
54
- BenchRawColumn.model_params.name,
55
- type="slider",
56
- min=100,
57
- max=10000,
58
- label="Select the number of parameters (M)",
59
- ),
60
- ColumnFilter(
61
- BenchRawColumn.embds_dim.name,
62
- type="slider",
63
- min=10,
64
- max=10000,
65
- label="Select the Embeddings Size",
66
- ),
67
- # ColumnFilter(BenchRawColumn.still_on_hub.name, type="boolean", label="Deleted/incomplete", default=True),
68
- ColumnFilter(
69
- BenchRawColumn.max_context_len.name,
70
- type="slider",
71
- min=10,
72
- max=10000000,
73
- label="Select the Max Context Length (bp)",
74
- ),
75
- ],
76
- search_columns=[BenchRawColumn.model.name],
77
- )
78
 
79
- # Create an elegant dark theme with custom colors for checkboxes and sliders
80
- theme = gr.themes.Soft(
81
- primary_hue="violet",
82
- secondary_hue="purple",
83
- neutral_hue="slate",
84
- spacing_size="md",
85
- radius_size="lg",
86
- text_size="md",
87
- font=[gr.themes.GoogleFont("Inter"), "ui-sans-serif", "system-ui", "sans-serif"],
88
- ).set(
89
- checkbox_background_color_selected="*primary_600",
90
- checkbox_border_color_selected="*primary_600",
91
- checkbox_background_color_selected_dark="*primary_500",
92
- checkbox_border_color_selected_dark="*primary_500",
93
- slider_color="*primary_600",
94
- slider_color_dark="*primary_500",
95
- button_primary_background_fill="*primary_600",
96
- button_primary_background_fill_hover="*primary_700",
97
- button_primary_background_fill_dark="*primary_500",
98
- button_primary_background_fill_hover_dark="*primary_600",
99
  )
100
 
101
- demo = gr.Blocks(
102
- css=custom_css,
103
- theme=theme,
104
- js="""
105
- () => {
106
- const theme = localStorage.getItem('theme');
107
- if (theme === null) {
108
- localStorage.setItem('theme', 'dark');
109
- document.body.classList.add('dark');
110
- }
111
- }
112
- """,
113
- title="DNA Benchmark",
114
- head="<link rel='icon' href='data:image/svg+xml,<svg xmlns=\"http://www.w3.org/2000/svg\" viewBox=\"0 0 100 100\"><text y=\".9em\" font-size=\"90\">🧬</text></svg>'>"
115
  )
116
- with demo:
117
- gr.HTML(TITLE)
118
- gr.Markdown(INTRODUCTION_TEXT, elem_classes="markdown-text")
119
 
120
- with gr.Tabs(elem_classes="tab-buttons") as tabs:
121
- with gr.TabItem("🏅 LLM Leaderboard", elem_id="llm-benchmark-tab-table", id=0):
122
- leaderboard = init_leaderboard(LEADERBOARD_DF)
123
 
124
- with gr.TabItem("📊 Performance Plots", elem_id="llm-benchmark-performance-plots", id=1):
125
- with gr.Accordion("🧬 Overview of Performance for a Specific Task", open=False):
126
- gr.Markdown("Visualize model performance for a specific task and metric")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
127
 
128
- task_choices = sorted(LEADERBOARD_DF["Task"].dropna().unique())
129
- dataset_choices = sorted(LEADERBOARD_DF["Dataset Name"].dropna().unique())
130
 
131
- with gr.Row():
132
- task_dropdown = gr.Dropdown(choices=task_choices, label="Select Task")
133
- metric_dropdown = gr.Dropdown(choices=METRICS_FOR_PLOTS, value="Accuracy", label="Select Metric")
134
- dataset_dropdown = gr.Dropdown(choices=dataset_choices, value="InstaDeepAI/nucleotide_transformer_downstream_tasks_revised", label="Select Dataset")
 
 
135
 
136
- performance_plot = gr.Plot()
 
 
 
 
 
 
 
 
 
 
137
 
138
- task_dropdown.change(fn=wrapped_task_plot, inputs=[task_dropdown, metric_dropdown, dataset_dropdown], outputs=performance_plot)
139
- metric_dropdown.change(fn=wrapped_task_plot, inputs=[task_dropdown, metric_dropdown, dataset_dropdown], outputs=performance_plot)
140
- dataset_dropdown.change(fn=wrapped_task_plot, inputs=[task_dropdown, metric_dropdown, dataset_dropdown], outputs=performance_plot)
141
 
142
- with gr.Accordion("🤖 Overview of Performance for a Specific Model", open=False):
143
- gr.Markdown("Visualize model performance across all tasks from all datasets")
 
 
 
 
 
 
 
 
 
 
 
144
 
145
- model_choices = sorted(LEADERBOARD_DF["Model"].dropna().unique())
 
 
146
 
147
- with gr.Row():
148
- model_dropdown = gr.Dropdown(choices=model_choices, label="Select Model")
149
- metric2_dropdown = gr.Dropdown(choices=METRICS_FOR_PLOTS, value="Accuracy", label="Select Metric")
150
 
151
- model_plot = gr.Plot()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
152
 
153
- model_dropdown.change(fn=wrapped_model_all_datasets_plot, inputs=[model_dropdown, metric2_dropdown], outputs=model_plot)
154
- metric2_dropdown.change(fn=wrapped_model_all_datasets_plot, inputs=[model_dropdown, metric2_dropdown], outputs=model_plot)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
155
 
156
- with gr.TabItem("📝 About", elem_id="llm-benchmark-about", id=2):
157
- gr.Markdown(LLM_BENCHMARKS_TEXT, elem_classes="markdown-text")
 
 
 
 
 
158
 
159
- scheduler = BackgroundScheduler()
160
- scheduler.add_job(restart_space, "interval", seconds=1800)
161
- scheduler.start()
162
 
163
- demo.queue(default_concurrency_limit=40).launch(share=True)
 
 
1
+ """DNA Foundation Model Benchmark Leaderboard — Gradio Blocks app."""
2
+
3
  import gradio as gr
4
+ import pandas as pd
 
 
5
 
6
+ from src.constants import CLASSIFICATION_METRICS
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
+ heatmap_variants,
13
+ ranking_table,
14
+ scatter_speed,
15
+ violin_plot,
16
+ )
17
 
18
+ # ---------------------------------------------------------------------------
19
+ # Startup: load & deduplicate once
20
+ # ---------------------------------------------------------------------------
21
 
22
+ print("Loading benchmark data from HF Hub …")
23
  try:
24
+ _LOADED_DF = load_data()
25
+ _RAW_DF = deduplicate(_LOADED_DF)
26
+ print(f" Loaded {len(_RAW_DF)} rows, {_RAW_DF['model_alias'].nunique()} models.")
27
+ except Exception as exc:
28
+ print(f" WARNING: Could not load data — {exc}")
29
+ _LOADED_DF = pd.DataFrame()
30
+ _RAW_DF = pd.DataFrame()
 
 
31
 
32
+ _METRIC_CHOICES = list(CLASSIFICATION_METRICS.values())
33
+ _METRIC_KEYS = list(CLASSIFICATION_METRICS.keys())
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
34
 
35
+ _CATEGORY_CHOICES = (
36
+ sorted(_RAW_DF["task_category"].unique().tolist()) if not _RAW_DF.empty else []
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
37
  )
38
 
39
+ _DATASET_CHOICES = (
40
+ sorted(_RAW_DF["dataset_name"].dropna().unique().tolist())
41
+ if not _RAW_DF.empty and "dataset_name" in _RAW_DF.columns
42
+ else []
 
 
 
 
 
 
 
 
 
 
43
  )
 
 
 
44
 
 
 
 
45
 
46
+ def _metric_key(label: str) -> str:
47
+ if not label or label not in _METRIC_CHOICES:
48
+ return _METRIC_KEYS[0]
49
+ return _METRIC_KEYS[_METRIC_CHOICES.index(label)]
50
+
51
+
52
+ def _numeric_range(df: pd.DataFrame, col: str) -> tuple[float, float]:
53
+ if col not in df.columns or df.empty:
54
+ return 0.0, 1.0
55
+ lo, hi = float(df[col].min()), float(df[col].max())
56
+ if lo == hi:
57
+ return lo, lo + 1.0
58
+ return lo, hi
59
+
60
+
61
+ # ---------------------------------------------------------------------------
62
+ # Render helpers
63
+ # ---------------------------------------------------------------------------
64
+
65
+ def render_leaderboard(df: pd.DataFrame, metric_label: str):
66
+ key = _metric_key(metric_label)
67
+ return ranking_table(df, key), violin_plot(df, key)
68
+
69
+
70
+ def render_per_category(df: pd.DataFrame, metric_label: str, datasets: list[str] | None = None):
71
+ if datasets and "dataset_name" in df.columns:
72
+ df = df[df["dataset_name"].isin(datasets)]
73
+ key = _metric_key(metric_label)
74
+ return bar_plot_per_category(df, key)
75
+
76
+
77
+ def render_per_task(df: pd.DataFrame, metric_label: str, category: str, datasets: list[str] | None = None):
78
+ if datasets and "dataset_name" in df.columns:
79
+ df = df[df["dataset_name"].isin(datasets)]
80
+ key = _metric_key(metric_label)
81
+ return bar_plot_per_task(df, key, category)
82
+
83
+
84
+ def render_speed(df: pd.DataFrame, metric_label: str):
85
+ key = _metric_key(metric_label)
86
+ return scatter_speed(df, key)
87
+
88
+
89
+ def render_context(df: pd.DataFrame, metric_label: str):
90
+ key = _metric_key(metric_label)
91
+ return bubble_context_length(df, key)
92
+
93
+
94
+ def render_variant_heatmap(
95
+ df: pd.DataFrame,
96
+ metric_label: str,
97
+ model_alias: str,
98
+ aggregate_by_dataset: bool = False,
99
+ ):
100
+ key = _metric_key(metric_label)
101
+ return heatmap_variants(df, key, model_alias, aggregate_by_dataset)
102
+
103
+
104
+ def update_category_choices(df: pd.DataFrame, datasets: list[str] | None = None):
105
+ if datasets and "dataset_name" in df.columns:
106
+ df = df[df["dataset_name"].isin(datasets)]
107
+ cats = sorted(df["task_category"].unique().tolist()) if not df.empty else []
108
+ return gr.update(choices=cats, value=cats[0] if cats else None)
109
+
110
+
111
+ def apply_and_preview(
112
+ df: pd.DataFrame,
113
+ prod_only: bool,
114
+ exclude_models: list[str],
115
+ params_min: float,
116
+ params_max: float,
117
+ dim_min: float,
118
+ dim_max: float,
119
+ ctx_min: float,
120
+ ctx_max: float,
121
+ ) -> tuple[pd.DataFrame, pd.DataFrame, str]:
122
+ filters = {
123
+ "subsample_train": "full" if prod_only else None,
124
+ "exclude_models": exclude_models or [],
125
+ "model_params_min": params_min,
126
+ "model_params_max": params_max,
127
+ "embedding_dim_min": dim_min,
128
+ "embedding_dim_max": dim_max,
129
+ "max_context_size_min": ctx_min,
130
+ "max_context_size_max": ctx_max,
131
+ }
132
+ result = apply_filters(df, filters)
133
+ preview = result.head(50) if not result.empty else pd.DataFrame()
134
+ return result, preview, f"{len(result)} rows active"
135
+
136
+
137
+ def _hash_label(row: pd.Series) -> str:
138
+ """Human-readable label for an embedding config hash dropdown entry."""
139
+ def _s(val) -> str:
140
+ return "" if val is None or (isinstance(val, float) and pd.isna(val)) else str(val)
141
+
142
+ category = _s(row.get("task_category"))
143
+ task = _s(row.get("task_name"))
144
+ pool = _s(row.get("pooling_strategy"))
145
+ strategy = _s(row.get("layer_selection_strategy"))
146
+ idx = row.get("layer_index")
147
+ layer = strategy + (f"({int(idx)})" if idx is not None and not (isinstance(idx, float) and pd.isna(idx)) else "")
148
+ h = _s(row.get("embedding_config_hash"))[:8]
149
+
150
+ parts = [f"{category}/{task}" if category else task, pool, layer]
151
+ return " · ".join(p for p in parts if p) + f" ({h})"
152
+
153
+
154
+ def get_hash_choices_for_model(loaded_df: pd.DataFrame, model: str) -> list[tuple[str, str]]:
155
+ if not model or loaded_df.empty or "embedding_config_hash" not in loaded_df.columns:
156
+ return []
157
+ sub = loaded_df[loaded_df["model_alias"] == model].drop_duplicates("embedding_config_hash")
158
+ return sorted(
159
+ [(_hash_label(row), row["embedding_config_hash"]) for _, row in sub.iterrows()],
160
+ key=lambda x: x[0],
161
+ )
162
+
163
+
164
+ def get_head_types_for_hash(loaded_df: pd.DataFrame, model: str, key: str) -> list[str]:
165
+ if not model or not key or loaded_df.empty or "embedding_config_hash" not in loaded_df.columns:
166
+ return []
167
+ sub = loaded_df[(loaded_df["model_alias"] == model) & (loaded_df["embedding_config_hash"] == key)]
168
+ return sorted(sub["head_type"].dropna().unique().tolist()) if "head_type" in sub.columns else []
169
+
170
+
171
+ _DEFAULT_INSP_COLS = ["kept", "run_id", "run_at", "subsample_train", "task_name", "mcc_test", "accuracy_test"]
172
+
173
+
174
+ def build_inspector_table(loaded_df: pd.DataFrame, raw_df: pd.DataFrame,
175
+ model: str, key: str, head_type: str,
176
+ selected_cols: list[str] | None = None) -> pd.DataFrame:
177
+ if not model or not key or not head_type or loaded_df.empty or "embedding_config_hash" not in loaded_df.columns:
178
+ return pd.DataFrame()
179
+ mask = (
180
+ (loaded_df["model_alias"] == model)
181
+ & (loaded_df["embedding_config_hash"] == key)
182
+ & (loaded_df["head_type"] == head_type)
183
+ )
184
+ sub = loaded_df[mask].copy()
185
+ kept_ids = set(raw_df["run_id"]) if "run_id" in raw_df.columns else set()
186
+ sub["kept"] = sub["run_id"].isin(kept_ids).map({True: "✓", False: "✗"})
187
+ cols = selected_cols if selected_cols else _DEFAULT_INSP_COLS
188
+ return sub[[c for c in cols if c in sub.columns or c == "kept"]].reset_index(drop=True)
189
+
190
+
191
+ def exclude_run_ids(df: pd.DataFrame, run_ids_text: str) -> tuple[pd.DataFrame, pd.DataFrame, str]:
192
+ ids = [r.strip() for r in run_ids_text.split(",") if r.strip()]
193
+ result = apply_filters(df, {"exclude_run_ids": ids})
194
+ preview = result.head(50) if not result.empty else pd.DataFrame()
195
+ return result, preview, f"{len(result)} rows active"
196
+
197
+
198
+ # ---------------------------------------------------------------------------
199
+ # Build UI
200
+ # ---------------------------------------------------------------------------
201
+
202
+ all_models = sorted(_RAW_DF["model_alias"].unique().tolist()) if not _RAW_DF.empty else []
203
+ params_min_v, params_max_v = _numeric_range(_RAW_DF, "model_params")
204
+ dim_min_v, dim_max_v = _numeric_range(_RAW_DF, "embedding_dim")
205
+ ctx_min_v, ctx_max_v = _numeric_range(_RAW_DF, "max_context_size")
206
+
207
+ _insp_default_model = all_models[0] if all_models else None
208
+ _insp_default_keys = get_hash_choices_for_model(_LOADED_DF, _insp_default_model)
209
+ _ALL_INSP_COLS = ["kept"] + _LOADED_DF.columns.tolist() if not _LOADED_DF.empty else _DEFAULT_INSP_COLS
210
+
211
+ with gr.Blocks(title="DNA Benchmark Leaderboard") as demo:
212
+
213
+ gr.Markdown(
214
+ "# 🧬 DNA Foundation Model Benchmark Leaderboard\n"
215
+ "Compare DNA language models across genomics classification tasks."
216
+ )
217
+
218
+ # ── Global controls ───────────────────────────────────────────────────────
219
+ with gr.Row():
220
+ metric_dd = gr.Dropdown(
221
+ choices=_METRIC_CHOICES, value=_METRIC_CHOICES[0], label="Metric", scale=2
222
+ )
223
+
224
+ # ── Shared state ──────────────────────────────────────────────────────────
225
+ loaded_df_state = gr.State(_LOADED_DF)
226
+ raw_df_state = gr.State(_RAW_DF)
227
+ filtered_df_state = gr.State(_RAW_DF)
228
 
229
+ # ── Tabs ──────────────────────────────────────────────────────────────────
230
+ with gr.Tabs():
231
 
232
+ # ── Tab 1: Leaderboard ────────────────────────────────────────────────
233
+ with gr.Tab("🏅 Leaderboard"):
234
+ leaderboard_table = gr.Dataframe(
235
+ label="Rankings (mean metric per category)", interactive=False, wrap=True
236
+ )
237
+ violin_fig = gr.Plot(label="Metric distribution per model")
238
 
239
+ gr.Markdown("---\n### Drill into model variants")
240
+ model_dd = gr.Dropdown(
241
+ choices=all_models,
242
+ value=all_models[0] if all_models else None,
243
+ label="Select model",
244
+ )
245
+ agg_dataset_cb = gr.Checkbox(
246
+ label="Aggregate by dataset (show mean per benchmark dataset instead of per task)",
247
+ value=False,
248
+ )
249
+ variant_violin_fig = gr.Plot(label="Variant comparison (metric per task)")
250
 
251
+ # ── Tab 2: Per-Category Performance ───────────────────────────────────
252
+ with gr.Tab("📊 Per-Category Performance"):
253
+ category_fig = gr.Plot(label="Mean performance per task category")
254
 
255
+ gr.Markdown("---\n### Drill into individual tasks")
256
+ category_dd = gr.Dropdown(
257
+ choices=_CATEGORY_CHOICES,
258
+ value=_CATEGORY_CHOICES[0] if _CATEGORY_CHOICES else None,
259
+ label="Select category",
260
+ )
261
+ dataset_dd = gr.Dropdown(
262
+ choices=_DATASET_CHOICES,
263
+ value=[],
264
+ multiselect=True,
265
+ label="Filter by dataset",
266
+ )
267
+ per_task_fig = gr.Plot(label="Performance per task within selected category")
268
 
269
+ # ── Tab 3: Speed vs Performance ───────────────────────────────────────
270
+ with gr.Tab("⚡ Speed vs Performance"):
271
+ speed_fig = gr.Plot(label="Throughput & embedding time vs performance")
272
 
273
+ # ── Tab 4: Context Length ─────────────────────────────────────────────
274
+ with gr.Tab("🔬 Context Length"):
275
+ context_fig = gr.Plot(label="Performance vs task sequence length")
276
 
277
+ # ── Tab 5: Settings ───────────────────────────────────────────────────
278
+ with gr.Tab("⚙️ Settings"):
279
+ prod_only_cb = gr.Checkbox(
280
+ label="Production runs only (subsample_train IS NULL)",
281
+ value=False,
282
+ info="Currently returns no data — all runs use subsample_train=0.05.",
283
+ )
284
+ exclude_models_ms = gr.Dropdown(
285
+ choices=all_models, multiselect=True, label="Exclude models", value=[]
286
+ )
287
+ with gr.Row():
288
+ params_min_sl = gr.Slider(
289
+ minimum=params_min_v, maximum=params_max_v,
290
+ value=params_min_v, label="Model params (min)"
291
+ )
292
+ params_max_sl = gr.Slider(
293
+ minimum=params_min_v, maximum=params_max_v,
294
+ value=params_max_v, label="Model params (max)"
295
+ )
296
+ with gr.Row():
297
+ dim_min_sl = gr.Slider(
298
+ minimum=dim_min_v, maximum=dim_max_v,
299
+ value=dim_min_v, label="Embedding dim (min)"
300
+ )
301
+ dim_max_sl = gr.Slider(
302
+ minimum=dim_min_v, maximum=dim_max_v,
303
+ value=dim_max_v, label="Embedding dim (max)"
304
+ )
305
+ with gr.Row():
306
+ ctx_min_sl = gr.Slider(
307
+ minimum=ctx_min_v, maximum=ctx_max_v,
308
+ value=ctx_min_v, label="Max context size (min bp)"
309
+ )
310
+ ctx_max_sl = gr.Slider(
311
+ minimum=ctx_min_v, maximum=ctx_max_v,
312
+ value=ctx_max_v, label="Max context size (max bp)"
313
+ )
314
+ apply_btn = gr.Button("Apply Filters", variant="primary")
315
 
316
+ gr.Markdown("---")
317
+ run_ids_tb = gr.Textbox(
318
+ label="Exclude run_ids (comma-separated)", placeholder="run_id_1, run_id_2, …"
319
+ )
320
+ exclude_runs_btn = gr.Button("Exclude run_ids")
321
+
322
+ gr.Markdown("---")
323
+ settings_count = gr.Label(label="Active rows")
324
+ settings_preview = gr.Dataframe(label="Preview (first 50 rows)", interactive=False)
325
+
326
+ gr.Markdown("---\n### Raw Run Inspector\nSelect a model and embedding config to audit which runs were kept or dropped by dedup.")
327
+ with gr.Row():
328
+ insp_model_dd = gr.Dropdown(choices=all_models, value=_insp_default_model, label="Model", scale=2)
329
+ insp_cache_dd = gr.Dropdown(choices=_insp_default_keys, value=None, label="Embedding config hash", scale=3)
330
+ insp_head_dd = gr.Dropdown(choices=[], value=None, label="Head type", scale=2)
331
+ insp_cols_dd = gr.Dropdown(
332
+ choices=_ALL_INSP_COLS,
333
+ value=_DEFAULT_INSP_COLS,
334
+ multiselect=True,
335
+ label="Columns to display",
336
+ )
337
+ insp_table = gr.Dataframe(
338
+ label="Runs for selection (✓ = kept by dedup, ✗ = dropped)",
339
+ interactive=False,
340
+ )
341
+
342
+ # ── Event wiring ──────────────────────────────────────────────────────────
343
+
344
+ _plot_triggers = [filtered_df_state, metric_dd]
345
+
346
+ for trigger in [filtered_df_state, metric_dd]:
347
+ trigger.change(render_leaderboard, _plot_triggers, [leaderboard_table, violin_fig])
348
+ trigger.change(render_speed, _plot_triggers, speed_fig)
349
+ trigger.change(render_context, _plot_triggers, context_fig)
350
+
351
+ # Category bar chart: reacts to df, metric, AND dataset filter
352
+ _cat_triggers = [filtered_df_state, metric_dd, dataset_dd]
353
+ for trigger in [filtered_df_state, metric_dd, dataset_dd]:
354
+ trigger.change(render_per_category, _cat_triggers, category_fig)
355
+
356
+ # Per-task drill-down: reacts to category/dataset dropdowns OR global controls
357
+ _per_task_triggers = [filtered_df_state, metric_dd, category_dd, dataset_dd]
358
+ for trigger in [filtered_df_state, metric_dd, category_dd, dataset_dd]:
359
+ trigger.change(render_per_task, _per_task_triggers, per_task_fig)
360
+
361
+ # Variant heatmap: reacts to model selector, metric, filtered data, or aggregation toggle
362
+ _variant_triggers = [filtered_df_state, metric_dd, model_dd, agg_dataset_cb]
363
+ for trigger in [filtered_df_state, metric_dd, model_dd, agg_dataset_cb]:
364
+ trigger.change(render_variant_heatmap, _variant_triggers, variant_violin_fig)
365
+
366
+ # Update category dropdown choices when global filter OR dataset filter changes
367
+ for trigger in [filtered_df_state, dataset_dd]:
368
+ trigger.change(update_category_choices, [filtered_df_state, dataset_dd], category_dd)
369
+
370
+ # Settings — Apply Filters
371
+ apply_btn.click(
372
+ apply_and_preview,
373
+ inputs=[
374
+ raw_df_state, prod_only_cb, exclude_models_ms,
375
+ params_min_sl, params_max_sl,
376
+ dim_min_sl, dim_max_sl,
377
+ ctx_min_sl, ctx_max_sl,
378
+ ],
379
+ outputs=[filtered_df_state, settings_preview, settings_count],
380
+ )
381
+
382
+ # Settings — Exclude run_ids
383
+ exclude_runs_btn.click(
384
+ exclude_run_ids,
385
+ inputs=[filtered_df_state, run_ids_tb],
386
+ outputs=[filtered_df_state, settings_preview, settings_count],
387
+ )
388
+
389
+ # Inspector — cascade: model → hash → head_type → table
390
+ insp_model_dd.change(
391
+ lambda df, m: gr.update(choices=get_hash_choices_for_model(df, m), value=None),
392
+ [loaded_df_state, insp_model_dd],
393
+ insp_cache_dd,
394
+ )
395
+ insp_cache_dd.change(
396
+ lambda df, m, k: gr.update(choices=get_head_types_for_hash(df, m, k), value=None),
397
+ [loaded_df_state, insp_model_dd, insp_cache_dd],
398
+ insp_head_dd,
399
+ )
400
+ _insp_inputs = [loaded_df_state, raw_df_state, insp_model_dd, insp_cache_dd, insp_head_dd, insp_cols_dd]
401
+ insp_head_dd.change(build_inspector_table, _insp_inputs, insp_table)
402
+ insp_cols_dd.change(build_inspector_table, _insp_inputs, insp_table)
403
 
404
+ # Initial render on load
405
+ demo.load(render_leaderboard, inputs=_plot_triggers, outputs=[leaderboard_table, violin_fig])
406
+ demo.load(render_per_category, inputs=_cat_triggers, outputs=category_fig)
407
+ demo.load(render_per_task, inputs=_per_task_triggers, outputs=per_task_fig)
408
+ demo.load(render_speed, inputs=_plot_triggers, outputs=speed_fig)
409
+ demo.load(render_context, inputs=_plot_triggers, outputs=context_fig)
410
+ demo.load(render_variant_heatmap, inputs=_variant_triggers, outputs=variant_violin_fig)
411
 
 
 
 
412
 
413
+ if __name__ == "__main__":
414
+ demo.launch(theme=gr.themes.Soft(), share=True)
docs/context-specs-plan.md ADDED
@@ -0,0 +1,245 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # bench-dna Leaderboard Webapp — Plan
2
+
3
+ > **Portability note**: This doc lives in the git repo and is the canonical reference.
4
+ > When starting a new session on a different machine, paste this file as context to Claude Code.
5
+
6
+ ---
7
+
8
+ ## Overview
9
+
10
+ A Gradio webapp deployed on HuggingFace Spaces that loads benchmark results from the HF Hub
11
+ dataset (`lokahq/genomic-benchmark-metrics`) and provides interactive visualizations to compare
12
+ DNA foundation models across tasks, metrics, and efficiency dimensions.
13
+
14
+ **Stack**: Gradio (Blocks), Plotly, Pandas, HF datasets library
15
+ **Location**: repo root (`app.py` + `src/`) — HF Spaces expects root-level `app.py`
16
+ **Status**: Implemented (March 2026)
17
+
18
+ ---
19
+
20
+ ## Project Context (for new sessions)
21
+
22
+ ### Data pipeline
23
+ - Benchmark runs are produced by `bDNA/run_benchmark.py` (Hydra-based)
24
+ - Each run generates one `BenchmarkRecord` (see `bDNA/io/results.py`)
25
+ - Records are pushed to HF Hub dataset (`lokahq/genomic-benchmark-metrics`), one split per model alias
26
+ - Parquet files also saved locally at `artifacts/results/<model_alias>.parquet`
27
+
28
+ ### Current data state (as of 2026-03)
29
+ - Models with data: `dnabert2`, `caduceus`, `nt`, `omniDNA`
30
+ - All runs use `subsample_train=0.05` (5% of training data — hardware constraint)
31
+ - All runs are binary classification (short-range tasks only)
32
+ - `subsample_train IS NULL` would mean "full dataset run" — none exist yet
33
+ - Long-range tasks (regression, variant effect) not benchmarked yet
34
+
35
+ ### Key BenchmarkRecord fields for the webapp
36
+ | Field | Use |
37
+ |---|---|
38
+ | `model_alias` | Primary grouping key |
39
+ | `task_name` | Task identifier (18 short-range tasks) |
40
+ | `task_type` | `"binary"`, `"multiclass"`, `"regression"` |
41
+ | `embedding_config_hash` | Deduplication key per config variant |
42
+ | `run_at` | ISO timestamp — used to pick latest run |
43
+ | `subsample_train` | Filter for "production" vs test runs (currently all 0.05) |
44
+ | `model_params` | Filter by model size |
45
+ | `embedding_dim` | Filter by embedding size |
46
+ | `max_context_size` | Model's max sequence length |
47
+ | `mcc_test`, `accuracy_test`, `weighted_f1_test`, `macro_f1_test` | Classification metrics |
48
+ | `test_throughput_seq_s`, `test_embedding_time_s` | Speed metrics |
49
+ | `vram_model_mb`, `peak_vram_extraction_mb` | Memory metrics |
50
+
51
+ ---
52
+
53
+ ## File Structure
54
+
55
+ ```
56
+ app.py # main Gradio Blocks app
57
+ requirements.txt
58
+ src/
59
+ constants.py # task categories, seq length lookup (from JSON), metric names
60
+ data.py # HF Hub loading, dedup, filtering, best-variant logic
61
+ plots.py # all Plotly figure factories
62
+ docs/
63
+ context-specs-plan.md # this file
64
+ tools/
65
+ compute_seq_lengths.py # one-time script to populate task_seq_lengths.json
66
+ task_seq_lengths.json # committed JSON — populated by compute_seq_lengths.py
67
+ README.md # HF Spaces config (title, sdk: gradio, etc.)
68
+ ```
69
+
70
+ ---
71
+
72
+ ## Data Layer (`src/data.py`)
73
+
74
+ ### `load_data() -> pd.DataFrame`
75
+ - Reads `LEADERBOARD_DATASET` env var (set to `lokahq/genomic-benchmark-metrics`)
76
+ - Uses `datasets.get_dataset_split_names()` to list all splits (one per model alias)
77
+ - Loads each split and concatenates into one DataFrame
78
+ - Adds derived columns:
79
+ - `task_category`: mapped from `task_name` via `TASK_TO_CATEGORY`
80
+ - `task_seq_length`: mapped from `task_name` via `TASK_SEQ_LENGTHS` (from JSON)
81
+ - Called once on app startup
82
+
83
+ ### `deduplicate(df) -> pd.DataFrame`
84
+ - Group by `(model_alias, task_name, embedding_config_hash)`
85
+ - Keep the row with the latest `run_at` per group
86
+ - Eliminates accidental duplicate reruns of the exact same config
87
+
88
+ ### `get_best_variants(df, metric) -> pd.DataFrame`
89
+ - For each `(model_alias, embedding_config_hash)`, compute mean of `metric` across all tasks
90
+ - Per `model_alias`, keep only the config_hash with the highest mean
91
+ - Used when user selects "Best config per model" view
92
+
93
+ ### `apply_filters(df, filters: dict) -> pd.DataFrame`
94
+ - `subsample_train`: None=all, "full"=production runs only (IS NULL / ==0), "subsampled"=test runs
95
+ - `model_params` range (min/max)
96
+ - `embedding_dim` range
97
+ - `max_context_size` range
98
+ - `exclude_models` list
99
+ - `exclude_run_ids` list (row-level, user-managed)
100
+
101
+ > **Important**: the Settings tab default shows ALL data. "Production only"
102
+ > (subsample_train IS NULL) currently returns nothing — all records have subsample=0.05.
103
+ > The filter is there for the future when full-dataset runs exist.
104
+
105
+ ---
106
+
107
+ ## Constants (`src/constants.py`)
108
+
109
+ ```python
110
+ TASK_CATEGORIES = {
111
+ "histone_marks": [
112
+ "H2AFZ", "H3K27ac", "H3K27me3", "H3K36me3",
113
+ "H3K4me1", "H3K4me2", "H3K4me3", "H3K9ac", "H3K9me3", "H4K20me1",
114
+ ],
115
+ "promoter": ["promoter_all", "promoter_tata", "promoter_no_tata"],
116
+ "enhancer": ["enhancers", "enhancers_types"],
117
+ "splice_site": ["splice_sites_all", "splice_sites_donors", "splice_sites_acceptors"],
118
+ "variant_effect":["variant_effect_pathogenic_clinvar", "variant_effect_pathogenic_omim"],
119
+ "rna_expression":["bulk_rna_expression"],
120
+ }
121
+
122
+ # Loaded from tools/task_seq_lengths.json at runtime
123
+ TASK_SEQ_LENGTHS: dict[str, int] = json.loads(...)
124
+
125
+ CLASSIFICATION_METRICS = {
126
+ "mcc_test": "MCC",
127
+ "accuracy_test": "Accuracy",
128
+ "weighted_f1_test": "Weighted F1",
129
+ "macro_f1_test": "Macro F1",
130
+ }
131
+ ```
132
+
133
+ ---
134
+
135
+ ## Visualizations (`src/plots.py`)
136
+
137
+ | Function | Chart type | Description |
138
+ |---|---|---|
139
+ | `ranking_table` | DataFrame | Rows=models, cols=task categories (mean metric), sorted by overall mean |
140
+ | `violin_plot` | Violin + points | Distribution of metric per model across all tasks |
141
+ | `bar_plot_per_category` | Grouped bar | X=task_category, groups=model, mean metric |
142
+ | `scatter_speed` | 2-subplot scatter | Left: throughput (seq/s) vs metric; Right: embedding time (s) vs metric |
143
+ | `bubble_context_length` | Bubble chart | X=task_seq_length, Y=metric, color=model, size=max_context_size |
144
+
145
+ ### Context length bubble chart — Option B (implemented)
146
+ X=task_seq_length, Y=metric, color=model, size=max_context_size.
147
+ More directly answers "does performance degrade as task length increases, and does context size help?"
148
+ Requires `tools/task_seq_lengths.json` to be populated.
149
+
150
+ ---
151
+
152
+ ## App Layout (`app.py`)
153
+
154
+ ```
155
+ gr.Blocks(theme=gr.themes.Soft())
156
+ │
157
+ ├── Global Controls (always visible, affect all tabs)
158
+ │ ├── Metric selector (Dropdown: MCC, Accuracy, Weighted F1, Macro F1)
159
+ │ └── View toggle: "Best config per model" vs "All configs" (Radio)
160
+ │
161
+ ├── gr.State: raw_df ← loaded & deduplicated at startup
162
+ ├── gr.State: filtered_df ← written by Settings Apply
163
+ │
164
+ └── gr.Tabs
165
+ ├── "Leaderboard"
166
+ │ ├── gr.Dataframe (ranking table, sortable)
167
+ │ └── gr.Plot (violin plot)
168
+ │
169
+ ├── "Per-Category Performance"
170
+ │ └── gr.Plot (grouped bar chart by task category)
171
+ │
172
+ ├── "Speed vs Performance"
173
+ │ └── gr.Plot (2-subplot: throughput + embedding time)
174
+ │
175
+ ├── "Context Length"
176
+ │ └── gr.Plot (bubble chart — Option B)
177
+ │
178
+ └── "Settings"
179
+ ├── Checkbox: "Production runs only" (subsample_train IS NULL)
180
+ ├── Multiselect: exclude specific models
181
+ ├── Slider: model_params range
182
+ ├── Slider: embedding_dim range
183
+ ├── Slider: max_context_size range
184
+ ├── Button: "Apply Filters" → writes to filtered_df
185
+ ├── Textbox: "Paste run_ids to exclude (comma-separated)"
186
+ ├── Button: "Exclude run_ids"
187
+ └── gr.Dataframe (preview of active rows) + row count label
188
+ ```
189
+
190
+ ### State flow
191
+ 1. On load: `load_data()` → `deduplicate()` → store in `raw_df` + `filtered_df` state
192
+ 2. Settings "Apply Filters" → `apply_filters()` → update `filtered_df`
193
+ 3. All tabs: listen to `filtered_df` + global controls → regenerate plots on change
194
+
195
+ ---
196
+
197
+ ## Pre-computation (one-time, run locally before deployment)
198
+
199
+ `tools/compute_seq_lengths.py`:
200
+ - Loads each benchmark HF dataset (`InstaDeepAI/nucleotide_transformer_downstream_tasks_revised`)
201
+ - Samples ~100 sequences per task_name
202
+ - Computes mean sequence length (bp) per task_name
203
+ - **Writes result to `tools/task_seq_lengths.json`** (commit this file to the repo)
204
+ - `constants.py` reads this JSON at runtime
205
+ - **Bubble chart is a no-op until this is run and `task_seq_lengths.json` is committed**
206
+
207
+ ---
208
+
209
+ ## Environment Variables (`.env`)
210
+
211
+ ```
212
+ HF_TOKEN=<your huggingface token>
213
+ LEADERBOARD_DATASET=lokahq/genomic-benchmark-metrics
214
+ ```
215
+
216
+ ---
217
+
218
+ ## Key Design Decisions
219
+
220
+ | Decision | Choice | Rationale |
221
+ |---|---|---|
222
+ | Framework | Gradio Blocks | Native HF Spaces support |
223
+ | App placement | Root (`app.py`) | HF Spaces expects root-level app; `spaces/leaderboard/` not used |
224
+ | Data source | HF Hub on startup | Auto-updates as new runs are pushed |
225
+ | Dataset env var | `LEADERBOARD_DATASET` in `.env` | Easy to change without code edits |
226
+ | Deduplication key | `(model, task, config_hash)` latest | No prod/test flag needed |
227
+ | Best variant | Best mean-metric config per model | Consistent cross-task comparison |
228
+ | Subsample filter | Use `subsample_train` field | `run_tag` not yet implemented in BenchmarkRecord |
229
+ | Settings default | Show ALL data | All current records have subsample=0.05; "prod only" would show nothing |
230
+ | TASK_SEQ_LENGTHS | JSON file (`tools/task_seq_lengths.json`) | Committed to repo, loaded at runtime — no copy-paste into constants |
231
+ | Context length axes | Option B (X=seq_len, Y=metric, size=ctx) | More directly answers the key question |
232
+ | Regression | Deferred | No data yet |
233
+ | Row-level exclusion | Paste run_ids | Gradio DataFrames lack native per-row checkboxes |
234
+ | Submission workflow | Dropped | Not relevant to new BenchmarkRecord data pipeline |
235
+
236
+ ---
237
+
238
+ ## Out of Scope (v1)
239
+
240
+ - Regression task visualizations (long-range tasks not yet benchmarked)
241
+ - Real-time data refresh (loads on startup; user refreshes page)
242
+ - Per-subtask bar chart (18 bars is too much; use category aggregation)
243
+ - Custom embedding analysis (PCA/UMAP of raw embedding vectors)
244
+ - `run_tag` field on BenchmarkRecord (discussed, not yet implemented in benchmark runner)
245
+ - Model submission workflow (dropped in v1 remake)
pyproject.toml CHANGED
@@ -1,3 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  [tool.ruff]
2
  select = ["E", "F"]
3
  ignore = ["E501"]
 
1
+ [project]
2
+ name = "dna-benchmark"
3
+ version = "0.1.0"
4
+ requires-python = ">=3.11"
5
+ dependencies = [
6
+ "datasets>=4.0.0",
7
+ "gradio>=6.0.0",
8
+ "huggingface-hub>=0.18.0",
9
+ "plotly>=5.0.0",
10
+ "pandas",
11
+ "python-dotenv>=1.0.0",
12
+ ]
13
+
14
  [tool.ruff]
15
  select = ["E", "F"]
16
  ignore = ["E501"]
requirements.txt DELETED
@@ -1,9 +0,0 @@
1
- APScheduler
2
- datasets
3
- gradio==5.41.0
4
- gradio[oauth]
5
- gradio_leaderboard==0.0.13
6
- huggingface-hub>=0.18.0
7
- plotly>=5.0.0
8
- pandas
9
- python-dotenv==1.1.0
 
 
 
 
 
 
 
 
 
 
src/about.py DELETED
@@ -1,57 +0,0 @@
1
- from dataclasses import dataclass
2
- from enum import Enum
3
-
4
-
5
- @dataclass
6
- class Task:
7
- benchmark: str
8
- metric: str
9
- col_name: str
10
-
11
-
12
- class Tasks(Enum):
13
- task0 = Task("anli_r1", "acc", "ANLI")
14
- task1 = Task("logiqa", "acc_norm", "LogiQA")
15
-
16
-
17
- NUM_FEWSHOT = 0
18
-
19
- TITLE = """
20
- <div style="text-align: center; background: linear-gradient(135deg, #667eea 0%, #764ba2 100%); padding: 3rem 2rem; border-radius: 15px; margin-bottom: 2rem; box-shadow: 0 10px 30px rgba(0,0,0,0.2);">
21
- <h1 style="color: white; font-size: 3.5rem; margin: 0; font-weight: 800; text-shadow: 2px 2px 4px rgba(0,0,0,0.3);">
22
- 🧬 DNA Benchmark
23
- </h1>
24
- <p style="color: rgba(255,255,255,0.95); font-size: 1.3rem; margin-top: 1rem; font-weight: 300;">
25
- Evaluating DNA Foundational Models Performance
26
- </p>
27
- </div>
28
- """
29
-
30
- INTRODUCTION_TEXT = """
31
- <div style="background: linear-gradient(135deg, #f5f7fa 0%, #c3cfe2 100%); padding: 1.5rem; border-radius: 10px; border-left: 5px solid #667eea; margin-bottom: 1.5rem;">
32
- <p style="font-size: 1.1rem; color: #2d3748; margin: 0; line-height: 1.6;">
33
- Compare and analyze the performance of state-of-the-art DNA foundational models across various genomics tasks including histone modification, splicing, promoter identification, enhancer detection, and SNP classification.
34
- </p>
35
- </div>
36
- """
37
-
38
- LLM_BENCHMARKS_TEXT = """
39
- <div style="background: var(--background-fill-primary, #fff);">
40
-
41
- ## 📊 About This Leaderboard
42
-
43
- This leaderboard provides comprehensive benchmarking of DNA foundational models across multiple genomics tasks:
44
-
45
- - **🧬 Histone Modifications**: H2AFZ, H3K27ac, H3K27me3, H3K36me3, H3K4me1/2/3, H3K9ac, H3K9me3, H4K20me1
46
- - **✂️ Splicing**: Donor sites, Acceptor sites, All splice sites
47
- - **📍 Promoter Identification**: TATA-containing, Non-TATA, All promoters
48
- - **🎯 Enhancer Detection**: Enhancer identification and classification
49
- - **🔬 SNP Classification**: eQTL causality, ClinVar pathogenic variants, OMIM pathogenic variants
50
-
51
- ### 📈 Key Metrics
52
- - **Accuracy**: Overall prediction accuracy
53
- - **MCC**: Matthews Correlation Coefficient
54
- - **Weighted F1**: F1-score weighted by class support
55
-
56
- </div>
57
- """
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
src/constants.py ADDED
@@ -0,0 +1,53 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import json
2
+ import pathlib
3
+
4
+ TASK_CATEGORIES: dict[str, list[str]] = {
5
+ "histone_marks": [
6
+ "h2afz", "h3k27ac", "h3k27me3", "h3k36me3",
7
+ "h3k4me1", "h3k4me2", "h3k4me3", "h3k9ac", "h3k9me3", "h4k20me1",
8
+ ],
9
+ "promoter": ["promoter_all", "promoter_tata", "promoter_no_tata"],
10
+ "enhancer": ["enhancers", "enhancers_types"],
11
+ "splice_site": ["splice_sites_all", "splice_sites_donors", "splice_sites_acceptors"],
12
+ "variant_effect": ["variant_effect_pathogenic_clinvar", "variant_effect_pathogenic_omim"],
13
+ "rna_expression": ["bulk_rna_expression"],
14
+ }
15
+
16
+ # Inverse map: task_name -> category label
17
+ TASK_TO_CATEGORY: dict[str, str] = {
18
+ task: cat for cat, tasks in TASK_CATEGORIES.items() for task in tasks
19
+ }
20
+
21
+ # Mean sequence lengths in bp per task, populated by tools/compute_seq_lengths.py
22
+ _SEQ_LEN_PATH = pathlib.Path(__file__).parent.parent / "tools" / "task_seq_lengths.json"
23
+ TASK_SEQ_LENGTHS: dict[str, int] = (
24
+ json.loads(_SEQ_LEN_PATH.read_text()) if _SEQ_LEN_PATH.exists() else {}
25
+ )
26
+
27
+ CLASSIFICATION_METRICS: dict[str, str] = {
28
+ "mcc_test": "MCC",
29
+ "accuracy_test": "Accuracy",
30
+ "weighted_f1_test": "Weighted F1",
31
+ "macro_f1_test": "Macro F1",
32
+ }
33
+
34
+ CATEGORY_DISPLAY_NAMES: dict[str, str] = {
35
+ "histone_marks": "Histone Marks",
36
+ "promoter": "Promoter",
37
+ "enhancer": "Enhancer",
38
+ "splice_site": "Splice Site",
39
+ "variant_effect": "Variant Effect",
40
+ "rna_expression": "RNA Expression",
41
+ "other": "Other",
42
+ }
43
+
44
+ # One color per task category — pastel, full-opacity friendly, distinct from _MODEL_PALETTE
45
+ CATEGORY_COLORS: dict[str, str] = {
46
+ "histone_marks": "#8EB4E8",
47
+ "promoter": "#F5956C",
48
+ "enhancer": "#60C4A8",
49
+ "splice_site": "#5AC8D8",
50
+ "variant_effect": "#C49AE0",
51
+ "rna_expression": "#F075A0",
52
+ "other": "#BCC4CC",
53
+ }
src/data.py ADDED
@@ -0,0 +1,103 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+
3
+ import pandas as pd
4
+ from datasets import get_dataset_split_names, load_dataset
5
+ from dotenv import load_dotenv
6
+
7
+ from src.constants import TASK_SEQ_LENGTHS, TASK_TO_CATEGORY
8
+
9
+ load_dotenv()
10
+
11
+ HF_TOKEN = os.environ.get("HF_TOKEN")
12
+ LEADERBOARD_DATASET = os.environ.get("LEADERBOARD_DATASET", "lokahq/genomic-benchmark-metrics")
13
+
14
+
15
+ def load_data() -> pd.DataFrame:
16
+ """Load all model splits from the HF Hub dataset and return a concatenated DataFrame.
17
+
18
+ Each split corresponds to one model alias. Derived columns task_category and
19
+ task_seq_length are added from the constants maps.
20
+ """
21
+ splits = get_dataset_split_names(LEADERBOARD_DATASET, token=HF_TOKEN)
22
+ frames = []
23
+ for split in splits:
24
+ ds = load_dataset(LEADERBOARD_DATASET, split=split, token=HF_TOKEN)
25
+ frames.append(ds.to_pandas())
26
+ df = pd.concat(frames, ignore_index=True)
27
+
28
+ # Normalize task names to lowercase so TASK_TO_CATEGORY mapping works regardless
29
+ # of how the benchmark runner capitalizes them (e.g. "H3K27ac" → "h3k27ac")
30
+ df["task_name"] = df["task_name"].str.lower()
31
+
32
+ df["task_category"] = df["task_name"].map(TASK_TO_CATEGORY).fillna("other")
33
+ df["task_seq_length"] = df["task_name"].map(TASK_SEQ_LENGTHS)
34
+
35
+ if "run_at" in df.columns:
36
+ df["run_at"] = pd.to_datetime(df["run_at"], utc=True, errors="coerce")
37
+
38
+ return df
39
+
40
+
41
+ def _subsample_priority(val) -> float:
42
+ """Map subsample_train value to sort priority. NULL/0/1 = production = highest priority."""
43
+ if val is None or (isinstance(val, float) and pd.isna(val)) or val == 0 or val == 1:
44
+ return float("inf")
45
+ return float(val)
46
+
47
+
48
+ def deduplicate(df: pd.DataFrame) -> pd.DataFrame:
49
+ """Keep best run per (embedding_config_hash, head_type).
50
+
51
+ embedding_config_hash encodes model, pooling, layer selection (+ task via the cache key it hashes).
52
+ 'Best' = production run preferred (subsample_train NULL/0/1 → inf priority),
53
+ then highest subsample fraction, then most recent run_at.
54
+ """
55
+ if "run_at" not in df.columns:
56
+ return df
57
+ dedup_cols = []
58
+ if "embedding_config_hash" in df.columns:
59
+ dedup_cols.append("embedding_config_hash")
60
+ if "head_type" in df.columns:
61
+ dedup_cols.append("head_type")
62
+ if not dedup_cols:
63
+ return df
64
+ df = df.copy()
65
+ df["_sub_priority"] = df.get("subsample_train", pd.Series(dtype=float)).map(_subsample_priority)
66
+ df = df.sort_values(["_sub_priority", "run_at"], ascending=[False, False])
67
+ return df.drop_duplicates(subset=dedup_cols).drop(columns=["_sub_priority"]).reset_index(drop=True)
68
+
69
+
70
+
71
+ def apply_filters(df: pd.DataFrame, filters: dict) -> pd.DataFrame:
72
+ """Apply user-selected filters and return the filtered DataFrame."""
73
+ result = df.copy()
74
+
75
+ subsample_mode = filters.get("subsample_train")
76
+ if subsample_mode == "full" and "subsample_train" in result.columns:
77
+ result = result[result["subsample_train"].isna() | (result["subsample_train"] == 0)]
78
+ elif subsample_mode == "subsampled" and "subsample_train" in result.columns:
79
+ result = result[result["subsample_train"].notna() & (result["subsample_train"] > 0)]
80
+
81
+ for col, min_key, max_key in [
82
+ ("model_params", "model_params_min", "model_params_max"),
83
+ ("embedding_dim", "embedding_dim_min", "embedding_dim_max"),
84
+ ("max_context_size", "max_context_size_min", "max_context_size_max"),
85
+ ]:
86
+ if col not in result.columns:
87
+ continue
88
+ lo = filters.get(min_key)
89
+ hi = filters.get(max_key)
90
+ if lo is not None:
91
+ result = result[result[col] >= lo]
92
+ if hi is not None:
93
+ result = result[result[col] <= hi]
94
+
95
+ excluded_models = filters.get("exclude_models") or []
96
+ if excluded_models:
97
+ result = result[~result["model_alias"].isin(excluded_models)]
98
+
99
+ excluded_runs = filters.get("exclude_run_ids") or []
100
+ if excluded_runs and "run_id" in result.columns:
101
+ result = result[~result["run_id"].isin(excluded_runs)]
102
+
103
+ return result.reset_index(drop=True)
src/display/css_html_js.py DELETED
@@ -1,64 +0,0 @@
1
- custom_css = """
2
- .gradio-container {
3
- max-width: 1400px !important;
4
- margin: 0 auto !important;
5
- font-family: 'Inter', -apple-system, BlinkMacSystemFont, 'Segoe UI', sans-serif !important;
6
- }
7
-
8
- .markdown-text {
9
- font-size: 16px !important;
10
- line-height: 1.7 !important;
11
- }
12
-
13
- body:not(.dark) .markdown-text {
14
- color: #2d3748 !important;
15
- }
16
-
17
- body:not(.dark) .markdown-text h2 {
18
- color: #667eea !important;
19
- font-weight: 700 !important;
20
- margin-top: 1.5rem !important;
21
- margin-bottom: 1rem !important;
22
- }
23
-
24
- /* Table: only font weights */
25
- table tbody td,
26
- table tbody th {
27
- font-weight: 400 !important;
28
- }
29
- table thead th {
30
- font-weight: 600 !important;
31
- }
32
- """
33
- custom_css = """
34
- .gradio-container {
35
- max-width: 1400px !important;
36
- margin: 0 auto !important;
37
- font-family: 'Inter', -apple-system, BlinkMacSystemFont, 'Segoe UI', sans-serif !important;
38
- }
39
-
40
- .markdown-text {
41
- font-size: 16px !important;
42
- line-height: 1.7 !important;
43
- }
44
-
45
- body:not(.dark) .markdown-text {
46
- color: #2d3748 !important;
47
- }
48
-
49
- body:not(.dark) .markdown-text h2 {
50
- color: #667eea !important;
51
- font-weight: 700 !important;
52
- margin-top: 1.5rem !important;
53
- margin-bottom: 1rem !important;
54
- }
55
-
56
- /* Table: only font weights */
57
- table tbody td,
58
- table tbody th {
59
- font-weight: 400 !important;
60
- }
61
- table thead th {
62
- font-weight: 600 !important;
63
- }
64
- """
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
src/display/formatting.py DELETED
@@ -1,27 +0,0 @@
1
- def model_hyperlink(link, model_name):
2
- return f'<a target="_blank" href="{link}" style="color: var(--link-text-color); text-decoration: underline;text-decoration-style: dotted;">{model_name}</a>'
3
-
4
-
5
- def make_clickable_model(model_name):
6
- link = f"https://huggingface.co/{model_name}"
7
- return model_hyperlink(link, model_name)
8
-
9
-
10
- def styled_error(error):
11
- return f"<p style='color: red; font-size: 20px; text-align: center;'>{error}</p>"
12
-
13
-
14
- def styled_warning(warn):
15
- return f"<p style='color: orange; font-size: 20px; text-align: center;'>{warn}</p>"
16
-
17
-
18
- def styled_message(message):
19
- return f"<p style='color: green; font-size: 20px; text-align: center;'>{message}</p>"
20
-
21
-
22
- def has_no_nan_values(df, columns):
23
- return df[columns].notna().all(axis=1)
24
-
25
-
26
- def has_nan_values(df, columns):
27
- return df[columns].isna().any(axis=1)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
src/display/plotting.py DELETED
@@ -1,205 +0,0 @@
1
- import plotly.express as px
2
- import plotly.graph_objects as go
3
- import pandas as pd
4
-
5
- METRICS_FOR_PLOTS = ["Accuracy", "MCC", "Weighted F1"]
6
-
7
- # Color scheme matching the UI theme
8
- COLORS = {
9
- 'primary': '#667eea',
10
- 'secondary': '#764ba2',
11
- 'accent': '#F97316',
12
- 'gradient': ['#667eea', '#764ba2', '#F97316', '#06b6d4', '#10b981']
13
- }
14
-
15
- def extract_mean_std(col):
16
- if isinstance(col, str) and '±' in col:
17
- try:
18
- mean, std = col.split('±')
19
- return float(mean.strip()), float(std.strip())
20
- except:
21
- return None, None
22
- elif isinstance(col, (int, float)):
23
- return col, 0
24
- return None, None
25
-
26
- def prepare_leaderboard_df(df):
27
- for metric in METRICS_FOR_PLOTS:
28
- means, stds = zip(*df[metric].apply(extract_mean_std))
29
- df[f"{metric}_mean"] = means
30
- df[f"{metric}_std"] = stds
31
- return df
32
-
33
- def make_plot_wrapper(leaderboard_df, group_by, filter_col, orientation):
34
- def plot_fn(filter_val, metric, dataset):
35
- return plot_metric_bar(
36
- df=leaderboard_df,
37
- group_by=group_by,
38
- filter_col=filter_col,
39
- dataset=str(dataset),
40
- filter_val=filter_val,
41
- metric=metric,
42
- orientation=orientation
43
- )
44
- return plot_fn
45
-
46
- def make_model_all_datasets_wrapper(leaderboard_df):
47
- def plot_fn(model, metric):
48
- return plot_metric_bar_all_datasets(
49
- df=leaderboard_df,
50
- model=model,
51
- metric=metric
52
- )
53
- return plot_fn
54
-
55
- def plot_metric_bar_all_datasets(df, model, metric, orientation="v"):
56
- df = df.copy()
57
- df = df[df['Model'] == model]
58
-
59
- y_col = f"{metric}_mean"
60
- std_col = f"{metric}_std"
61
-
62
- if df.empty or y_col not in df.columns:
63
- return px.bar(title=f"No data found for model: {model}")
64
-
65
- df = df.sort_values(by=y_col, ascending=False)
66
-
67
- x_axis, y_axis = (y_col, "Task") if orientation == "h" else ("Task", y_col)
68
- height = 20 * len(df) + 200 if orientation == "h" else None
69
-
70
- fig = px.bar(
71
- df,
72
- x=x_axis,
73
- y=y_axis,
74
- color='Dataset Name',
75
- orientation=orientation,
76
- text_auto=".2f",
77
- labels={y_col: metric, "Task": "Task"},
78
- title=f"{metric} for {model} across all datasets",
79
- height=height,
80
- color_discrete_sequence=COLORS['gradient']
81
- )
82
-
83
- if orientation == "h":
84
- hovertemplate = (
85
- "<b>%{y}</b><br>"
86
- f"Mean {metric}: %{{x:.2f}}<br>"
87
- "Std Dev: %{customdata[0]:.3f}"
88
- "<extra></extra>"
89
- )
90
- else:
91
- hovertemplate = (
92
- "<b>%{x}</b><br>"
93
- f"Mean {metric}: %{{y:.2f}}<br>"
94
- "Std Dev: %{customdata[0]:.3f}"
95
- "<extra></extra>"
96
- )
97
-
98
- fig.update_traces(
99
- hovertemplate=hovertemplate,
100
- customdata=df[[std_col]].values
101
- )
102
-
103
- layout_args = dict(
104
- title={
105
- 'text': f"{metric} for {model} across all tasks",
106
- 'x': 0.5,
107
- 'xanchor': 'center',
108
- 'yanchor': 'top',
109
- 'font': dict(size=22, color='#2d3748', family='Inter, sans-serif'),
110
- 'pad': dict(t=20, b=10),
111
- },
112
- margin=dict(t=80, b=100, l=150, r=10),
113
- plot_bgcolor='rgba(0,0,0,0)',
114
- paper_bgcolor='white',
115
- font=dict(family='Inter, sans-serif', color='#2d3748'),
116
- )
117
-
118
- if orientation == "h":
119
- layout_args["xaxis"] = dict(range=[0, 1], gridcolor='#e2e8f0')
120
- layout_args["yaxis_tickfont_size"] = 10
121
- layout_args["yaxis"] = dict(gridcolor='#e2e8f0')
122
- else:
123
- layout_args["yaxis"] = dict(range=[0, 1], gridcolor='#e2e8f0')
124
- layout_args["xaxis"] = dict(gridcolor='#e2e8f0')
125
-
126
- fig.update_layout(**layout_args)
127
-
128
- return fig
129
-
130
- def plot_metric_bar(df, group_by, filter_col, filter_val, dataset, metric, orientation="v"):
131
- df = df.copy()
132
- df = df[df['Dataset Name'] == dataset]
133
- df = df[df[filter_col] == filter_val]
134
-
135
- y_col = f"{metric}_mean"
136
- std_col = f"{metric}_std"
137
-
138
- if df.empty or y_col not in df.columns:
139
- return px.bar(title="No data found.")
140
-
141
- df = df.sort_values(by=y_col, ascending=False)
142
-
143
- x_axis, y_axis = (y_col, group_by) if orientation == "h" else (group_by, y_col)
144
- height = 20 * len(df) + 200 if orientation == "h" else None
145
-
146
- fig = px.bar(
147
- df,
148
- x=x_axis,
149
- y=y_axis,
150
- orientation=orientation,
151
- text_auto=".2f",
152
- labels={y_col: metric, group_by: group_by},
153
- title=f"{metric} for {filter_val} - {dataset}",
154
- height=height,
155
- )
156
-
157
- if orientation == "h":
158
- hovertemplate = (
159
- "<b>%{y}</b><br>"
160
- f"Mean {metric}: %{{x:.2f}}<br>"
161
- "Std Dev: %{customdata[0]:.3f}"
162
- "<extra></extra>"
163
- )
164
- else:
165
- hovertemplate = (
166
- "<b>%{x}</b><br>"
167
- f"Mean {metric}: %{{y:.2f}}<br>"
168
- "Std Dev: %{customdata[0]:.3f}"
169
- "<extra></extra>"
170
- )
171
-
172
- fig.update_traces(
173
- marker_color=COLORS['primary'],
174
- marker_line_color=COLORS['secondary'],
175
- marker_line_width=1.5,
176
- hovertemplate=hovertemplate,
177
- customdata=df[[std_col]].values
178
- )
179
-
180
- layout_args = dict(
181
- title={
182
- 'text': f"{metric} for {filter_val} - {dataset}",
183
- 'x': 0.5,
184
- 'xanchor': 'center',
185
- 'yanchor': 'top',
186
- 'font': dict(size=22, color='#2d3748', family='Inter, sans-serif'),
187
- 'pad': dict(t=20, b=10),
188
- },
189
- margin=dict(t=80, b=100, l=150, r=10),
190
- plot_bgcolor='rgba(0,0,0,0)',
191
- paper_bgcolor='white',
192
- font=dict(family='Inter, sans-serif', color='#2d3748'),
193
- )
194
-
195
- if orientation == "h":
196
- layout_args["xaxis"] = dict(range=[0, 1], gridcolor='#e2e8f0')
197
- layout_args["yaxis_tickfont_size"] = 10
198
- layout_args["yaxis"] = dict(gridcolor='#e2e8f0')
199
- else:
200
- layout_args["yaxis"] = dict(range=[0, 1], gridcolor='#e2e8f0')
201
- layout_args["xaxis"] = dict(gridcolor='#e2e8f0')
202
-
203
- fig.update_layout(**layout_args)
204
-
205
- return fig
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
src/display/utils.py DELETED
@@ -1,65 +0,0 @@
1
- from dataclasses import dataclass
2
-
3
- from src.about import Tasks
4
-
5
-
6
- def fields(raw_class):
7
- return [v for k, v in raw_class.__dict__.items() if k[:2] != "__" and k[-2:] != "__"]
8
-
9
-
10
- # These classes are for user facing column names,
11
- # to avoid having to change them all around the code
12
- # when a modif is needed
13
- @dataclass
14
- class ColumnContent:
15
- raw_name: str
16
- name: str
17
- type: str
18
- displayed_by_default: bool
19
- hidden: bool = False
20
- never_hidden: bool = False
21
-
22
-
23
- ## For the queue columns in the submission tab
24
- @dataclass(frozen=True)
25
- class BenchRawColumn: # Queue column
26
- model = ColumnContent("hf_model", "Model", "str", True)
27
- task = ColumnContent("task", "Task", "markdown", False)
28
- train_len = ColumnContent("train_len", "Train Samples", "number", False)
29
- test_len = ColumnContent("test_len", "Test Samples", "number", False)
30
- train_proc_time = ColumnContent("train_proc_time", "Train Processing (s)", "number", False)
31
- test_proc_time = ColumnContent("test_proc_time", "Test Processing (s)", "number", False)
32
- input_length = ColumnContent("input_length", "Number of Samples", "number", False)
33
- acc = ColumnContent("accuracy", "Accuracy", "number", True)
34
- mcc = ColumnContent("mcc", "MCC", "number", True)
35
- infer_time = ColumnContent("inference_time", "Inference Time (s)", "number", False)
36
- embds_dim = ColumnContent("embeddings_dim", "Embds Dim", "number", False)
37
- model_params = ColumnContent("model_params", "Model Params (M)", "number", True)
38
- vram_model = ColumnContent("vram_model", "VRAM Model (MB)", "number", False)
39
- classification_report = ColumnContent("classification_report", "Classification Report", "str", False, hidden=True)
40
- config = ColumnContent("config", "Config", "str", False, hidden=True)
41
- max_context_len = ColumnContent("max_context_size", "Max Context Length (bp)", "number", False)
42
- dataset_name = ColumnContent("dataset", "Dataset Name", "str", False)
43
- weighted_f1 = ColumnContent("weighted f1", "Weighted F1", "number", True)
44
-
45
-
46
- BenchTableModel = BenchRawColumn()
47
-
48
- # Column selection
49
- COLS = [c.name for c in fields(BenchRawColumn) if not c.hidden]
50
- METRICS_COLS = [BenchRawColumn.acc.name, BenchRawColumn.mcc.name, BenchRawColumn.weighted_f1.name]
51
- BENCHMARK_COLS = [t.value.col_name for t in Tasks]
52
-
53
- # For summarizing model performance per task type
54
- COLS_TO_AVERAGE = ['Accuracy', 'MCC', 'Weighted F1', 'Train Processing (s)', 'Test Processing (s)', 'Inference Time (s)',
55
- 'Number of Samples', 'Train Samples', 'Test Samples']
56
- COLS_DEPEND_ON_MODEL = ['Model Params (M)', 'Embds Dim', 'VRAM Model (MB)', "Max Context Length (bp)", "Dataset Name"]
57
- TASK_TYPE_MAP = {
58
- "H2AFZ": "Histone", "H3K27ac": "Histone", "H3K27me3": "Histone", "H3K36me3": "Histone",
59
- "H3K4me1": "Histone", "H3K4me2": "Histone", "H3K4me3": "Histone", "H3K9ac": "Histone",
60
- "H3K9me3": "Histone", "H4K20me1": "Histone",
61
- "splice_sites_donors": "Splicing", "splice_sites_acceptors": "Splicing", "splice_sites_all": "Splicing",
62
- "promoter_no_tata": "Promoter", "promoter_tata": "Promoter", "promoter_all": "Promoter",
63
- "enhancers": "Enhancer", "enhancers_types": "Enhancer", "variant_effect_causal_eqtl": "SNP Classification",
64
- "variant_effect_pathogenic_clinvar": "SNP Classification", "variant_effect_pathogenic_omim": "SNP Classification"
65
- }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
src/envs.py DELETED
@@ -1,29 +0,0 @@
1
- import os
2
-
3
- from dotenv import load_dotenv
4
- from huggingface_hub import HfApi
5
-
6
- # Load environment variables from .env file
7
- load_dotenv()
8
-
9
- # Info to change for your repository
10
- # ----------------------------------
11
- TOKEN = os.environ.get("HF_TOKEN") # A read/write token for your org
12
-
13
- OWNER = "lokahq" # Change to your org - don't forget to create a results and request dataset, with the correct format!
14
- # ----------------------------------
15
-
16
- REPO_ID = f"{OWNER}/dna-benchmark"
17
- QUEUE_REPO = f"{OWNER}/requests"
18
- RESULTS_REPO = f"{OWNER}/bench-dna-results"
19
-
20
- # If you setup a cache later, just change HF_HOME
21
- CACHE_PATH = os.getenv("HF_HOME", ".")
22
-
23
- # Local caches
24
- EVAL_REQUESTS_PATH = os.path.join(CACHE_PATH, "eval-queue")
25
- EVAL_RESULTS_PATH = os.path.join(CACHE_PATH, "eval-results")
26
- EVAL_REQUESTS_PATH_BACKEND = os.path.join(CACHE_PATH, "eval-queue-bk")
27
- EVAL_RESULTS_PATH_BACKEND = os.path.join(CACHE_PATH, "eval-results-bk")
28
-
29
- API = HfApi(token=TOKEN)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
src/plots.py ADDED
@@ -0,0 +1,598 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Plotly figure factories for the DNA benchmark leaderboard."""
2
+
3
+ import json
4
+ import numpy as np
5
+ import pandas as pd
6
+ import plotly.graph_objects as go
7
+ from plotly.subplots import make_subplots
8
+
9
+ from src.constants import (
10
+ CATEGORY_COLORS,
11
+ CATEGORY_DISPLAY_NAMES,
12
+ CLASSIFICATION_METRICS,
13
+ TASK_SEQ_LENGTHS,
14
+ )
15
+
16
+ # Model colour palette — pastel, full-opacity friendly, distinct from CATEGORY_COLORS
17
+ _MODEL_PALETTE = ["#F08E9E", "#64BDAD", "#F7C96A", "#7FAACC", "#B09AD8", "#80C89A", "#F5A978"]
18
+
19
+ # Shared horizontal legend layout used across multiple figures.
20
+ # Explicit defaults: single-click toggles a trace, double-click isolates it.
21
+ _LEGEND_H = dict(
22
+ orientation="h", yanchor="bottom", y=1.02, xanchor="right", x=1,
23
+ itemclick="toggle", itemdoubleclick="toggleothers",
24
+ )
25
+
26
+
27
+ def _model_color_map(models: list[str]) -> dict[str, str]:
28
+ return {m: _MODEL_PALETTE[i % len(_MODEL_PALETTE)] for i, m in enumerate(sorted(models))}
29
+
30
+
31
+ def _fmt_layer(params_str) -> str:
32
+ """Short label for layer selection from layer_selection_params JSON string."""
33
+ if params_str is None or (isinstance(params_str, float) and pd.isna(params_str)):
34
+ return "default"
35
+ try:
36
+ params = json.loads(params_str) if isinstance(params_str, str) else params_str
37
+ if not params:
38
+ return "default"
39
+ idx = params.get("layer_index")
40
+ strategy = params.get("strategy") or params.get("layer_selection_strategy")
41
+ parts = []
42
+ if strategy:
43
+ parts.append(str(strategy))
44
+ if idx is not None:
45
+ parts.append(f"({int(idx)})")
46
+ return " ".join(parts) if parts else "default"
47
+ except (json.JSONDecodeError, TypeError):
48
+ return "default"
49
+
50
+
51
+ def _variant_label(row: pd.Series) -> str:
52
+ """Human-readable config label from a data row: 'short-hf-id · pooling · layer · head'."""
53
+ parts = []
54
+ hf = row.get("model_hf_id")
55
+ if hf is not None and not (isinstance(hf, float) and pd.isna(hf)):
56
+ parts.append(str(hf).rsplit("/", 1)[-1])
57
+ pool = row.get("pooling_strategy")
58
+ if pool is not None and not (isinstance(pool, float) and pd.isna(pool)):
59
+ parts.append(str(pool))
60
+ parts.append(_fmt_layer(row.get("layer_selection_params")))
61
+ head = row.get("head_type")
62
+ if head is not None and not (isinstance(head, float) and pd.isna(head)):
63
+ parts.append(str(head))
64
+ return " · ".join(parts)
65
+
66
+
67
+ def _variant_hover_lines(vdf: pd.DataFrame) -> str:
68
+ """Multi-line config block for heatmap hover, sourced from first row of variant group."""
69
+ row = vdf.iloc[0]
70
+ lines = []
71
+ hf = row.get("model_hf_id")
72
+ if hf is not None and not (isinstance(hf, float) and pd.isna(hf)):
73
+ lines.append(f"HF model: {hf}")
74
+ pool = row.get("pooling_strategy")
75
+ if pool is not None and not (isinstance(pool, float) and pd.isna(pool)):
76
+ lines.append(f"Pooling: {pool}")
77
+ lines.append(f"Layer: {_fmt_layer(row.get('layer_selection_params'))}")
78
+ head = row.get("head_type")
79
+ if head is not None and not (isinstance(head, float) and pd.isna(head)):
80
+ lines.append(f"Head: {head}")
81
+ return "<br>".join(lines)
82
+
83
+
84
+ # Hover label for config customdata columns
85
+ _CONFIG_HOVER_LABELS: dict[str, str] = {
86
+ "model_hf_id": "HF model",
87
+ "pooling_strategy": "Pooling",
88
+ "_layer_info": "Layer",
89
+ }
90
+
91
+
92
+
93
+ # ---------------------------------------------------------------------------
94
+ # 1. Ranking table
95
+ # ---------------------------------------------------------------------------
96
+
97
+ def ranking_table(df: pd.DataFrame, metric: str) -> pd.DataFrame:
98
+ """Return a summary DataFrame: rows=models, cols=task categories + Overall."""
99
+ if df.empty or metric not in df.columns:
100
+ return pd.DataFrame()
101
+
102
+ records = []
103
+ for model, group in df.groupby("model_alias"):
104
+ row: dict = {"Model": model}
105
+ for cat_key, cat_label in CATEGORY_DISPLAY_NAMES.items():
106
+ if cat_key == "other":
107
+ continue
108
+ cat_data = group[group["task_category"] == cat_key][metric]
109
+ row[cat_label] = round(cat_data.mean(), 3) if not cat_data.empty else None
110
+ row["Overall"] = round(group[metric].mean(), 3)
111
+ records.append(row)
112
+
113
+ table = pd.DataFrame(records)
114
+ # Drop columns where all values are None (no data for that category yet)
115
+ cat_cols = [col for col in table.columns if col not in ("Model", "Overall")]
116
+ table = table.drop(columns=[col for col in cat_cols if table[col].isna().all()])
117
+ if "Overall" in table.columns:
118
+ table = table.sort_values("Overall", ascending=False).reset_index(drop=True)
119
+ return table
120
+
121
+
122
+ # ---------------------------------------------------------------------------
123
+ # Internal helper: violin shape + category-coloured scatter overlay
124
+ # ---------------------------------------------------------------------------
125
+
126
+ def _violin_with_category_points(
127
+ fig: go.Figure,
128
+ x_pos: int,
129
+ subset: pd.DataFrame,
130
+ metric: str,
131
+ metric_label: str,
132
+ violin_color: str,
133
+ shown_cats: set[str],
134
+ legend_name: str,
135
+ customdata_cols: list[str] | None = None,
136
+ ) -> None:
137
+ """Add one violin (no built-in points) + jittered scatter overlay coloured by task category.
138
+
139
+ Uses numeric x positions so scatter points can be spread horizontally with jitter.
140
+ Callers must set xaxis tickvals/ticktext after all traces are added.
141
+ """
142
+ # Violin shape only (no built-in points)
143
+ fig.add_trace(
144
+ go.Violin(
145
+ y=subset[metric],
146
+ x=[x_pos] * len(subset),
147
+ name=legend_name,
148
+ box_visible=True,
149
+ meanline=dict(visible=True, color="white", width=3),
150
+ points=False,
151
+ line_color=violin_color,
152
+ fillcolor=violin_color,
153
+ opacity=0.5,
154
+ showlegend=True,
155
+ )
156
+ )
157
+
158
+ rng = np.random.default_rng(seed=x_pos)
159
+
160
+ # Category-coloured scatter points with horizontal jitter
161
+ for cat, cat_df in subset.groupby("task_category"):
162
+ color = CATEGORY_COLORS.get(cat, CATEGORY_COLORS["other"])
163
+ cat_label = CATEGORY_DISPLAY_NAMES.get(cat, cat)
164
+ show_legend = cat not in shown_cats
165
+ shown_cats.add(cat)
166
+
167
+ n = len(cat_df)
168
+ x_jitter = rng.uniform(x_pos - 0.18, x_pos + 0.18, n)
169
+
170
+ hover = f"<b>%{{text}}</b><br>{metric_label}: %{{y:.3f}}"
171
+ cd = None
172
+ if customdata_cols:
173
+ available = [c for c in customdata_cols if c in cat_df.columns]
174
+ if available:
175
+ cd = cat_df[available].to_numpy(dtype=object)
176
+ for i, col in enumerate(available):
177
+ label = _CONFIG_HOVER_LABELS.get(col, col)
178
+ hover += f"<br>{label}: %{{customdata[{i}]}}"
179
+ hover += "<extra></extra>"
180
+
181
+ fig.add_trace(
182
+ go.Scatter(
183
+ x=x_jitter,
184
+ y=cat_df[metric],
185
+ mode="markers",
186
+ name=cat_label,
187
+ marker=dict(color=color, size=9, opacity=0.9, line=dict(width=1, color="white")),
188
+ text=cat_df["task_name"],
189
+ hovertemplate=hover,
190
+ customdata=cd,
191
+ legendgroup="cat_" + cat,
192
+ showlegend=show_legend,
193
+ )
194
+ )
195
+
196
+
197
+ # ---------------------------------------------------------------------------
198
+ # 2. Violin plot — all models, points coloured by task category
199
+ # ---------------------------------------------------------------------------
200
+
201
+ def violin_plot(df: pd.DataFrame, metric: str) -> go.Figure:
202
+ """Violin plot per model with points coloured by task category and task hover text."""
203
+ if df.empty or metric not in df.columns:
204
+ return go.Figure()
205
+
206
+ models = sorted(df["model_alias"].unique())
207
+ color_map = _model_color_map(models)
208
+ metric_label = CLASSIFICATION_METRICS.get(metric, metric)
209
+
210
+ # Config columns to expose in hover
211
+ _CONFIG_RAW = ["model_hf_id", "pooling_strategy", "layer_selection_params"]
212
+ config_raw = [c for c in _CONFIG_RAW if c in df.columns]
213
+
214
+ fig = go.Figure()
215
+ shown_cats: set[str] = set()
216
+
217
+ for i, model in enumerate(models):
218
+ cols = ["task_name", "task_category", metric] + config_raw
219
+ subset = (
220
+ df[df["model_alias"] == model][cols]
221
+ .dropna(subset=[metric])
222
+ .reset_index(drop=True)
223
+ .copy()
224
+ )
225
+
226
+ customdata_cols: list[str] = []
227
+ if "model_hf_id" in subset.columns:
228
+ customdata_cols.append("model_hf_id")
229
+ if "pooling_strategy" in subset.columns:
230
+ customdata_cols.append("pooling_strategy")
231
+ if "layer_selection_params" in subset.columns:
232
+ subset["_layer_info"] = subset["layer_selection_params"].apply(_fmt_layer)
233
+ customdata_cols.append("_layer_info")
234
+
235
+ _violin_with_category_points(
236
+ fig, x_pos=i, subset=subset, metric=metric,
237
+ metric_label=metric_label, violin_color=color_map[model],
238
+ shown_cats=shown_cats, legend_name=model,
239
+ customdata_cols=customdata_cols or None,
240
+ )
241
+
242
+ fig.update_layout(
243
+ title=f"{metric_label} distribution per model (points coloured by task category)",
244
+ yaxis_title=metric_label,
245
+ height=500,
246
+ template="plotly_white",
247
+ legend=_LEGEND_H,
248
+ )
249
+ fig.update_xaxes(
250
+ tickmode="array",
251
+ tickvals=list(range(len(models))),
252
+ ticktext=models,
253
+ )
254
+ return fig
255
+
256
+
257
+ # ---------------------------------------------------------------------------
258
+ # 3. Heatmap — single model, rows = config variants, cols = task categories
259
+ # ---------------------------------------------------------------------------
260
+
261
+ def heatmap_variants(
262
+ df: pd.DataFrame,
263
+ metric: str,
264
+ model_alias: str,
265
+ aggregate_by_dataset: bool = False,
266
+ ) -> go.Figure:
267
+ """Heatmap: rows = config variants (or single row if only one), cols = tasks (or datasets) + Overall.
268
+
269
+ By default columns are individual task_names ordered by category then alphabetically.
270
+ When aggregate_by_dataset=True columns collapse to dataset_key values (mean per dataset).
271
+ Rows sorted best→worst by Overall when multiple variants exist.
272
+ """
273
+ metric_label = CLASSIFICATION_METRICS.get(metric, metric)
274
+
275
+ if not model_alias or df.empty or metric not in df.columns:
276
+ return go.Figure()
277
+
278
+ mdf = df[df["model_alias"] == model_alias]
279
+ if mdf.empty:
280
+ return go.Figure()
281
+
282
+ # Group by config fields that are task-agnostic (embedding_cache_key includes task name)
283
+ _KEY_COLS = ["pooling_strategy", "layer_selection_params", "head_type"]
284
+ key_cols = [c for c in _KEY_COLS if c in mdf.columns]
285
+ mdf = mdf.copy()
286
+ mdf["_embed_key"] = mdf[key_cols].apply(lambda r: str(tuple(r.values)), axis=1)
287
+ configs = sorted(mdf["_embed_key"].unique())
288
+
289
+ if len(configs) > 1:
290
+ rows_iter = [
291
+ (_variant_label(mdf[mdf["_embed_key"] == c].iloc[0]), mdf[mdf["_embed_key"] == c])
292
+ for c in configs
293
+ ]
294
+ else:
295
+ rows_iter = [(_variant_label(mdf.iloc[0]), mdf)]
296
+
297
+ # Build column list depending on aggregation mode
298
+ if aggregate_by_dataset:
299
+ # Collapse to dataset_name (full HF path, globally unique); one column per dataset
300
+ col_keys = sorted(mdf["dataset_name"].dropna().unique()) if "dataset_name" in mdf.columns else []
301
+ col_names = col_keys
302
+
303
+ def _cell_vals(vdf: pd.DataFrame, col_key: str) -> pd.Series:
304
+ return vdf[vdf["dataset_name"] == col_key][metric].dropna()
305
+ else:
306
+ # Individual task_names ordered by category then alphabetically
307
+ col_keys, col_names = [], []
308
+ for cat in CATEGORY_DISPLAY_NAMES:
309
+ for task in sorted(mdf[mdf["task_category"] == cat]["task_name"].unique()):
310
+ col_keys.append(task)
311
+ col_names.append(task)
312
+
313
+ def _cell_vals(vdf: pd.DataFrame, col_key: str) -> pd.Series:
314
+ return vdf[vdf["task_name"] == col_key][metric].dropna()
315
+
316
+ col_names = col_names + ["Overall"]
317
+
318
+ rows_z, rows_text, rows_hover, row_labels = [], [], [], []
319
+ for label, vdf in rows_iter:
320
+ z_row, t_row, h_row = [], [], []
321
+ config_lines = _variant_hover_lines(vdf)
322
+ for col_key, col_name in zip(col_keys, col_names[:-1]):
323
+ vals = _cell_vals(vdf, col_key)
324
+ v = round(float(vals.mean()), 3) if not vals.empty else None
325
+ z_row.append(v)
326
+ t_row.append(f"{v:.3f}" if v is not None else "—")
327
+ v_str = f"{v:.3f}" if v is not None else "N/A"
328
+ h_row.append(f"<b>{col_name}</b><br>{metric_label}: {v_str}<br>{config_lines}")
329
+ overall_vals = vdf[metric].dropna()
330
+ overall = round(float(overall_vals.mean()), 3) if not overall_vals.empty else None
331
+ z_row.append(overall)
332
+ t_row.append(f"{overall:.3f}" if overall is not None else "—")
333
+ ov_str = f"{overall:.3f}" if overall is not None else "N/A"
334
+ h_row.append(f"<b>Overall</b><br>{metric_label}: {ov_str}<br>{config_lines}")
335
+ rows_z.append(z_row)
336
+ rows_text.append(t_row)
337
+ rows_hover.append(h_row)
338
+ row_labels.append(label)
339
+
340
+ # Sort best→worst by Overall; heatmap y is bottom-up so reverse after sort
341
+ if len(rows_z) > 1:
342
+ sort_idx = sorted(range(len(rows_z)), key=lambda i: rows_z[i][-1] or -1, reverse=True)
343
+ rows_z = [rows_z[i] for i in sort_idx][::-1]
344
+ rows_text = [rows_text[i] for i in sort_idx][::-1]
345
+ rows_hover = [rows_hover[i] for i in sort_idx][::-1]
346
+ row_labels = [row_labels[i] for i in sort_idx][::-1]
347
+
348
+ zmin = -1 if metric == "mcc_test" else 0
349
+ n_data_cols = len(col_keys)
350
+ # Smaller font when many columns; angled labels in task mode for readability
351
+ cell_font_size = 9 if n_data_cols > 10 else 12
352
+ tick_angle = -45 if not aggregate_by_dataset else 0
353
+
354
+ fig = go.Figure(
355
+ go.Heatmap(
356
+ z=rows_z,
357
+ x=col_names,
358
+ y=row_labels,
359
+ text=rows_text,
360
+ texttemplate="%{text}",
361
+ textfont=dict(size=cell_font_size, color="black"),
362
+ hovertext=rows_hover,
363
+ hovertemplate="%{hovertext}<extra></extra>",
364
+ colorscale="RdYlGn",
365
+ zmin=zmin,
366
+ zmax=1,
367
+ colorbar=dict(title=metric_label, thickness=14, nticks=5, tickformat=".2f", lenmode="pixels", len=220),
368
+ hoverongaps=False,
369
+ )
370
+ )
371
+
372
+ # Highlight the Overall column with a border
373
+ fig.add_shape(
374
+ type="rect",
375
+ x0=n_data_cols - 0.5, x1=n_data_cols + 0.5,
376
+ y0=-0.5, y1=len(rows_z) - 0.5,
377
+ line=dict(color="#333", width=2),
378
+ fillcolor="rgba(0,0,0,0)",
379
+ )
380
+
381
+ n_rows = len(rows_z)
382
+ left_margin = min(max(len(lbl) for lbl in row_labels) * 7, 320)
383
+ # Extra bottom margin for angled tick labels in task mode
384
+ bottom_margin = 120 if not aggregate_by_dataset else 20
385
+ t_margin = 120 if not aggregate_by_dataset else 80
386
+ v_overhead = t_margin + bottom_margin
387
+ height = max(v_overhead + 80, n_rows * 48 + v_overhead + 20)
388
+ fig.update_layout(
389
+ title=f"{metric_label} — {model_alias}",
390
+ height=height,
391
+ template="plotly_white",
392
+ xaxis=dict(side="top", tickfont=dict(size=9 if not aggregate_by_dataset else 11), tickangle=tick_angle),
393
+ yaxis=dict(tickfont=dict(size=10), automargin=True),
394
+ margin=dict(l=left_margin, r=20, t=t_margin, b=bottom_margin),
395
+ )
396
+ return fig
397
+
398
+
399
+ # ---------------------------------------------------------------------------
400
+ # 4 & 5. Grouped bar chart (shared helper + two public wrappers)
401
+ # ---------------------------------------------------------------------------
402
+
403
+ def _grouped_bar(
404
+ df: pd.DataFrame,
405
+ metric: str,
406
+ x_items: list,
407
+ x_labels: list[str],
408
+ get_vals, # Callable[[model_df, item], pd.Series]
409
+ title: str,
410
+ ) -> go.Figure:
411
+ """Build a grouped bar chart shared by bar_plot_per_category and bar_plot_per_task."""
412
+ models = sorted(df["model_alias"].unique())
413
+ color_map = _model_color_map(models)
414
+ metric_label = CLASSIFICATION_METRICS.get(metric, metric)
415
+ fig = go.Figure()
416
+ for model in models:
417
+ model_df = df[df["model_alias"] == model]
418
+ y_vals = [
419
+ round(float(vals.mean()), 3) if not (vals := get_vals(model_df, item).dropna()).empty else None
420
+ for item in x_items
421
+ ]
422
+ fig.add_trace(go.Bar(name=model, x=x_labels, y=y_vals, marker_color=color_map[model]))
423
+ fig.update_layout(
424
+ title=title, yaxis_title=metric_label, barmode="group",
425
+ height=440, template="plotly_white", legend=_LEGEND_H,
426
+ )
427
+ return fig
428
+
429
+
430
+ def bar_plot_per_category(df: pd.DataFrame, metric: str) -> go.Figure:
431
+ """Grouped bar: X = task category, bars = models, height = mean metric."""
432
+ if df.empty or metric not in df.columns:
433
+ return go.Figure()
434
+ active_cats = df["task_category"].unique()
435
+ categories = [k for k in CATEGORY_DISPLAY_NAMES if k != "other" and k in active_cats]
436
+ metric_label = CLASSIFICATION_METRICS.get(metric, metric)
437
+ return _grouped_bar(
438
+ df, metric,
439
+ x_items=categories,
440
+ x_labels=[CATEGORY_DISPLAY_NAMES[c] for c in categories],
441
+ get_vals=lambda mdf, cat: mdf[mdf["task_category"] == cat][metric],
442
+ title=f"{metric_label} per task category",
443
+ )
444
+
445
+
446
+ def bar_plot_per_task(df: pd.DataFrame, metric: str, category: str) -> go.Figure:
447
+ """Grouped bar: X = individual tasks within category, bars = models."""
448
+ if df.empty or metric not in df.columns or not category:
449
+ return go.Figure()
450
+ task_df = df[df["task_category"] == category]
451
+ if task_df.empty:
452
+ fig = go.Figure()
453
+ fig.add_annotation(text=f"No data for category '{category}'", showarrow=False,
454
+ xref="paper", yref="paper", x=0.5, y=0.5)
455
+ return fig
456
+ tasks = sorted(task_df["task_name"].unique())
457
+ metric_label = CLASSIFICATION_METRICS.get(metric, metric)
458
+ cat_label = CATEGORY_DISPLAY_NAMES.get(category, category)
459
+ return _grouped_bar(
460
+ task_df, metric,
461
+ x_items=tasks,
462
+ x_labels=tasks,
463
+ get_vals=lambda mdf, task: mdf[mdf["task_name"] == task][metric],
464
+ title=f"{metric_label} per task — {cat_label}",
465
+ )
466
+
467
+
468
+ # ---------------------------------------------------------------------------
469
+ # 6. Speed vs Performance scatter (2-subplot)
470
+ # ---------------------------------------------------------------------------
471
+
472
+ def scatter_speed(df: pd.DataFrame, metric: str) -> go.Figure:
473
+ """Two side-by-side subplots: throughput vs metric, embedding time vs metric."""
474
+ if df.empty or metric not in df.columns:
475
+ return go.Figure()
476
+
477
+ has_throughput = "test_throughput_seq_s" in df.columns
478
+ has_embed_time = "test_embedding_time_s" in df.columns
479
+
480
+ if not has_throughput and not has_embed_time:
481
+ fig = go.Figure()
482
+ fig.add_annotation(text="Speed metrics not available in data", showarrow=False)
483
+ return fig
484
+
485
+ fig = make_subplots(
486
+ rows=1, cols=2,
487
+ subplot_titles=("Throughput (seq/s) vs Performance", "Embedding Time (s) vs Performance"),
488
+ )
489
+ models = sorted(df["model_alias"].unique())
490
+ color_map = _model_color_map(models)
491
+ metric_label = CLASSIFICATION_METRICS.get(metric, metric)
492
+
493
+ for model in models:
494
+ mdf = df[df["model_alias"] == model]
495
+ color = color_map[model]
496
+
497
+ if has_throughput:
498
+ sub = mdf[["test_throughput_seq_s", metric]].dropna()
499
+ fig.add_trace(
500
+ go.Scatter(x=sub["test_throughput_seq_s"], y=sub[metric], mode="markers",
501
+ name=model, marker=dict(color=color, size=10, opacity=0.75),
502
+ legendgroup=model, showlegend=True),
503
+ row=1, col=1,
504
+ )
505
+
506
+ if has_embed_time:
507
+ sub = mdf[["test_embedding_time_s", metric]].dropna()
508
+ fig.add_trace(
509
+ go.Scatter(x=sub["test_embedding_time_s"], y=sub[metric], mode="markers",
510
+ name=model, marker=dict(color=color, size=10, opacity=0.75),
511
+ legendgroup=model, showlegend=not has_throughput),
512
+ row=1, col=2,
513
+ )
514
+
515
+ fig.update_xaxes(title_text="Throughput (seq/s)", row=1, col=1)
516
+ fig.update_xaxes(title_text="Embedding Time (s)", row=1, col=2)
517
+ fig.update_yaxes(title_text=metric_label, row=1, col=1)
518
+ fig.update_yaxes(title_text=metric_label, row=1, col=2)
519
+ fig.update_layout(height=480, template="plotly_white")
520
+ return fig
521
+
522
+
523
+ # ---------------------------------------------------------------------------
524
+ # 7. Context length bubble chart
525
+ # ---------------------------------------------------------------------------
526
+
527
+ def bubble_context_length(df: pd.DataFrame, metric: str) -> go.Figure:
528
+ """Bubble chart: X=task_seq_length, Y=metric, color=model, size=max_context_size."""
529
+ if not TASK_SEQ_LENGTHS:
530
+ fig = go.Figure()
531
+ fig.add_annotation(
532
+ text=(
533
+ "Sequence length data not available.<br>"
534
+ "Run <b>python tools/compute_seq_lengths.py</b> to populate it."
535
+ ),
536
+ showarrow=False, font=dict(size=14), xref="paper", yref="paper", x=0.5, y=0.5,
537
+ )
538
+ fig.update_layout(height=480, template="plotly_white")
539
+ return fig
540
+
541
+ if df.empty or metric not in df.columns or "task_seq_length" not in df.columns:
542
+ return go.Figure()
543
+
544
+ _CONFIG_COLS = ["model_hf_id", "pooling_strategy", "layer_selection_params", "head_type"]
545
+ extra_cols = [c for c in _CONFIG_COLS if c in df.columns]
546
+ base_cols = ["model_alias", "task_name", "task_seq_length", metric, "max_context_size"]
547
+ plot_df = df[base_cols + extra_cols].dropna(subset=["task_seq_length", metric]).copy()
548
+ if "layer_selection_params" in plot_df.columns:
549
+ plot_df["_layer_info"] = plot_df["layer_selection_params"].apply(_fmt_layer)
550
+
551
+ def _config_str(row: pd.Series) -> str:
552
+ parts = []
553
+ for col, label in [("model_hf_id", "HF model"), ("pooling_strategy", "Pooling"),
554
+ ("_layer_info", "Layer"), ("head_type", "Head")]:
555
+ if col not in row.index:
556
+ continue
557
+ v = row[col]
558
+ if v is not None and not (isinstance(v, float) and pd.isna(v)):
559
+ parts.append(f"{label}: {v}")
560
+ return "<br>".join(parts)
561
+
562
+ plot_df["_config_str"] = plot_df.apply(_config_str, axis=1)
563
+
564
+ models = sorted(plot_df["model_alias"].unique())
565
+ color_map = _model_color_map(models)
566
+ max_ctx = plot_df["max_context_size"].max() or 1
567
+ metric_label = CLASSIFICATION_METRICS.get(metric, metric)
568
+ fig = go.Figure()
569
+
570
+ for model in models:
571
+ mdf = plot_df[plot_df["model_alias"] == model]
572
+ sizes = (mdf["max_context_size"].fillna(max_ctx) / max_ctx * 40 + 8).clip(8, 50)
573
+ cd = mdf[["max_context_size", "_config_str"]].to_numpy(dtype=object)
574
+ fig.add_trace(
575
+ go.Scatter(
576
+ x=mdf["task_seq_length"], y=mdf[metric], mode="markers", name=model,
577
+ marker=dict(size=sizes, color=color_map[model], opacity=0.7,
578
+ line=dict(width=1, color="white")),
579
+ text=mdf["task_name"],
580
+ hovertemplate=(
581
+ "<b>%{text}</b><br>Seq length: %{x} bp<br>"
582
+ f"{metric_label}: %{{y:.3f}}<br>"
583
+ "Max context: %{customdata[0]} bp<br>"
584
+ "%{customdata[1]}<extra></extra>"
585
+ ),
586
+ customdata=cd,
587
+ )
588
+ )
589
+
590
+ fig.update_layout(
591
+ title=f"{metric_label} vs task sequence length (bubble size = model max context)",
592
+ xaxis_title="Task sequence length (bp)",
593
+ yaxis_title=metric_label,
594
+ height=520,
595
+ template="plotly_white",
596
+ legend=_LEGEND_H,
597
+ )
598
+ return fig
src/populate.py DELETED
@@ -1,53 +0,0 @@
1
- import pandas as pd
2
- from datasets import load_dataset
3
-
4
- from src.display.utils import COLS, METRICS_COLS, BenchRawColumn, BenchTableModel, fields
5
- from src.display.utils import COLS_TO_AVERAGE, COLS_DEPEND_ON_MODEL, TASK_TYPE_MAP
6
-
7
-
8
- def get_leaderboard_df_from_hf_dataset(path: str) -> pd.DataFrame:
9
- dataset = load_dataset(path, split="train")
10
- df = pd.DataFrame(dataset)
11
- headers = {e.raw_name: e.name for e in fields(BenchRawColumn)}
12
- df = df.rename(columns=headers)
13
- cparams = BenchRawColumn.model_params.name
14
- df[cparams] = (df[cparams] / 1000000).round() # mandatory
15
- cinfer = BenchRawColumn.infer_time.name
16
- df[cinfer] = (df[cinfer] * 1000).round(2) # infer time per 1k samples
17
- BenchTableModel.infer_time.name = "Infer Time (s) - 1k samples"
18
- df = df[COLS]
19
- df = df.round(5)
20
- df[METRICS_COLS] = df[METRICS_COLS].round(2)
21
-
22
- return df
23
-
24
- def summarize_model_task_type_performance(df):
25
- """Summarizes model performance across task types"""
26
- df = df.copy()
27
- df['Task'] = df['Task'].map(TASK_TYPE_MAP)
28
- grouped = df.groupby(['Model', 'Task'])
29
-
30
- avg_aggs = {col: ['mean', 'std'] for col in COLS_TO_AVERAGE}
31
- model_aggs = {col: 'first' for col in COLS_DEPEND_ON_MODEL}
32
- agg_dict = {**avg_aggs, **model_aggs}
33
- agg_df = grouped.agg(agg_dict).reset_index()
34
-
35
- agg_df.columns = [
36
- f"{col[0]}_{col[1]}" if isinstance(col, tuple) and col[1] in ['mean', 'std']
37
- else col[0] if isinstance(col, tuple)
38
- else col
39
- for col in agg_df.columns
40
- ]
41
-
42
- for col in COLS_TO_AVERAGE:
43
- mean_col = f"{col}_mean"
44
- std_col = f"{col}_std"
45
-
46
- agg_df[col] = agg_df[mean_col].round(3).astype(str) + ' ± ' + agg_df[std_col].round(3).astype(str)
47
- agg_df.drop(columns=[mean_col, std_col], inplace=True)
48
-
49
-
50
- final_cols = ['Model', 'Task'] + COLS_TO_AVERAGE + COLS_DEPEND_ON_MODEL
51
- agg_df = agg_df[final_cols]
52
-
53
- return agg_df
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
src/submission/check_validity.py DELETED
@@ -1,109 +0,0 @@
1
- import json
2
- import os
3
- import re
4
- from collections import defaultdict
5
- from datetime import datetime, timedelta, timezone
6
-
7
- import huggingface_hub
8
- from huggingface_hub import ModelCard
9
- from huggingface_hub.hf_api import ModelInfo
10
- from transformers import AutoConfig
11
- from transformers.models.auto.tokenization_auto import AutoTokenizer
12
-
13
-
14
- def check_model_card(repo_id: str) -> tuple[bool, str]:
15
- """Checks if the model card and license exist and have been filled"""
16
- try:
17
- card = ModelCard.load(repo_id)
18
- except huggingface_hub.utils.EntryNotFoundError:
19
- return False, "Please add a model card to your model to explain how you trained/fine-tuned it."
20
-
21
- # Enforce license metadata
22
- if card.data.license is None:
23
- if not ("license_name" in card.data and "license_link" in card.data):
24
- return False, (
25
- "License not found. Please add a license to your model card using the `license` metadata or a"
26
- " `license_name`/`license_link` pair."
27
- )
28
-
29
- # Enforce card content
30
- if len(card.text) < 200:
31
- return False, "Please add a description to your model card, it is too short."
32
-
33
- return True, ""
34
-
35
-
36
- def is_model_on_hub(
37
- model_name: str, revision: str, token: str = None, trust_remote_code=False, test_tokenizer=False
38
- ) -> tuple[bool, str]:
39
- """Checks if the model model_name is on the hub, and whether it (and its tokenizer) can be loaded with AutoClasses."""
40
- try:
41
- config = AutoConfig.from_pretrained(
42
- model_name, revision=revision, trust_remote_code=trust_remote_code, token=token
43
- )
44
- if test_tokenizer:
45
- try:
46
- tk = AutoTokenizer.from_pretrained(
47
- model_name, revision=revision, trust_remote_code=trust_remote_code, token=token
48
- )
49
- except ValueError as e:
50
- return (False, f"uses a tokenizer which is not in a transformers release: {e}", None)
51
- except Exception as e:
52
- return (
53
- False,
54
- "'s tokenizer cannot be loaded. Is your tokenizer class in a stable transformers release, and correctly configured?",
55
- None,
56
- )
57
- return True, None, config
58
-
59
- except ValueError:
60
- return (
61
- False,
62
- "needs to be launched with `trust_remote_code=True`. For safety reason, we do not allow these models to be automatically submitted to the leaderboard.",
63
- None,
64
- )
65
-
66
- except Exception as e:
67
- return False, "was not found on hub!", None
68
-
69
-
70
- def get_model_size(model_info: ModelInfo, precision: str):
71
- """Gets the model size from the configuration, or the model name if the configuration does not contain the information."""
72
- try:
73
- model_size = round(model_info.safetensors["total"] / 1e9, 3)
74
- except (AttributeError, TypeError):
75
- return 0 # Unknown model sizes are indicated as 0, see NUMERIC_INTERVALS in app.py
76
-
77
- size_factor = 8 if (precision == "GPTQ" or "gptq" in model_info.modelId.lower()) else 1
78
- model_size = size_factor * model_size
79
- return model_size
80
-
81
-
82
- def get_model_arch(model_info: ModelInfo):
83
- """Gets the model architecture from the configuration"""
84
- return model_info.config.get("architectures", "Unknown")
85
-
86
-
87
- def already_submitted_models(requested_models_dir: str) -> set[str]:
88
- """Gather a list of already submitted models to avoid duplicates"""
89
- depth = 1
90
- file_names = []
91
- users_to_submission_dates = defaultdict(list)
92
-
93
- for root, _, files in os.walk(requested_models_dir):
94
- current_depth = root.count(os.sep) - requested_models_dir.count(os.sep)
95
- if current_depth == depth:
96
- for file in files:
97
- if not file.endswith(".json"):
98
- continue
99
- with open(os.path.join(root, file), "r") as f:
100
- info = json.load(f)
101
- file_names.append(f"{info['model']}_{info['revision']}_{info['precision']}")
102
-
103
- # Select organisation
104
- if info["model"].count("/") == 0 or "submitted_time" not in info:
105
- continue
106
- organisation, _ = info["model"].split("/")
107
- users_to_submission_dates[organisation].append(info["submitted_time"])
108
-
109
- return set(file_names), users_to_submission_dates
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
src/submission/submit.py DELETED
@@ -1,117 +0,0 @@
1
- import json
2
- import os
3
- from datetime import datetime, timezone
4
-
5
- from src.display.formatting import styled_error, styled_message, styled_warning
6
- from src.envs import API, EVAL_REQUESTS_PATH, QUEUE_REPO, TOKEN
7
- from src.submission.check_validity import already_submitted_models, check_model_card, get_model_size, is_model_on_hub
8
-
9
- REQUESTED_MODELS = None
10
- USERS_TO_SUBMISSION_DATES = None
11
-
12
-
13
- def add_new_eval(
14
- model: str,
15
- base_model: str,
16
- revision: str,
17
- precision: str,
18
- weight_type: str,
19
- model_type: str,
20
- ):
21
- global REQUESTED_MODELS
22
- global USERS_TO_SUBMISSION_DATES
23
- if not REQUESTED_MODELS:
24
- REQUESTED_MODELS, USERS_TO_SUBMISSION_DATES = already_submitted_models(EVAL_REQUESTS_PATH)
25
-
26
- user_name = ""
27
- model_path = model
28
- if "/" in model:
29
- user_name = model.split("/")[0]
30
- model_path = model.split("/")[1]
31
-
32
- precision = precision.split(" ")[0]
33
- current_time = datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ")
34
-
35
- if model_type is None or model_type == "":
36
- return styled_error("Please select a model type.")
37
-
38
- # Does the model actually exist?
39
- if revision == "":
40
- revision = "main"
41
-
42
- # Is the model on the hub?
43
- if weight_type in ["Delta", "Adapter"]:
44
- base_model_on_hub, error, _ = is_model_on_hub(
45
- model_name=base_model, revision=revision, token=TOKEN, test_tokenizer=True
46
- )
47
- if not base_model_on_hub:
48
- return styled_error(f'Base model "{base_model}" {error}')
49
-
50
- if not weight_type == "Adapter":
51
- model_on_hub, error, _ = is_model_on_hub(model_name=model, revision=revision, token=TOKEN, test_tokenizer=True)
52
- if not model_on_hub:
53
- return styled_error(f'Model "{model}" {error}')
54
-
55
- # Is the model info correctly filled?
56
- try:
57
- model_info = API.model_info(repo_id=model, revision=revision)
58
- except Exception:
59
- return styled_error("Could not get your model information. Please fill it up properly.")
60
-
61
- model_size = get_model_size(model_info=model_info, precision=precision)
62
-
63
- # Were the model card and license filled?
64
- try:
65
- license = model_info.cardData["license"]
66
- except Exception:
67
- return styled_error("Please select a license for your model")
68
-
69
- modelcard_OK, error_msg = check_model_card(model)
70
- if not modelcard_OK:
71
- return styled_error(error_msg)
72
-
73
- # Seems good, creating the eval
74
- print("Adding new eval")
75
-
76
- eval_entry = {
77
- "model": model,
78
- "base_model": base_model,
79
- "revision": revision,
80
- "precision": precision,
81
- "weight_type": weight_type,
82
- "status": "PENDING",
83
- "submitted_time": current_time,
84
- "model_type": model_type,
85
- "likes": model_info.likes,
86
- "params": model_size,
87
- "license": license,
88
- "private": False,
89
- }
90
-
91
- # Check for duplicate submission
92
- if f"{model}_{revision}_{precision}" in REQUESTED_MODELS:
93
- return styled_warning("This model has been already submitted.")
94
-
95
- print("Creating eval file")
96
- OUT_DIR = f"{EVAL_REQUESTS_PATH}/{user_name}"
97
- os.makedirs(OUT_DIR, exist_ok=True)
98
- out_path = f"{OUT_DIR}/{model_path}_eval_request_False_{precision}_{weight_type}.json"
99
-
100
- with open(out_path, "w") as f:
101
- f.write(json.dumps(eval_entry))
102
-
103
- print("Uploading eval file")
104
- API.upload_file(
105
- path_or_fileobj=out_path,
106
- path_in_repo=out_path.split("eval-queue/")[1],
107
- repo_id=QUEUE_REPO,
108
- repo_type="dataset",
109
- commit_message=f"Add {model} to eval queue",
110
- )
111
-
112
- # Remove the local file
113
- os.remove(out_path)
114
-
115
- return styled_message(
116
- "Your request has been submitted to the evaluation queue!\nPlease wait for up to an hour for the model to show in the PENDING list."
117
- )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
tools/compute_seq_lengths.py ADDED
@@ -0,0 +1,91 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """One-time script to compute mean sequence lengths (bp) per genomics benchmark task.
2
+
3
+ Loads each task from the InstaDeepAI/nucleotide_transformer_downstream_tasks HF dataset,
4
+ samples up to SAMPLE_SIZE sequences, measures their lengths, and writes the result to
5
+ tools/task_seq_lengths.json.
6
+
7
+ Usage:
8
+ python tools/compute_seq_lengths.py
9
+
10
+ After running, commit tools/task_seq_lengths.json to the repository.
11
+ The bubble chart in the leaderboard becomes functional once this file is populated.
12
+ """
13
+
14
+ import json
15
+ import os
16
+ import pathlib
17
+ from typing import Any
18
+
19
+ from datasets import DatasetDict, load_dataset
20
+ from dotenv import load_dotenv
21
+
22
+ load_dotenv()
23
+
24
+ HF_TOKEN = os.environ.get("HF_TOKEN")
25
+ HF_DATASET = "InstaDeepAI/nucleotide_transformer_downstream_tasks_revised"
26
+ SAMPLE_SIZE = 100
27
+
28
+ TASKS = [
29
+ "H2AFZ", "H3K27ac", "H3K27me3", "H3K36me3",
30
+ "H3K4me1", "H3K4me2", "H3K4me3", "H3K9ac", "H3K9me3", "H4K20me1",
31
+ "promoter_all", "promoter_tata", "promoter_no_tata",
32
+ "enhancers", "enhancers_types",
33
+ "splice_sites_all", "splice_sites_donors", "splice_sites_acceptors",
34
+ "variant_effect_pathogenic_clinvar", "variant_effect_pathogenic_omim",
35
+ "bulk_rna_expression",
36
+ ]
37
+
38
+ OUTPUT_PATH = pathlib.Path(__file__).parent / "task_seq_lengths.json"
39
+
40
+
41
+ def load_dataset_for_task(hf_dataset: str, task_name: str) -> DatasetDict:
42
+ """Load HuggingFace dataset for specific task.
43
+ Directly from the bdna package
44
+ """
45
+ kwargs: dict[str, Any] = {}
46
+ if HF_TOKEN:
47
+ kwargs["token"] = HF_TOKEN
48
+
49
+ if hf_dataset == "InstaDeepAI/genomics-long-range-benchmark":
50
+ return load_dataset(hf_dataset, task_name=task_name, **kwargs)
51
+
52
+ return load_dataset(hf_dataset, data_dir=task_name, **kwargs)
53
+
54
+
55
+ def compute_mean_length(task_name: str) -> int | None:
56
+ try:
57
+ ds_dict = load_dataset_for_task(HF_DATASET, task_name)
58
+ ds = ds_dict["train"]
59
+ # Sample a subset to keep it fast
60
+ sample = ds.select(range(min(SAMPLE_SIZE, len(ds))))
61
+ seq_col = next(
62
+ (c for c in sample.column_names if "sequence" in c.lower() or "seq" in c.lower()),
63
+ None,
64
+ )
65
+ if seq_col is None:
66
+ print(f" [{task_name}] No sequence column found — skipping.")
67
+ return None
68
+ lengths = [len(row[seq_col]) for row in sample]
69
+ mean_len = int(sum(lengths) / len(lengths))
70
+ print(f" [{task_name}] mean length = {mean_len} bp (n={len(lengths)})")
71
+ return mean_len
72
+ except Exception as exc:
73
+ print(f" [{task_name}] ERROR: {exc}")
74
+ return None
75
+
76
+
77
+ def main():
78
+ print(f"Computing sequence lengths from {HF_DATASET} …\n")
79
+ result: dict[str, int] = {}
80
+ for task in TASKS:
81
+ length = compute_mean_length(task)
82
+ if length is not None:
83
+ result[task] = length
84
+
85
+ OUTPUT_PATH.write_text(json.dumps(result, indent=2))
86
+ print(f"\nWrote {len(result)} entries to {OUTPUT_PATH}")
87
+ print("Commit tools/task_seq_lengths.json to enable the Context Length bubble chart.")
88
+
89
+
90
+ if __name__ == "__main__":
91
+ main()
tools/task_seq_lengths.json ADDED
@@ -0,0 +1,20 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "H2AFZ": 1000,
3
+ "H3K27ac": 1000,
4
+ "H3K27me3": 1000,
5
+ "H3K36me3": 1000,
6
+ "H3K4me1": 1000,
7
+ "H3K4me2": 1000,
8
+ "H3K4me3": 1000,
9
+ "H3K9ac": 1000,
10
+ "H3K9me3": 1000,
11
+ "H4K20me1": 1000,
12
+ "promoter_all": 300,
13
+ "promoter_tata": 300,
14
+ "promoter_no_tata": 300,
15
+ "enhancers": 400,
16
+ "enhancers_types": 400,
17
+ "splice_sites_all": 600,
18
+ "splice_sites_donors": 600,
19
+ "splice_sites_acceptors": 600
20
+ }