oncodsl / scripts /multiseed_hpv.py
govindbalki's picture
Upload folder using huggingface_hub
0fff343 verified
Raw
History Blame Contribute Delete
11.3 kB
"""Multi-seed HNSC/HPV stability report.
Runs the same HNSC/HPV genetic-programming pipeline the Lab launches
across a small grid of seeds, then reports (a) per-seed held-out AUROC
+ permutation p + GSE65858 transfer AUROC + p + gene coverage, (b) a
gene-recurrence tally across the winners with each symbol tagged by
its HNSC reference set (p16 / cell_cycle), and (c) held-out and
transfer AUROC ranges across seeds.
Read-only orchestration; no engine / API / airgap change. `scripts/`
is biology-aware — the same layer that owns `run_h2.py` — so this
file is allowed to import `airgap.reveal` and the named-side
validation helpers.
Airgap discipline
-----------------
Only each seed's WINNER's own opaque IDs are ever revealed (via the
bounded ``airgap.reveal`` — the same call ``/evaluate`` uses). The
sealed map is never dumped and the GSE65858 gene list never crosses
back into the engine. GSE65858 lives on the named (reveal) side; the
engine only sees TCGA HNSC via the anonymised matrix.
Usage
-----
python -m scripts.multiseed_hpv
python -m scripts.multiseed_hpv --seeds 1 3 7
python -m scripts.multiseed_hpv --generations 30 --population 300
"""
from __future__ import annotations
import argparse
import sys
import time
from typing import Sequence
from airgap import reveal
from api.app import (
COHERENCE_DEFAULT_WEIGHT,
REFERENCE_SETS_BY_DATASET,
_prepare_lab_data,
)
from engine_v2 import run_v2_pipeline
from engine_v2.fitness import objective_from_spec as v2_objective_from_spec
from validate.transfer_gse65858 import transfer_score
DEFAULT_SEEDS = [1, 3, 7, 11, 13, 17, 23, 29]
# Same numbers the Lab writes into RunParamsModel via DEFAULT_PARAMS.
DEFAULT_GENERATIONS = 60
DEFAULT_POPULATION = 300
DEFAULT_GENES_PER_SET = 8
DEFAULT_LAMBDA = 0.005
DEFAULT_PERMUTATIONS = 200
DEFAULT_PREFILTER_N: int | None = None # Match the Lab's "off" default.
def _tag_reference(sym: str) -> str:
"""Return the HNSC reference-set name that contains ``sym``, or ``-``.
``p16`` wins over ``cell_cycle`` if a symbol appears in both.
"""
hnsc = REFERENCE_SETS_BY_DATASET.get("hnsc", {})
for name in ("p16", "cell_cycle"):
if sym in hnsc.get(name, ()):
return name
return "-"
def _fmt(v: float | None, digits: int = 3) -> str:
if v is None:
return "—"
try:
f = float(v)
except (TypeError, ValueError):
return "—"
return f"{f:.{digits}f}" if f == f and f not in (float("inf"), float("-inf")) else "—"
def _run_one(
M, y, clinical, extra_labels,
*,
seed: int,
generations: int,
population: int,
genes_per_set: int,
lambda_size: float,
permutations: int,
prefilter_n: int | None,
coherence_weight: float,
confounders: tuple[str, ...],
transfer_perms: int,
) -> dict:
"""Run the pipeline for one seed and score the winner on GSE65858.
Returns a summary dict."""
objective_v2 = v2_objective_from_spec({"target": "hpv", "metric": "auroc"})
t0 = time.time()
_log, result = run_v2_pipeline(
M, y,
objective=objective_v2,
seed=seed,
prefilter_n=prefilter_n,
population_size=population,
n_generations=generations,
n_permutations=permutations,
lambda_size=lambda_size,
max_genes_per_set=genes_per_set,
clinical=clinical,
extra_labels=extra_labels,
coherence_weight=coherence_weight,
confounders=confounders,
)
elapsed = time.time() - t0
winning = result.get("winning", {}) or {}
gene_ids: list[str] = list(winning.get("gene_ids") or [])
# Bounded reveal — same discipline as /evaluate. Only this seed's
# winner's opaque IDs are ever translated; the sealed map is
# never dumped.
symbols: list[str] = reveal(gene_ids) if gene_ids else []
symbols = [s for s in symbols if isinstance(s, str) and s]
if symbols:
try:
transfer = transfer_score(
symbols, n_permutations=transfer_perms, seed=seed,
)
except FileNotFoundError:
transfer = None
else:
transfer = None
return {
"seed": seed,
"elapsed_s": round(elapsed, 1),
"holdout": winning.get("holdout_score"),
"permutation_p": winning.get("permutation_p"),
"gene_ids": gene_ids,
"symbols": symbols,
"transfer": transfer,
}
def main(argv: Sequence[str] | None = None) -> int:
ap = argparse.ArgumentParser(
description=(
"Run HNSC/HPV across several seeds and report per-seed stability "
"+ gene recurrence. Read-only; no engine / API / airgap change."
),
)
ap.add_argument(
"--seeds", type=int, nargs="+", default=None,
help=f"Seeds to run (default {DEFAULT_SEEDS}).",
)
ap.add_argument("--generations", type=int, default=DEFAULT_GENERATIONS)
ap.add_argument("--population", type=int, default=DEFAULT_POPULATION)
ap.add_argument("--genes-per-set", type=int, default=DEFAULT_GENES_PER_SET)
ap.add_argument("--lambda-size", type=float, default=DEFAULT_LAMBDA)
ap.add_argument("--permutations", type=int, default=DEFAULT_PERMUTATIONS)
ap.add_argument(
"--prefilter-n", type=int, default=DEFAULT_PREFILTER_N,
help=(
"Prefilter top-N (default: off, matches the Lab). Pass an int to "
"narrow the initial pool."
),
)
ap.add_argument(
"--transfer-perms", type=int, default=500,
help="Permutation-null size for the GSE65858 transfer AUROC.",
)
ap.add_argument(
"--no-coherence", action="store_true",
help="Skip the coherence prior (default: on, matches the Lab).",
)
args = ap.parse_args(argv or sys.argv[1:])
seeds = list(args.seeds) if args.seeds else list(DEFAULT_SEEDS)
coherence_weight = 0.0 if args.no_coherence else COHERENCE_DEFAULT_WEIGHT
print("=" * 78)
print("Multi-seed HNSC/HPV stability report")
print("=" * 78)
print(f"Seeds : {seeds}")
print(f"Generations × pop : {args.generations} × {args.population}")
print(f"Prefilter top-N : {args.prefilter_n if args.prefilter_n else 'off (all genes)'}")
print(f"Coherence prior : {'ON' if coherence_weight > 0 else 'OFF'} (weight {coherence_weight})")
print(f"Permutations (GP) : {args.permutations}")
print(f"Permutations (xfer): {args.transfer_perms}")
print(
"Airgap : only each winner's own genes are revealed "
"(bounded reveal via airgap.reveal — same discipline as /evaluate)."
)
print(
" The sealed map is never dumped; the GSE65858 "
"gene list never crosses back into the engine."
)
print("=" * 78)
# Load HNSC HPV cohort ONCE and reuse across seeds — the anonymised
# M / y / clinical / extra_labels are seed-independent.
print("\nLoading HNSC HPV cohort (anonymised matrix + clinical) …")
M, y, clinical, extra_labels = _prepare_lab_data("hpv", "hnsc")
print(f" M: {M.shape[0]} patients × {M.shape[1]} opaque columns")
print(f" y: HPV+ {int((y == 1).sum())} · HPV− {int((y == 0).sum())}")
if clinical is not None:
cols = list(clinical.columns)
print(f" clinical columns: {cols}")
# Match the API worker's HNSC confounder set.
extras = [c for c in ("sex", "race") if c in cols]
confounders = ("stage", "age", *extras)
else:
confounders = ("stage", "age")
print(f" Effect confounders: {confounders}")
# ----- Per-seed runs -----------------------------------------------------
rows: list[dict] = []
for i, seed in enumerate(seeds, start=1):
print(f"\n[{i}/{len(seeds)}] seed={seed} — running …", flush=True)
row = _run_one(
M, y, clinical, extra_labels,
seed=seed,
generations=args.generations,
population=args.population,
genes_per_set=args.genes_per_set,
lambda_size=args.lambda_size,
permutations=args.permutations,
prefilter_n=args.prefilter_n,
coherence_weight=coherence_weight,
confounders=confounders,
transfer_perms=args.transfer_perms,
)
symbols_short = ", ".join(row["symbols"][:6])
if len(row["symbols"]) > 6:
symbols_short += f", … (+{len(row['symbols']) - 6})"
print(
f" held-out AUROC={_fmt(row['holdout'])} "
f"p={_fmt(row['permutation_p'])} "
f"elapsed={row['elapsed_s']:.1f}s "
f"winner=[{symbols_short}]"
)
rows.append(row)
# ----- Per-seed table ----------------------------------------------------
print("\n" + "-" * 78)
print("PER-SEED SUMMARY")
print("-" * 78)
header = (
f"{'seed':>5} {'held-out':>9} {'GP p':>6} "
f"{'xfer AUROC':>10} {'xfer p':>7} {'genes found':>12}"
)
print(header)
print("-" * len(header))
for r in rows:
t = r["transfer"] or {}
found = int(t.get("n_found", 0)) if t else 0
total_req = (
int(t.get("n_found", 0)) + int(t.get("n_missing", 0)) if t else 0
)
if total_req == 0:
total_req = len(r["symbols"])
print(
f"{r['seed']:>5} {_fmt(r['holdout']):>9} "
f"{_fmt(r['permutation_p']):>6} "
f"{_fmt(t.get('auroc') if t else None):>10} "
f"{_fmt(t.get('p') if t else None):>7} "
f"{found:>4} / {total_req:<5}"
)
# ----- Range summary -----------------------------------------------------
def _range(vs: list[float | None]) -> str:
ok = [v for v in vs if v is not None]
if not ok:
return "—"
return f"{min(ok):.3f}{max(ok):.3f}"
holdouts = [r["holdout"] for r in rows]
xfer_aurocs = [
(r["transfer"] or {}).get("auroc") if r["transfer"] else None
for r in rows
]
print("\n" + "-" * 78)
print("RANGES ACROSS SEEDS")
print("-" * 78)
print(f" held-out AUROC : {_range(holdouts)}")
print(f" transfer AUROC : {_range(xfer_aurocs)}")
# ----- Gene recurrence tally --------------------------------------------
tally: dict[str, int] = {}
for r in rows:
for sym in dict.fromkeys(r["symbols"]): # dedup within a winner
tally[sym] = tally.get(sym, 0) + 1
print("\n" + "-" * 78)
print("GENE RECURRENCE ACROSS WINNERS")
print(f" ({len(seeds)} seeds; symbols listed by seeds they appear in, desc)")
print("-" * 78)
if not tally:
print(" (no winners had revealed symbols — nothing to tally)")
else:
header = f" {'symbol':<14} {'seeds':>6} {'reference':<11}"
print(header)
for sym, count in sorted(
tally.items(),
key=lambda kv: (-kv[1], kv[0]),
):
tag = _tag_reference(sym)
print(f" {sym:<14} {count:>4} / {len(seeds):<3} {tag:<11}")
print("\n" + "=" * 78)
print("Done.")
return 0
if __name__ == "__main__":
sys.exit(main())