dna-benchmark / tools /compute_seq_lengths.py
github-actions[bot]
Claude Sonnet 5
Fix long-range task categorization and seq-length lookup bugs
6fa59b9
Raw History Blame Contribute Delete
3.53 kB
"""One-time script to compute mean sequence lengths (bp) per genomics benchmark task.
Loads each task from the InstaDeepAI/nucleotide_transformer_downstream_tasks HF dataset,
samples up to SAMPLE_SIZE sequences, measures their lengths, and writes the result to
tools/task_seq_lengths.json.
Usage:
python tools/compute_seq_lengths.py
After running, commit tools/task_seq_lengths.json to the repository.
The bubble chart in the leaderboard becomes functional once this file is populated.
"""
import json
import os
import pathlib
from typing import Any
from datasets import load_dataset
from dotenv import load_dotenv
load_dotenv()
HF_TOKEN = os.environ.get("HF_TOKEN")
HF_DATASET = "InstaDeepAI/nucleotide_transformer_downstream_tasks_revised"
SAMPLE_SIZE = 100
# Names match the HF dataset's own data_dir casing (e.g. "H2AFZ"), not the lowercased
# task_name used elsewhere in the dash — see the .lower() when writing results below.
TASKS = [
"H2AFZ", "H3K27ac", "H3K27me3", "H3K36me3",
"H3K4me1", "H3K4me2", "H3K4me3", "H3K9ac", "H3K9me3", "H4K20me1",
"promoter_all", "promoter_tata", "promoter_no_tata",
"enhancers", "enhancers_types",
"splice_sites_all", "splice_sites_donors", "splice_sites_acceptors",
]
# variant_effect_causal_eqtl, variant_effect_pathogenic_clinvar, and bulk_rna_expression are
# long-range tasks (100,000bp anchor-centered windows, see PR #35 in bench-dna). They don't
# exist in HF_DATASET so they can't be measured the same way as the short-range tasks above,
# and their raw window length (100,000bp for all of them) wouldn't be a meaningful x-axis value
# anyway. These are the anchor-centered pooling windows the Genomics LRB paper actually reads
# from that window -- the portion of the sequence that's relevant to the label.
LONG_RANGE_SEQ_LENGTHS: dict[str, int] = {
"variant_effect_causal_eqtl": 1536,
"variant_effect_pathogenic_clinvar": 1536,
"bulk_rna_expression": 383 + 256,
}
OUTPUT_PATH = pathlib.Path(__file__).parent / "task_seq_lengths.json"
def compute_mean_length(task_name: str) -> int | None:
try:
kwargs: dict[str, Any] = {}
if HF_TOKEN:
kwargs["token"] = HF_TOKEN
ds_dict = load_dataset(HF_DATASET, data_dir=task_name, **kwargs)
ds = ds_dict["train"]
# Sample a subset to keep it fast
sample = ds.select(range(min(SAMPLE_SIZE, len(ds))))
seq_col = next(
(c for c in sample.column_names if "sequence" in c.lower() or "seq" in c.lower()),
None,
)
if seq_col is None:
print(f" [{task_name}] No sequence column found — skipping.")
return None
lengths = [len(row[seq_col]) for row in sample]
mean_len = int(sum(lengths) / len(lengths))
print(f" [{task_name}] mean length = {mean_len} bp (n={len(lengths)})")
return mean_len
except Exception as exc:
print(f" [{task_name}] ERROR: {exc}")
return None
def main():
print(f"Computing sequence lengths from {HF_DATASET} …\n")
result: dict[str, int] = {}
for task in TASKS:
length = compute_mean_length(task)
if length is not None:
result[task.lower()] = length
result.update(LONG_RANGE_SEQ_LENGTHS)
OUTPUT_PATH.write_text(json.dumps(result, indent=2))
print(f"\nWrote {len(result)} entries to {OUTPUT_PATH}")
print("Commit tools/task_seq_lengths.json to enable the Context Length bubble chart.")
if __name__ == "__main__":
main()