MitoInteract / recovery /scripts /dataset_smoke.py
ethanolivertroy's picture
feat: build leakage-aware MitoInteract recovery pipeline
6019d52
Raw
History Blame Contribute Delete
3.94 kB
#!/usr/bin/env python3
"""Validate dataset target semantics and compute a deterministic mean baseline."""
from __future__ import annotations
import argparse
import json
import math
import random
import statistics
import urllib.parse
import urllib.request
from pathlib import Path
from mitointeract_recovery import micromolar_to_paffinity
def fetch_rows(dataset: str, rows: int) -> list[dict]:
if not 2 <= rows <= 100:
raise ValueError("Dataset Viewer smoke checks support 2-100 rows")
query = urllib.parse.urlencode(
{
"dataset": dataset,
"config": "default",
"split": "train",
"offset": 0,
"length": rows,
}
)
url = f"https://datasets-server.huggingface.co/rows?{query}"
with urllib.request.urlopen(url, timeout=60) as response:
payload = json.load(response)
return [entry["row"] for entry in payload["rows"]]
def pearson(xs: list[float], ys: list[float]) -> float | None:
if len(xs) != len(ys) or len(xs) < 2:
return None
xm, ym = statistics.mean(xs), statistics.mean(ys)
numerator = sum((x - xm) * (y - ym) for x, y in zip(xs, ys))
denominator = math.sqrt(
sum((x - xm) ** 2 for x in xs) * sum((y - ym) ** 2 for y in ys)
)
return numerator / denominator if denominator else None
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--dataset", default="jglaser/binding_affinity")
parser.add_argument("--rows", type=int, default=100)
parser.add_argument("--seed", type=int, default=42)
parser.add_argument(
"--output", type=Path, default=Path("artifacts/dataset-smoke.json")
)
args = parser.parse_args()
rows = fetch_rows(args.dataset, args.rows)
errors = []
for index, row in enumerate(rows):
calculated = micromolar_to_paffinity(float(row["affinity_uM"]))
published = float(row["neg_log10_affinity_M"])
delta = abs(calculated - published)
if delta > 1e-5:
errors.append(
{
"row": index,
"calculated": calculated,
"published": published,
"delta": delta,
}
)
shuffled = list(rows)
random.Random(args.seed).shuffle(shuffled)
cut = max(1, int(len(shuffled) * 0.8))
train, test = shuffled[:cut], shuffled[cut:]
train_targets = [float(row["neg_log10_affinity_M"]) for row in train]
test_targets = [float(row["neg_log10_affinity_M"]) for row in test]
mean_prediction = statistics.mean(train_targets)
predictions = [mean_prediction] * len(test_targets)
rmse = math.sqrt(
statistics.mean(
(pred - target) ** 2 for pred, target in zip(predictions, test_targets)
)
)
mae = statistics.mean(
abs(pred - target) for pred, target in zip(predictions, test_targets)
)
report = {
"dataset": args.dataset,
"rows": len(rows),
"seed": args.seed,
"target_column": "neg_log10_affinity_M",
"target_unit": "pAffinity (-log10 of the dataset's mixed affinity value in M)",
"unit_check": {"passed": not errors, "errors": errors},
"split": {
"train_rows": len(train),
"test_rows": len(test),
"kind": "seeded smoke only",
},
"train_mean_baseline": {
"prediction": mean_prediction,
"rmse": rmse,
"mae": mae,
"pearson_r": pearson(predictions, test_targets),
},
"warning": "This first-row smoke sample is not a leakage-aware scientific evaluation.",
}
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(json.dumps(report, indent=2) + "\n")
print(json.dumps(report, indent=2))
if errors:
raise SystemExit(1)
if __name__ == "__main__":
main()