MitoInteract / recovery /tests /test_bindingdb_benchmark.py
Ethan Troy
feat: add source-aware exact-Kd model evaluation with honest no-publish verdict
f9c6388
Raw
History Blame Contribute Delete
12 kB
import hashlib
import json
import math
from copy import deepcopy
from pathlib import Path
import pytest
from mitointeract_recovery.bindingdb_benchmark import (
MANIFEST_FILENAMES,
aggregate_strata,
build_manifests,
prepare_benchmark,
validate_gold_record,
)
from mitointeract_recovery.chemistry import canonicalize_smiles, stable_id
RAW_SMILES = (
"c1ccccc1",
"c1ccncc1",
"C1CCCCC1",
"C1CCCC1",
"c1ccoc1",
"c1ccsc1",
)
SEQUENCES = tuple("ACDEFGHIKLMNPQRSTVWY" + "A" * index for index in range(6))
YEARS = (2010, 2011, 2012, 2013, 2014, 2015, 2017, 2017, 2017, 2021, 2023, 2023)
def make_record(
index: int,
*,
kd_nm: float | None = None,
assay_id: str | None = None,
doi: str | None = None,
pmid: str | None = None,
publication_date: str | None = None,
source_record_id: str | None = None,
main_row_number: int | None = None,
) -> dict:
sequence = SEQUENCES[(index // 2) % len(SEQUENCES)]
smiles = canonicalize_smiles(RAW_SMILES[index % len(RAW_SMILES)])
protein_id = stable_id("protein", sequence)
ligand_id = stable_id("ligand", smiles)
pair_id = stable_id("pair", f"{sequence}\0{smiles}")
value = float(kd_nm if kd_nm is not None else 10 + index)
year = YEARS[index % len(YEARS)]
assay_value = assay_id or f"{100 + index}_1"
return {
"schema_version": "bindingdb-source-envelope/v1",
"source_database": "BindingDB",
"source_release": "202607",
"source_record_id": source_record_id or f"rs-{index}",
"reactant_set_id": f"reactant-{index}",
"main_row_number": main_row_number or index + 1,
"protein_id": protein_id,
"ligand_id": ligand_id,
"pair_id": pair_id,
"sequence": sequence,
"smiles": smiles,
"measurement_type": "Kd",
"relation": "=",
"kd_nm": value,
"pkd": -math.log10(value * 1e-9),
"assay": {
"entryid_assayid": assay_value,
"entry_id": assay_value.split("_")[0],
"assay_id": assay_value.split("_")[-1],
"assay_name": f"assay-{assay_value}",
"assay_description": f"description-{assay_value}",
"joined": True,
},
"citation": {
"article_doi": doi if doi is not None else f"10.1000/{index}",
"pmid": pmid if pmid is not None else str(1000 + index),
"publication_date": publication_date or f"1/2/{year}",
},
}
def write_inputs(tmp_path: Path, records: list[dict], *, release: str = "202607"):
gold = tmp_path / "gold_exact_kd.jsonl"
gold.write_text("".join(json.dumps(row, sort_keys=True) + "\n" for row in records))
audit = {
"source_database": "BindingDB",
"source_release": release,
"outputs": {
"gold_exact_kd.jsonl": {
"sha256": hashlib.sha256(gold.read_bytes()).hexdigest(),
"bytes": gold.stat().st_size,
"records": len(records),
}
},
}
audit_path = tmp_path / "bindingdb_audit.json"
audit_path.write_text(json.dumps(audit, sort_keys=True) + "\n")
return gold, audit_path
def normalized(records: list[dict]) -> list[dict]:
return [
validate_gold_record(record, line_number=index)
for index, record in enumerate(records, 1)
]
def read_jsonl(path: Path) -> list[dict]:
return [json.loads(line) for line in path.read_text().splitlines() if line]
def test_prepare_benchmark_aggregates_only_identical_strata_and_audits(tmp_path):
records = [make_record(index) for index in range(12)]
replicate = make_record(
0,
kd_nm=30,
source_record_id="rs-0-replicate",
main_row_number=99,
)
records.append(replicate)
gold, source_audit = write_inputs(tmp_path, records)
output = tmp_path / "output"
report = prepare_benchmark(gold, source_audit, output, seed=42)
observations = read_jsonl(output / "sample.jsonl")
assert report["counts"]["source_records"] == 13
assert report["counts"]["observations"] == 12
assert report["counts"]["unique_pairs"] == 12
assert sum(row["replicate_count"] for row in observations) == 13
assert [row["observation_id"] for row in observations] == sorted(
row["observation_id"] for row in observations
)
aggregated = next(row for row in observations if row["replicate_count"] == 2)
assert aggregated["kd_nm"] == 20
assert aggregated["pkd"] == pytest.approx(-math.log10(20e-9))
assert aggregated["replicate_kd_nm_min"] == 10
assert aggregated["replicate_kd_nm_max"] == 30
assert aggregated["replicate_kd_nm_iqr"] == 10
assert aggregated["source_record_ids"] == ["rs-0", "rs-0-replicate"]
assert aggregated["source_row_numbers"] == [1, 99]
assert aggregated["reactant_set_ids"] == ["reactant-0"]
assert aggregated["assay"]["assay_description"].startswith("description-")
assert aggregated["citation"]["publication_date"] == "1/2/2010"
singleton = next(row for row in observations if row["replicate_count"] == 1)
assert singleton["replicate_kd_nm_iqr"] == 0
assert report["pairs_spanning_multiple_years"]["count"] == 0
assert report["overlap_assertions"]["all_zero"] is True
for name, filename in MANIFEST_FILENAMES.items():
manifest = read_jsonl(output / filename)
assert len(manifest) == 12
assert len({row["pair_id"] for row in manifest}) == 12
assert {row["split"] for row in manifest} == {
"train",
"validation",
"test",
}
assert sum(report["manifests"][name]["pairs"].values()) == 12
assert sum(report["manifests"][name]["observations"].values()) == 12
def test_checksum_size_and_record_count_are_fail_closed(tmp_path):
records = [make_record(index) for index in range(12)]
gold, source_audit = write_inputs(tmp_path, records)
audit = json.loads(source_audit.read_text())
for field, value in (
("sha256", "0" * 64),
("bytes", gold.stat().st_size + 1),
("records", len(records) + 1),
):
changed = deepcopy(audit)
changed["outputs"]["gold_exact_kd.jsonl"][field] = value
source_audit.write_text(json.dumps(changed))
with pytest.raises(ValueError, match="mismatch"):
prepare_benchmark(gold, source_audit, tmp_path / field)
def test_source_release_mismatch_fails_before_preparation(tmp_path):
gold, source_audit = write_inputs(
tmp_path, [make_record(index) for index in range(12)], release="202606"
)
with pytest.raises(ValueError, match="source audit release"):
prepare_benchmark(gold, source_audit, tmp_path / "output")
@pytest.mark.parametrize(
("mutation", "message"),
[
(lambda row: row.update(measurement_type="Ki"), "measurement_type"),
(lambda row: row.update(relation="<"), "relation"),
(lambda row: row.update(kd_nm=0), "kd_nm"),
(lambda row: row.update(pkd=999), "inconsistent"),
(lambda row: row["assay"].update(joined=False), "assay join"),
(lambda row: row.update(source_release="202606"), "source_release"),
(lambda row: row.update(pair_id="pair-wrong"), "pair_id"),
(
lambda row: row["citation"].update(publication_date="not-a-date"),
"publication_date",
),
],
)
def test_malformed_gold_contract_is_rejected(mutation, message):
row = make_record(0)
mutation(row)
with pytest.raises(ValueError, match=message):
validate_gold_record(row, line_number=1)
def test_duplicate_source_record_ids_are_rejected(tmp_path):
records = [make_record(index) for index in range(12)]
records[1]["source_record_id"] = records[0]["source_record_id"]
gold, source_audit = write_inputs(tmp_path, records)
with pytest.raises(ValueError, match="duplicate source_record_id"):
prepare_benchmark(gold, source_audit, tmp_path / "output")
def test_assay_and_citation_boundaries_prevent_cross_aggregation():
base = make_record(0)
same = make_record(0, kd_nm=30, source_record_id="same", main_row_number=30)
other_assay = make_record(
0,
kd_nm=50,
assay_id="other_1",
source_record_id="assay",
main_row_number=31,
)
other_citation = make_record(
0,
kd_nm=70,
doi="10.2000/other",
pmid="9999",
source_record_id="citation",
main_row_number=32,
)
observations = aggregate_strata(
normalized([base, same, other_assay, other_citation])
)
assert len(observations) == 3
assert sorted(row["replicate_count"] for row in observations) == [1, 1, 2]
def test_missing_doi_and_pmid_records_remain_singleton_citation_strata():
first = make_record(0, source_record_id="missing-a", main_row_number=40)
second = make_record(0, kd_nm=30, source_record_id="missing-b", main_row_number=41)
for row in (first, second):
row["citation"]["article_doi"] = None
row["citation"]["pmid"] = None
observations = aggregate_strata(normalized([first, second]))
assert len(observations) == 2
assert {row["replicate_count"] for row in observations} == {1}
assert all(row["citation"]["article_doi"] is None for row in observations)
assert all(row["citation"]["pmid"] is None for row in observations)
def test_observation_ids_and_provenance_are_input_order_invariant():
first = make_record(0)
second = make_record(0, kd_nm=30, source_record_id="z-source", main_row_number=90)
forward = aggregate_strata(normalized([first, second]))
reverse = aggregate_strata(normalized([second, first]))
assert forward == reverse
def test_publication_split_uses_pair_maximum_year_and_counts_multi_year_pairs():
records = [make_record(index) for index in range(12)]
later_same_pair = make_record(
0,
assay_id="later_1",
doi="10.3000/later",
pmid="3000",
publication_date="2/3/2023",
source_record_id="later-source",
main_row_number=100,
)
observations = aggregate_strata(normalized(records + [later_same_pair]))
manifests, rejected = build_manifests(observations, seed=42)
assert rejected == {}
assert manifests["publication_time"][records[0]["pair_id"]] == "test"
pair_years = {}
for row in observations:
pair_years.setdefault(row["pair_id"], set()).add(row["publication_year"])
assert sum(len(years) > 1 for years in pair_years.values()) == 1
def test_publication_year_gap_fails_closed():
records = [make_record(index) for index in range(12)]
records[9]["citation"]["publication_date"] = "1/1/2019"
observations = aggregate_strata(normalized(records))
with pytest.raises(ValueError, match="outside the declared buckets"):
build_manifests(observations, seed=42)
def test_empty_publication_bucket_fails_closed():
records = [make_record(index) for index in range(12)]
for row in records:
row["citation"]["publication_date"] = "1/1/2015"
observations = aggregate_strata(normalized(records))
with pytest.raises(ValueError, match="split 'validation' is empty"):
build_manifests(observations, seed=42)
def test_random_manifest_is_deterministic_and_grouped_manifests_are_disjoint():
observations = aggregate_strata(
normalized([make_record(index) for index in range(12)])
)
first, _ = build_manifests(observations, seed=42)
second, _ = build_manifests(list(reversed(observations)), seed=42)
assert first == second
for manifest_name, group_key in (
("cold_protein_exact", "protein_id"),
("cold_scaffold", "scaffold_id"),
):
seen = {}
for row in observations:
split = first[manifest_name][row["pair_id"]]
previous = seen.setdefault(row[group_key], split)
assert previous == split