Spaces:
Sleeping
Sleeping
Major refactor - V2 (#7)
Browse files- refactored the whole webapp (a486179cf332b50a0585aff1550b344815676d91)
- added raw input inspector (a011d958b616b93695e594aefaa0f9ae6bc7e336)
- Makefile +5 -2
- README.md +24 -26
- app.py +390 -139
- docs/context-specs-plan.md +245 -0
- pyproject.toml +13 -0
- requirements.txt +0 -9
- src/about.py +0 -57
- src/constants.py +53 -0
- src/data.py +103 -0
- src/display/css_html_js.py +0 -64
- src/display/formatting.py +0 -27
- src/display/plotting.py +0 -205
- src/display/utils.py +0 -65
- src/envs.py +0 -29
- src/plots.py +598 -0
- src/populate.py +0 -53
- src/submission/check_validity.py +0 -109
- src/submission/submit.py +0 -117
- tools/compute_seq_lengths.py +91 -0
- tools/task_seq_lengths.json +20 -0
Makefile
CHANGED
|
@@ -1,8 +1,11 @@
|
|
| 1 |
-
.PHONY: style
|
| 2 |
|
| 3 |
style:
|
| 4 |
python -m black --line-length 119 .
|
| 5 |
ruff check --fix .
|
| 6 |
|
| 7 |
-
|
| 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 |
-
|
| 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 |
-
|
| 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 |
-
|
| 3 |
-
from gradio_leaderboard import ColumnFilter, Leaderboard, SelectColumns
|
| 4 |
-
from huggingface_hub import snapshot_download
|
| 5 |
|
| 6 |
-
from src.
|
| 7 |
-
from src.
|
| 8 |
-
from src.
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 12 |
|
| 13 |
-
|
| 14 |
-
|
|
|
|
| 15 |
|
|
|
|
| 16 |
try:
|
| 17 |
-
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
|
| 24 |
-
token=TOKEN,
|
| 25 |
-
)
|
| 26 |
|
| 27 |
-
|
| 28 |
-
|
| 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 |
-
|
| 80 |
-
|
| 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 |
-
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
|
| 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 |
-
|
| 125 |
-
|
| 126 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 127 |
|
| 128 |
-
|
| 129 |
-
|
| 130 |
|
| 131 |
-
|
| 132 |
-
|
| 133 |
-
|
| 134 |
-
|
|
|
|
|
|
|
| 135 |
|
| 136 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 137 |
|
| 138 |
-
|
| 139 |
-
|
| 140 |
-
|
| 141 |
|
| 142 |
-
|
| 143 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 144 |
|
| 145 |
-
|
|
|
|
|
|
|
| 146 |
|
| 147 |
-
|
| 148 |
-
|
| 149 |
-
|
| 150 |
|
| 151 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 152 |
|
| 153 |
-
|
| 154 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 155 |
|
| 156 |
-
|
| 157 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 158 |
|
| 159 |
-
scheduler = BackgroundScheduler()
|
| 160 |
-
scheduler.add_job(restart_space, "interval", seconds=1800)
|
| 161 |
-
scheduler.start()
|
| 162 |
|
| 163 |
-
|
|
|
|
|
|
| 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 |
+
}
|