diff --git a/README.md b/README.md index 52788303..bdefeba7 100644 --- a/README.md +++ b/README.md @@ -140,7 +140,7 @@ uv run python -m odyssey.inference.steering --run-dir --held-out-shard-dir uv run python -m odyssey.inference.counterfactual --run-dir --held-out-shard-dir ``` -Paired subject-clustered intervals come from `scripts/alerts_cis.py` and `scripts/intervention_cis.py` (`odyssey/inference/uncertainty.py`); `scripts/gbm_feature_ablation.py` refits the GBM with each feature group dropped; `scripts/long_history_compare.py` compares two backbones on identical rows split by truncation; `scripts/concept_atlas.py` renders what each concept promotes and suppresses. The `scripts/make_*_table.py` generators turn banked JSON into the paper's tables. `scripts/eval_run.sh` chains the full evaluation of one run. +Paired subject-clustered intervals come from `scripts/alerts_cis.py` and `scripts/intervention_cis.py` (`odyssey/inference/uncertainty.py`); `scripts/gbm_feature_ablation.py` refits the GBM with each feature group dropped; `scripts/long_history_compare.py` compares two backbones on identical rows split by truncation; `scripts/compare_runs_paired.py` gives the paired hazard-head AUROC delta between two runs' alert row dumps on identical landmark rows (refusing on a row-set mismatch); `scripts/concept_atlas.py` renders what each concept promotes and suppresses. The `scripts/make_*_table.py` generators turn banked JSON into the paper's tables. `scripts/eval_run.sh` chains the full evaluation of one run. Every run is registered in [`docs/experiments.md`](docs/experiments.md) (host, data, commit, purpose, outcome) and its aggregate outputs are banked under `research_journal/figure_data///`. Environment fingerprints and numeric canaries are written with every checkpoint; the landmark protocol is versioned and stamped on every row dump. diff --git a/docs/ml4h2026_rebuttal_plan.md b/docs/ml4h2026_rebuttal_plan.md index bbfd0c84..589a795f 100644 --- a/docs/ml4h2026_rebuttal_plan.md +++ b/docs/ml4h2026_rebuttal_plan.md @@ -38,6 +38,16 @@ Each package says what exists, what to build, who runs it, and where the result - Exists: `odyssey/inference/interventions.py` runs truth/flip/none/calibrated/zeroing but scores only `top1_accuracy` and `mean_task_loss`. `odyssey/inference/steering.py` already scores hazard heads and knows the landmark rows. - Build: add hazard-head outputs to `interventions.py` (per event, per horizon, at landmark positions, paired subject-clustered bootstrap of truth minus none and truth minus flip on AUROC and on mean hazard). Reuse the landmark row set and the bootstrap from `scripts/alerts_cis.py` so the rows match Tables 12 to 14. - Run: on the banked flagship checkpoints, MIMIC `full_run_v10` on VM1 and the eICU joint mixture on VM2, band 0.15, all held-out shards (full-data rule). +- Built 2026-09-25 (branch `rebuttal/hazard-override`): `--hazard-heads` adds a `hazard` block per mode ({event: {horizon: {auroc, mean_risk, n_at_risk, n_positive, n_censored}}}) read at the alert protocol's landmark rows (same mask and outcome rule as `alerts.py`, so n_at_risk matches Tables 12 to 14 on the hybrid flagships), and a `hazard_paired` block on the left-hand mode's entry (truth carries truth_minus_none and truth_minus_flip, flip carries flip_minus_none) with 95% subject-clustered paired bootstrap intervals on the AUROC difference and on the mean P(event within h) difference. Off by default; the existing top-1/loss JSON is unchanged. Command to run on each VM (one pass per mode, the hazard readout is free; the bootstrap runs on the landmark table only): + + ``` + python -m odyssey.inference.interventions --run-dir R --held-out-shard-dir D/held_out \ + --output-json R/interventions_band15_hazard.json --max-shards 37 --num-lanes 64 \ + --chunk-size 512 --uncertain-band 0.15 --modes none truth flip random \ + --hazard-heads --hazard-boot 1000 --checkpoint checkpoint_best.pt + ``` + + Set `--max-shards` to the split's full shard count on each VM (37 is MIMIC's). The transformer arm keeps the lever test's TBTT view (whole history, no context truncation), so its landmark rows are the model-free set, not the packed-context set `alerts.py` scores for that backbone. - Bank: `research_journal/figure_data/{vm1,vm2}//interventions_hazard.json`. New table generator `scripts/make_lever_hazard_table.py`. - Rebuttal use: if truth moves the 24 h death or vasopressor hazard the right way, the lever verdict changes and the paper's Q3 gets a real endpoint. If it does not, the negative result becomes like-for-like with the edit test, which is what the review asked for. Either way it answers W1. - GEMINI: same script, run by Amrit, only if time. @@ -50,6 +60,7 @@ Each package says what exists, what to build, who runs it, and where the result - Bank: `research_journal/figure_data/{vm1,vm2}/full_baseline_v13/`. Add rows to `docs/experiments.md` and a column to `make_comparator_tables.py`. - Rebuttal use: reports the bottleneck's cost as a paired AUROC difference per cell, which is the number Q4 promises. Also settles whether the TabICLv2 result means "we lack the panel" or "the sequence model is undertrained". - Rule: full-data numbers only. No subset numbers go in the response. +- Paired scoring (branch `rebuttal/integration`): `scripts/compare_runs_paired.py` takes the two runs' `alerts_rows.parquet` dumps, inner-joins them per event on the row key `(subject_id, visit_id, time_hours)` that `index_row_table` writes, and refuses (exit 2, both row counts and the unmatched count printed) when more than 0.1% of either dump's rows are unmatched, so a row-set mismatch cannot pass silently. On the joined rows it reports, per (event, horizon), both hazard-head AUROCs and the paired subject-clustered bootstrap of `baseline minus bottleneck` (`bootstrap_auroc_delta`, 1000 draws), the two GBM refits side by side without a bootstrap (their difference is refit variance, for scale), and with `--inference-a/--inference-b` the next-event set top-1, exact top-1 and cross-entropy from each `inference_results.json`. VM1 command, from `~/odyssey` after the baseline alert chain has written its dump: `uv run python scripts/compare_runs_paired.py --dump-a ~/runs/full_run_v10/alerts_rows.parquet --dump-b ~/runs/full_run_baseline_v10/alerts_rows.parquet --label-a bottleneck --label-b baseline --inference-a ~/runs/full_run_v10/inference_results.json --inference-b ~/runs/full_run_baseline_v10/inference_results.json --output-json ~/runs/full_run_baseline_v10/paired_vs_v10.json`. VM2: the same with `~/runs/eicu_full_v10` and `~/runs/eicu_full_baseline_v10`. The dump file names must be the ones the chain actually wrote (the flagships' dumps may carry a `_v4` suffix; `ls ~/runs//alerts_rows*.parquet` first). Copy the JSON to `research_journal/figure_data/{vm1,vm2}//paired_vs_v10.json`; the printed markdown table is the rebuttal table for Q4. ### WP3. Completeness and leakage probes (W2). Inference only. Days 2 to 4. @@ -58,20 +69,31 @@ Each package says what exists, what to build, who runs it, and where the result - Run: flagship MIMIC and eICU checkpoints, all held-out shards or a stated subsample if the dump is too large (say so in the table). - Bank: `research_journal/figure_data//channel_probes.json`. - Rebuttal use: the review's strongest technical point is that the poles, not k, carry the forecast. This measures it directly. If k alone recovers most of the accuracy, W2 is answered. If not, the Q2 sentence in the abstract must be reworded to "runs through the named embeddings", and we say so in the response rather than defend it. +- Script (branch `rebuttal/channel-probes`): `scripts/probe_channel.py` dumps, at every scored held-out position (a seeded, shard-stratified subsample capped by `--max-positions`, default 2,000,000), the concept probabilities k, the unnamed slot's probability u, the poles w+/w- and the unnamed embedding z_u, then fits fresh linear readouts on a split BY SUBJECT (70/10/20 train/tune/test) and scores exact next-event top-1 on the held-out subjects: `k_only`, `k_plus_u`, `z_named` (the named embeddings), `poles_mean_k` (z rebuilt with every k held at its training mean), `poles_both` (raw [w+, w-]), `h_bar` (the full bottleneck, the ceiling) and `model_head` (the model's own head, no refit). Each carries a subject-clustered bootstrap CI (1000 resamples) and its ratio to the model's own accuracy on the same rows ("retained", with a paired CI); the CTL leakage probes of `odyssey/inference/leakage.py` run on the same split from the same dump. Hazard heads at landmarks are not dumped: the landmark protocol lives in `alerts.py`, and a second implementation here would risk a number that disagrees with the alert tables. VM1 command, from `~/odyssey` after `git fetch && git reset --hard origin/main` (never `uv sync` on a VM): `setsid nohup uv run python scripts/probe_channel.py --run-dir ~/runs/full_run_v10 --held-out-shard-dir ~/data/mimiciv_3.1_v1/data/held_out --output-json ~/runs/full_run_v10/channel_probes.json --num-lanes 16 --chunk-size 512 --max-positions 2000000 --seed 0 > ~/runs/full_run_v10/channel_probes.log 2>&1 & disown`. VM2: the same with `~/runs/eicu_full_v10` and `~/data/eicu_2.0_v1/data/held_out`. Runtime guess: the dump is one streaming pass at 16 lanes, about 20 to 30 minutes on MIMIC's 37 shards (the Sep 1 intervention rerun took about 20 minutes per mode) and less on eICU's 17; then six readout fits of 5 epochs over about 1.4M rows against the full vocabulary (a few minutes each on the A100) and four CTL probes on a 500k-row cap. Budget one hour per flagship. Memory: the bank keeps both poles per concept, 2M x 29 x 32 x 2 x 2 bytes = 7.4 GB in fp16 on the GPU; pass `--bank-on-cpu` if training shares the card. Copy the JSON to `research_journal/figure_data//channel_probes.json`. Reading it: if `k_only` recovers most of `h_bar` (retained near 1), the forecast runs through the named STATES and W2 is answered; if `poles_mean_k` recovers most of it while `k_only` does not, it runs through the poles and the Q2 sentence must say "through the named embeddings". `k_plus_u` against `k_only` prices the unnamed slot's probability; `z_named` against `h_bar` prices the unnamed embedding. ### WP4. RandInt arm (W5). Training. Days 2 to 8. - Exists: `randint_prob` in `TrainingConfig` (default 0.25; flagships set 0.0). Steerling-style steering epochs `eicu_full_DEC_v13_steer2` and a control epoch `_ctrl` are finished, banked under `research_journal/figure_data/vm2/`, and recorded in `docs/experiments.md` (eval rows): neither epoch makes the override help (truth minus none, top-1 points: v13 -0.16, steer2 -0.18, ctrl -0.12), the extra epoch rather than the steering losses carries the accuracy gains, and no epoch strengthens the lever. The paper body says only that the retrofit "does not reliably strengthen the dials". -- Build: nothing for training. The steer2/ctrl result is already a banked intervention-aware arm and goes straight into the rebuttal block for W5. +- Also exists, at subset scale only (docs/experiments.md rows 36, 60, 65; journal entries 08 and 25): `subset_run_v4` (MIMIC 30 shards, RandInt 0.25, early architecture without hazard heads: lever correctly signed but +0.1 top-1 point, exact top-1 -8.8 vs no RandInt, retired), `subset_run_L2` (MIMIC 30 shards, RandInt 1.0: none 25.9, truth 25.9, flip 25.8, no separation) and `eicu_subset_indep_b` (eICU 30 shards, stage-B independent training with RandInt 1.0: truth minus flip +0.78 points, truth ties none, set top-1 collapsed to 47.1). So the "standard remedy" was applied three times and never produced a usable lever. Under the paper's rules the lever numbers are admissible (interventions are exempt from the full-data rule), the accuracy costs are not. +- Build: nothing for training. The steer2/ctrl result is already a banked intervention-aware arm and goes straight into the rebuttal block for W5, together with the three subset RandInt arms as development-scale history. The full-scale eICU RandInt run below turns that history into one full-data number; the subset runs tell us what to expect (correct sign, tiny magnitude, an accuracy cost). - Run: eICU joint mixture with `randint_prob: 0.25`, full scale, on VM2 after WP2 finishes there. MIMIC only if VM1 is free by day 5. Score with `interventions.py` (top-1/loss and the new hazard output from WP1) and the readout table. - Bank: `research_journal/figure_data/vm2/eicu_full_RI_v13/`. - Rebuttal use: the reviewer will say "CEM's RandInt fixes this and you turned it off". We answer with the RandInt arm's truth-vs-none numbers and with the already-banked steering epochs. If RandInt makes truth beat none, that is a positive finding and goes in the camera-ready as the predicted remedy confirmed. If not, the negative result is stronger. ### WP5. GEMINI items (W3, W11). Needs Amrit. Days 3 to 9, one session inside the secure environment. -- (a) Paired bootstrap against the current refit: `scripts/alerts_cis.py --dump --scorers hazard gbm` on the current GEMINI GBM refit. `scripts/gemini/run.sh` has no alerts_cis stage; the lead adds one before Amrit's session. -- (b) Panel coverage: add a one-line report to `odyssey/data/signal_panel.py` that prints how many of the 48 signals resolve per source, and run it for all three sources. The review's guess is that the blood-pressure panel does not map on GEMINI; we need the number either way. -- (c) Cohort counts: held-out subjects, admissions, hospitals, years, and the split policy, for the GEMINI arm. Only aggregate numbers leave the environment. +- (a) Paired bootstrap against the current refit: `scripts/alerts_cis.py --dump --scorers hazard gbm` on the current GEMINI GBM refit. Built as the `alerts-cis` step of `scripts/gemini/run.sh` (branch `rebuttal/gemini-stages`): it reads `~/runs/gemini_full_DEC_v12/alerts_rows_allshards.parquet`, the dump the `_allshards` alerts pass left on the node (GBM on all 894 train shards at 10% row thinning, all 112 held-out shards), prints its row and subject counts, runs 1000 paired subject-clustered draws, and exports `scripts/gemini/out/evals/gemini_full_DEC_v12_allshards_alerts_cis.json` (per event x horizon: hazard AUROC, GBM AUROC, hazard minus GBM with a 95% interval, n_at_risk, n_positive, n_subjects). The full dump is about 25M rows and the bootstrap takes 15 to 20 h; set `GEMINI_CIS_MAX_SUBJECTS` if the session cannot afford that, and the output records the subsample size. +- (b) Panel coverage: `scripts/panel_coverage.py` resolves the 48 panel signals against a source's LOINC table in `odyssey/data/code_mapping.py` and, given `metadata/codes.parquet`, says which resolved prefixes were actually charted. The in-repo table already answers the review's guess: GEMINI resolves 15 of 48, and the non-invasive blood-pressure panel is among the unresolved. The `panel-coverage` step confirms it on the node's code inventory and exports the names. For MIMIC-IV and eICU run `uv run python scripts/panel_coverage.py --source mimic_iv` (or `eicu`) anywhere. +- (c) Cohort counts: `scripts/cohort_counts.py` over the MEDS shards a run used (train and tuning from its `config.json`, held-out from the eval dir): subjects, admissions, hospitals from `metadata/hadm_id_hospital.parquet`, admission year range, length of stay, sex and age where the source charts them (GEMINI extracts neither, so those read "not available"), and the share of subjects with each hazard event under the alerts leg's own onset definitions (sepsis3 does not resolve on GEMINI and is listed under `events_dropped`). Every count under 10 leaves as `"<10"`. The `cohort-counts` step exports `scripts/gemini/out/evals/gemini_full_DEC_v12_cohort_counts.json`. The same script runs on the VMs for MIMIC-IV and eICU with `--run-dir --split held_out=`. +- Commands for Amrit's session, in this order, each inside tmux, each self-syncing and idempotent (re-running with the output present only re-exports it): + + ``` + tmux new -s wp5 'scripts/gemini/run.sh panel-coverage gemini_full_DEC_v12' + tmux new -s wp5 'scripts/gemini/run.sh cohort-counts gemini_full_DEC_v12' + tmux new -s wp5 'scripts/gemini/run.sh alerts-cis gemini_full_DEC_v12' + ``` + + `alerts-cis` goes last because it is the long one. If the printed dump size makes the full bootstrap unaffordable, use `GEMINI_CIS_MAX_SUBJECTS=200000 scripts/gemini/run.sh alerts-cis gemini_full_DEC_v12` instead, and the rebuttal cites the subsample size the JSON records. If the `_allshards` dump is missing, the step prints the exact `alerts` command that regenerates it. - (d) GEMINI GBM count-feature ablation (Table 15 on GEMINI) if the session has time. This is the one that would let us keep "the deficit belongs to those sites" or force us to drop it. - Rebuttal use: the abstract's "the deficit belongs to those sites, not the design" is the sentence most at risk. If (a) separates the 9 wins against the current refit and (b) shows the panel is not handicapped, the sentence stays. If (b) shows the panel lost blood pressure on GEMINI, we soften the sentence in the response and in the camera-ready to "reverses by point estimate on GEMINI; the panel there lacks N of 48 signals". @@ -112,6 +134,67 @@ Keep each block under 150 words. The response box on OpenReview is short. - Run `figures/pagecheck.py` and the awk comment check after every edit (the build has lost prose to `%` lines before). - The submitted source is `paper/ml4h/main_mixture.tex`; the old `main.tex` and its aux files were retired to `paper/ml4h/retired/` on 2026-09-25. `make_steering_table.py` and `make_specificity_table.py` now feed no table in the paper; keep them for the steering follow-up. +## Results so far (updated 2026-09-29, 01:00 UTC) + +Every number below is banked under `research_journal/figure_data/` in the run directory named, and was computed on all held-out shards unless stated. Runs launched 2026-09-25 from the `rebuttal/integration` branch; VM recipe and logs are in the session memory and the chain scripts under the VM home directories. + +### WP1, override scored on the hazard heads: done, negative on both databases + +`interventions_band15_hazard.json` under `vm1/full_run_v10/` and `vm2/eicu_full_v10/`. Band 0.15, modes none / truth / flip / random, hazards read at the landmark rows of Tables 12 and 13 (the no-intervention hazard AUROCs reproduce those tables to three decimals), paired subject-clustered bootstrap with 1,000 resamples. + +- eICU-CRD (655,415 landmark rows): truth minus none on hazard AUROC is negative and separated in all 12 cells (vasopressor -0.006 to -0.008, AKI -0.003 to -0.006, death -0.002 to -0.006, ICU -0.001 to -0.003). Mean predicted risk moves by at most 0.004 absolute. +- MIMIC-IV (1,214,849 landmark rows): truth minus none between -0.0002 and -0.006 across 15 cells, negative and separated in 9, never positive. Mean risk moves by at most 0.0015. +- Reading: the label override is inert on the clinical hazards as it is on next-event accuracy. The two halves of Q3 are now scored on the same endpoint. + +### WP2, the cost of the bottleneck: done + +No-bottleneck arms (`model_kind=baseline`, flagship recipe, RandInt 0) trained at full scale: `vm1/full_run_baseline_v10` (49,375 steps, best val 1.7985) and `vm2/eicu_full_baseline_v10` (30,250 steps, early stop, best val 1.5362). Because the submitted flagships were trained on 2026-08-31 code, the bottleneck recipe was also retrained on current code (`vm1/full_run_v10_re`, `vm2/eicu_full_v10_re`), so the like-for-like pair is retrain versus baseline. All arms scored with the standard chain (GBM refit on every train shard) and compared on identical landmark rows with `scripts/compare_runs_paired.py` (`paired_re_vs_baseline.json`, `paired_flagship_vs_re.json` under each `_re` run; `paired_vs_v10*.json` under each baseline run for the flagship pair). + +- Paired intervals (`paired_re_vs_baseline.json`): MIMIC-IV 13 of 15 cells separated in the no-bottleneck arm's favour, 2 ties (death 24 h and 72 h), none the other way; eICU-CRD 8 of 15 separated (AKI at all horizons, ICU 8 and 24 h, vasopressor at all horizons), 7 ties (death at all horizons, ICU 72 h, Sepsis-3 at all horizons), none the other way. Code drift on eICU-CRD (`paired_flagship_vs_re.json`): the retrain beats the rescored August flagship in 11 of 12 cells. The MIMIC-IV flagship-versus-retrain file was written on VM1 after the VM was listed and stays on its disk (9 retrain wins, 6 ties by its log line). +- Like-for-like, MIMIC-IV: the no-bottleneck arm is ahead in 15 of 15 cells by point estimate, +0.004 to +0.009 AUROC (death +0.004 at every horizon, vasopressor +0.005 / +0.006 / +0.009, AKI +0.006 / +0.007 / +0.004, ICU +0.006 / +0.007 / +0.009, Sepsis-3 +0.005 / +0.006 / +0.005). Code drift on MIMIC-IV (retrain minus flagship) is +0.000 to +0.009. +- Like-for-like, eICU-CRD: tied on death (-0.004 / -0.000 / +0.005), no-bottleneck ahead by +0.012 to +0.014 on AKI, +0.005 to +0.008 on vasopressor, +0.007 / +0.003 / -0.000 on ICU, mixed on Sepsis-3. Code drift on eICU-CRD is large: retrain minus flagship +0.033 / +0.034 / +0.034 on death and +0.095 / +0.085 / +0.089 on AKI (current AKI label), +0.003 to +0.009 on vasopressor, +0.004 on ICU. Current main trains a much better eICU model than the one in the paper. +- Against the panel: the MIMIC-IV retrain keeps the flagship's two wins (death 8 h, vasopressor 8 h) and turns vasopressor 24 h into a tie; the eICU-CRD retrain ties on death at 8 h (-0.006, interval crosses zero), wins Sepsis-3 at 8 and 24 h (a new event there; the panel has 469 positives), and loses the other cells by about half the flagship's margin on death (24 h: -0.043 versus -0.078). The no-bottleneck arms lose every original cell too, with the same two MIMIC-IV wins. +- Next-event, current code: MIMIC-IV exact top-1 37.8 (bottleneck) versus 37.8 (no bottleneck), set top-1 82.0 versus 83.3; eICU-CRD 55.2 versus 55.9, set top-1 88.4 versus 88.7. +- Reading: the bottleneck's cost is real but under one AUROC point in every cell, at most a fifth of the gap to the panel. The earlier flagship-based estimate (up to 0.04 on eICU death, 0.10 on AKI) was mostly code drift, not the bottleneck. +- Retrain levers (band 0.15, top-1 points): MIMIC-IV truth minus none -0.22, zero_known retains 5.3%; eICU-CRD -0.04, zero_known retains 39.2%. Same verdict as the paper on both. + +Two facts found on the way. First, the banked eICU Table 13 AKI cells are under an older AKI label: rescoring the same checkpoint with current code (`vm2/eicu_full_v10/alerts_rescore.json`) leaves death, vasopressor and ICU AUROCs identical and drops AKI from 0.741 to 0.649 at 8 h (at-risk rows 451,747 to 308,700; the GBM drops from 0.891 to 0.820). The paper must state which label Table 13 uses, and the camera-ready should carry the retrained eICU model. Second, the panel's death-at-8 h AUROC on MIMIC-IV moved from 0.944 to 0.896 to 0.909 across three refits of the same recipe (1,987 positives), larger than the 0.029 refit variance the paper reports; the other cells agree within 0.005. + +### WP3, which channel carries the forecast: done, the poles carry it + +`channel_probes.json` under both flagship run directories (banks `channel_probes.bank.pt`, about 8 GB each, stay on the VMs). Subject-split held-out sample of 2,000,000 positions; fresh linear readouts scored on strict next-event top-1 over about 390,000 (MIMIC) and 400,000 (eICU) test positions; 95% subject-clustered intervals within 0.005. + +| readout | MIMIC-IV | eICU-CRD | +|---|---|---| +| model's own head | 0.372 | 0.539 | +| concept probabilities k only | 0.261 | 0.565 | +| k plus the unnamed slot | 0.264 | 0.585 | +| named embeddings z | 0.584 | 0.790 | +| poles with k fixed at its mean | 0.583 | 0.790 | +| poles alone (raw w+, w-) | 0.569 | 0.786 | +| full bottleneck | 0.584 | 0.788 | + +- The poles with the probabilities held constant recover everything the bottleneck carries. The probabilities alone recover 45% of it on MIMIC-IV and 72% on eICU-CRD. The forecast runs through the named embeddings; the reviewer's W2 point stands and Q2 must be reworded. +- Fresh readouts beat the model's own head because they optimise strict top-1 while the model trains bundle-invariant; compare readouts with the full-bottleneck readout, not with the model head. +- CTL leakage probes on the next-token code family (9 classes): probabilities only 0.935 / 0.967, embeddings only 0.956 / 0.973, unnamed slot only 0.878 / 0.953, random-projected probabilities 0.936 / 0.968 (MIMIC / eICU). + +### WP4, RandInt: done, no lever at full scale either + +`vm2/eicu_full_RI_v10` (flagship recipe, `randint_prob 0.25`, current code, 2 epochs): truth minus none -0.26 top-1 points (loss +0.006 nats), truth minus flip -0.44 points; zero_known retains 54.5%. No accuracy cost: set top-1 88.4 versus 88.1, exact 53.6 versus 54.0, readout mean 0.873 versus 0.872 over the 25 shared concepts. Against the panel it wins the three Sepsis-3 cells and loses the twelve original ones. With the three subset arms and the two steering epochs, the intervention-aware remedy has now been applied five times at two scales without producing a usable override. + +### WP5, GEMINI: code ready, the node session is Amrit's + +`scripts/panel_coverage.py` on the in-repo mapping table: GEMINI resolves 15 of the 48 panel signals; the non-invasive systolic, diastolic and mean pressures are among the unresolved (only arterial systolic maps). The three run.sh steps (`panel-coverage`, `cohort-counts`, `alerts-cis`) are built and tested; commands are listed under WP5. + +### WP6, free items: done + +- Parameters: 32,233,101 (MIMIC-IV flagship) and 20,384,390 (eICU-CRD), the difference being the per-source next-event vocabulary head; `research_journal/figure_data/param_counts.json`. `paper/ml4h/tables/hparams.tex` generated by `scripts/make_hparams_table.py` with the GBM grid, estimator and panel sizes read from the code. +- Cohort counts (`cohort_counts.json` under both flagship run directories): MIMIC-IV 291,702 / 36,463 / 36,462 subjects (train / tuning / held-out), 435,803 / 54,898 / 55,328 admissions, median stay 2.8 days, 53% female; per-subject prevalence vasopressor 5.0%, ICU admission 17.9%, AKI 18.2%, death 10.5%, Sepsis-3 12.6%, 30-day readmission 13.8%. eICU-CRD 133,084 / 16,636 / 16,635 subjects, 160,643 / 20,094 / 20,122 stays, years 2014 to 2016, median stay 5.5 days; prevalence vasopressor 27.8%, AKI 55.9%, death 8.8%, Sepsis-3 0.28%, readmission 10.5%. + +### Code landed on `rebuttal/integration` + +Event pinning for old checkpoints (eICU checkpoints trained before PR #222 have five hazard heads; the registry now builds six because Sepsis-3 resolves; `run_pins.json` plus a legacy table), hazard-head scoring in `interventions.py` (`--hazard-heads`), `scripts/probe_channel.py` (with bank save and resume), the three GEMINI run.sh steps, `scripts/make_hparams_table.py`, `scripts/cohort_counts.py`, `scripts/panel_coverage.py`, `scripts/compare_runs_paired.py`. 1,503 tests pass. Not merged to main yet. + ## Schedule | Day | VM1 (MIMIC) | VM2 (eICU) | Lead / Amrit | diff --git a/odyssey/inference/alerts.py b/odyssey/inference/alerts.py index af60293a..c3231b2d 100644 --- a/odyssey/inference/alerts.py +++ b/odyssey/inference/alerts.py @@ -322,6 +322,30 @@ def _landmark_mask( # --------------------------------------------------------------------------- +def position_visit_starts( + sids: torch.Tensor, + vids: torch.Tensor, + times: torch.Tensor, + visit_start: dict[tuple[int, int], float], +) -> torch.Tensor: + """Visit start hours per position, shaped like ``times``. + + Looks up the unique (subject, visit) keys of the chunk once (a few + hundred lookups, not one per token); a key with no visit start maps + to 0. Shared by :func:`collect_model_scores` and the intervention + pass in :mod:`odyssey.inference.interventions`, so both select the + same landmark rows. + """ + keys = torch.stack([sids, vids], dim=-1).reshape(-1, 2) + unique_keys, inverse = torch.unique(keys, dim=0, return_inverse=True) + unique_starts = torch.tensor( + [visit_start.get((int(s), int(v)), 0.0) for s, v in unique_keys.tolist()], + dtype=times.dtype, + device=times.device, + ) + return unique_starts[inverse].view_as(times) + + def _event_token_mask( vocab: Vocabulary, alert: AlertEvent, device: str ) -> torch.Tensor: @@ -480,19 +504,7 @@ def collect_model_scores( times = chunk.batch.aux.time_stamps # Packed-path timestamps are already in the true frame (see # packed_context._truncate_head): no un-rebasing. - # visit start per position, via the unique (subject, visit) - # keys in this chunk (a few hundred lookups, not one per token) - keys = torch.stack([sids, vids], dim=-1).reshape(-1, 2) - unique_keys, inverse = torch.unique(keys, dim=0, return_inverse=True) - unique_starts = torch.tensor( - [ - visit_start.get((int(s), int(v)), 0.0) - for s, v in unique_keys.tolist() - ], - dtype=times.dtype, - device=times.device, - ) - starts = unique_starts[inverse].view_as(times) + starts = position_visit_starts(sids, vids, times, visit_start) # Static/demographic tokens build_patient_sequence prepends # (GENDER etc., stamped with the first real event's time) carry # visit_id=-1, not a real encounter -- _landmark_mask's own diff --git a/odyssey/inference/interventions.py b/odyssey/inference/interventions.py index 9c1c12cd..38b2c90f 100644 --- a/odyssey/inference/interventions.py +++ b/odyssey/inference/interventions.py @@ -69,6 +69,21 @@ visit (see :func:`_position_labels`). Everything is gated by the observed mask: unobserved concepts keep the model's own probability, in every mode -- there is no ground truth to feed there. + +Hazard heads (``--hazard-heads``). The next-event scores above say +nothing about the clinical alert heads. With the flag on, the same pass +also reads each per-event hazard head at the landmark rows the alert +protocol scores (:mod:`odyssey.inference.alerts`: the first event of +every 4 h bucket in each visit, at-risk rows only, right-censored rows +dropped per horizon) and reports, per mode, event and horizon, the AUROC +against the landmark outcome and the mean predicted ``P(event within +h)``. Because the override is applied at every position, the readout is +what the alert would have said had the concept been overridden all +along. Three paired, subject-clustered bootstrap differences (truth +minus none, truth minus flip, flip minus none) are attached to the +left-hand mode's entry. For ``flip_gated`` the heads are read from the +flip-intervened features: the logit gate has no hazard analogue. Off by +default, so an existing run reproduces byte for byte. """ import json @@ -76,17 +91,33 @@ from collections.abc import Sequence from dataclasses import asdict, dataclass, field, replace from pathlib import Path +from typing import Any +import numpy as np import polars as pl import torch import torch.nn.functional as F # noqa: N812 +from sklearn.metrics import roc_auc_score +from odyssey.data.alert_events import ( + AlertEvent, + EventTimes, + alert_events_for, + all_event_times, +) from odyssey.data.code_normalization import maybe_normalize from odyssey.data.history_recap import maybe_history_recap from odyssey.data.sidecars import activate_sidecars from odyssey.data.streaming import NO_SUBJECT, PackedLaneSampler, StreamingChunk from odyssey.data.value_binning import add_value_tokens from odyssey.data.vocabulary import PAD_ID, Vocabulary +from odyssey.inference.alerts import ( + HORIZONS_HOURS, + LandmarkState, + _landmark_mask, + _visit_starts, + position_visit_starts, +) from odyssey.inference.concept_attribution import ( calibrated_gammas, mean_concept_directions, @@ -98,6 +129,7 @@ load_run, refuse_existing_output, ) +from odyssey.inference.uncertainty import BootstrapAUROCDelta, bootstrap_auroc_delta from odyssey.models.concept_bottleneck import ( BottleneckIntervention, intervention_apply_mask, @@ -107,6 +139,7 @@ ConceptLabelDict, ConceptSupervision, ) +from odyssey.models.time_to_event import EventHazardHeads, probability_within from odyssey.training.data import ( build_concept_first_times, build_concept_label_dicts, @@ -136,6 +169,13 @@ CALIBRATED_MODES = ("truth_calibrated", "flip_calibrated") +#: Paired hazard differences, (left, right), reported on the left mode's entry. +HAZARD_PAIRS: tuple[tuple[str, str], ...] = ( + ("truth", "none"), + ("truth", "flip"), + ("flip", "none"), +) + @dataclass(frozen=True) class InterventionResult: @@ -183,6 +223,372 @@ class InterventionResult: step ``tau / peak_i`` (attached with concept names by :func:`evaluate_interventions`).""" + hazard: dict[str, Any] | None = None + """``{event: {"8h": {auroc, mean_risk, n_at_risk, n_positive, + n_censored}}}`` from the hazard heads at landmark rows under this + mode (``--hazard-heads``); None when hazard scoring is off.""" + + hazard_paired: dict[str, Any] | None = None + """``{"truth_minus_none": {event: {"8h": {auroc, mean_risk, ...}}}}``: + the :data:`HAZARD_PAIRS` differences whose left mode is this one, + each with a subject-clustered paired bootstrap interval. ``{}`` for a + mode that is never a left operand; None when hazard scoring is off.""" + + +def result_to_json(result: InterventionResult) -> dict[str, Any]: + """Serialise one result; the hazard keys appear only when scored. + + With ``--hazard-heads`` off the dict is exactly what earlier versions + wrote, so banked JSONs stay byte for byte reproducible. + """ + out = asdict(result) + for key in ("hazard", "hazard_paired"): + if out[key] is None: + del out[key] + return out + + +# --------------------------------------------------------------------------- +# Hazard heads at landmark rows +# --------------------------------------------------------------------------- + + +@dataclass(frozen=True) +class LandmarkRiskTable: + """Hazard-head risk at every landmark row of one intervention pass.""" + + event_names: list[str] + horizons: list[float] + subject_ids: np.ndarray + visit_ids: np.ndarray + time_hours: np.ndarray + risk: np.ndarray + """``(n_rows, n_events, n_horizons)`` P(event within horizon).""" + + @property + def n_rows(self) -> int: + """Number of landmark rows.""" + return int(self.subject_ids.shape[0]) + + def same_rows(self, other: "LandmarkRiskTable") -> bool: + """Whether both tables hold the same (subject, visit, time) rows in order.""" + return ( + np.array_equal(self.subject_ids, other.subject_ids) + and np.array_equal(self.visit_ids, other.visit_ids) + and np.array_equal(self.time_hours, other.time_hours) + ) + + +class HazardLandmarkScorer: + """Read P(event within h) from the hazard heads at landmark rows. + + Landmark rows are chosen with the alert protocol's own mask + (:func:`odyssey.inference.alerts._landmark_mask`), state threaded + across chunks, so the rows are the ones Tables 12 to 14 score. Only + the landmark rows are kept: subject, visit, time and one risk per + event and horizon. + """ + + def __init__( + self, + event_heads: EventHazardHeads, + alerts: Sequence[AlertEvent], + visit_start: dict[tuple[int, int], float], + *, + landmark_hours: float = 4.0, + horizons: Sequence[float] = HORIZONS_HOURS, + ) -> None: + """Bind the heads, the events to read and the visit envelope.""" + missing = [a.name for a in alerts if a.name not in event_heads.event_names] + if missing: + raise ValueError(f"no hazard head for alert events {missing}") + self.event_heads = event_heads + self.event_names = [a.name for a in alerts] + self.head_index = [event_heads.event_names.index(n) for n in self.event_names] + self.visit_start = visit_start + self.landmark_hours = landmark_hours + self.horizons = [float(h) for h in horizons] + self._state: LandmarkState | None = None + self._sids: list[np.ndarray] = [] + self._vids: list[np.ndarray] = [] + self._times: list[np.ndarray] = [] + self._risk: list[np.ndarray] = [] + + def add_chunk(self, chunk: StreamingChunk, features: torch.Tensor) -> None: + """Score this chunk's landmark positions from ``features``.""" + sids, vids = chunk.subject_ids, chunk.visit_ids + times = chunk.batch.aux.time_stamps + starts = position_visit_starts(sids, vids, times, self.visit_start) + keep, self._state = _landmark_mask( + times, sids, vids, self.landmark_hours, starts, state=self._state + ) + if not keep.any(): + return + keep = keep.to(features.device) + logits = self.event_heads(features[keep])[:, self.head_index] + risk = torch.stack( + [ + probability_within(logits, self.event_heads.edges, h) + for h in self.horizons + ], + dim=-1, + ) + self._sids.append(sids[keep].cpu().numpy().astype(np.int64)) + self._vids.append(vids[keep].cpu().numpy().astype(np.int64)) + self._times.append(times[keep].cpu().numpy().astype(np.float64)) + self._risk.append(risk.float().cpu().numpy()) + + def table(self) -> LandmarkRiskTable: + """Concatenate everything accumulated so far.""" + n_events, n_h = len(self.event_names), len(self.horizons) + return LandmarkRiskTable( + event_names=list(self.event_names), + horizons=list(self.horizons), + subject_ids=( + np.concatenate(self._sids) if self._sids else np.zeros(0, np.int64) + ), + visit_ids=( + np.concatenate(self._vids) if self._vids else np.zeros(0, np.int64) + ), + time_hours=( + np.concatenate(self._times) if self._times else np.zeros(0, np.float64) + ), + risk=( + np.concatenate(self._risk) + if self._risk + else np.zeros((0, n_events, n_h), np.float32) + ), + ) + + +def landmark_labels( + table: LandmarkRiskTable, times: dict[str, EventTimes] +) -> np.ndarray: + """``(n_rows, n_events, n_horizons)`` int8 labels: 1, 0, or -1. + + -1 marks a row the alert protocol does not score for that horizon: + not at risk (the event already happened) or censored (follow-up ends + before ``t + h``). The rule is :func:`odyssey.inference.alerts.outcome_at_horizon`, + vectorised per event. + """ + n = table.n_rows + labels = np.full((n, len(table.event_names), len(table.horizons)), -1, np.int8) + t = table.time_hours + for e, name in enumerate(table.event_names): + ev = times[name] + vids = np.full(n, -1, np.int64) if ev.subject_scoped else table.visit_ids + keys = zip(table.subject_ids.tolist(), vids.tolist()) + onset = np.empty(n, np.float64) + censor = np.empty(n, np.float64) + for i, key in enumerate(keys): + o = ev.onset.get(key) + c = ev.censor.get(key) + onset[i] = np.inf if o is None else o + censor[i] = -np.inf if c is None else c + at_risk = ~(onset <= t) + for j, h in enumerate(table.horizons): + positive = at_risk & (onset <= t + h) + negative = at_risk & ~positive & (censor >= t + h) + labels[positive, e, j] = 1 + labels[negative, e, j] = 0 + return labels + + +def _horizon_key(table: LandmarkRiskTable, j: int) -> str: + return f"{table.horizons[j]:g}h" + + +def hazard_summary(table: LandmarkRiskTable, labels: np.ndarray) -> dict[str, Any]: + """Per event and horizon: AUROC, mean risk and the row counts.""" + out: dict[str, Any] = {} + for e, name in enumerate(table.event_names): + out[name] = {} + for j in range(len(table.horizons)): + y = labels[:, e, j] + ok = y >= 0 + p = table.risk[ok, e, j].astype(np.float64) + y_ok = y[ok] + two_class = ok.any() and y_ok.min() != y_ok.max() + out[name][_horizon_key(table, j)] = { + "auroc": float(roc_auc_score(y_ok, p)) if two_class else None, + "mean_risk": float(p.mean()) if ok.any() else None, + "n_at_risk": int(ok.sum()), + "n_positive": int(y_ok.sum()), + "n_censored": int((~ok).sum()), + } + return out + + +@dataclass(frozen=True) +class PairedMeanDelta: + """Mean over rows of ``a - b`` with a subject-clustered bootstrap interval.""" + + point: float + ci_low: float + ci_high: float + n_rows: int + n_subjects: int + + @property + def separated(self) -> bool: + """Whether the interval excludes zero.""" + return self.ci_low > 0.0 or self.ci_high < 0.0 + + +def subject_bootstrap_means( + diff: np.ndarray, + subject_ids: np.ndarray, + *, + n_boot: int, + seed: int, + block: int = 64, +) -> np.ndarray: + """``(n_boot,)`` resampled means of ``diff``, drawing whole subjects. + + Rows are summed per subject once; each resample draws ``n_subjects`` + subjects with replacement and divides the drawn sums by the drawn + counts, so a subject's rows always move together. Draws are made in + blocks to bound memory. + """ + diff = np.asarray(diff, dtype=np.float64) + _, inverse = np.unique(np.asarray(subject_ids), return_inverse=True) + n_subjects = int(inverse.max()) + 1 if inverse.size else 0 + if n_subjects == 0: + return np.full(n_boot, np.nan) + sums = np.bincount(inverse, weights=diff, minlength=n_subjects) + counts = np.bincount(inverse, minlength=n_subjects).astype(np.float64) + rng = np.random.default_rng(seed) + out = np.empty(n_boot, dtype=np.float64) + for start in range(0, n_boot, block): + k = min(block, n_boot - start) + drawn = rng.integers(0, n_subjects, size=(k, n_subjects)) + out[start : start + k] = sums[drawn].sum(axis=1) / counts[drawn].sum(axis=1) + return out + + +def paired_mean_delta( + a: np.ndarray, + b: np.ndarray, + subject_ids: np.ndarray, + *, + n_boot: int = 1000, + seed: int = 0, +) -> PairedMeanDelta: + """Subject-clustered paired bootstrap of ``mean(a) - mean(b)`` on the same rows.""" + diff = np.asarray(a, dtype=np.float64) - np.asarray(b, dtype=np.float64) + subjects = np.asarray(subject_ids) + n_subjects = int(np.unique(subjects).shape[0]) + if diff.size == 0: + return PairedMeanDelta(float("nan"), float("nan"), float("nan"), 0, 0) + boots = subject_bootstrap_means(diff, subjects, n_boot=n_boot, seed=seed) + return PairedMeanDelta( + float(diff.mean()), + float(np.percentile(boots, 2.5)), + float(np.percentile(boots, 97.5)), + int(diff.size), + n_subjects, + ) + + +def _auroc_delta_json(d: BootstrapAUROCDelta | None) -> dict[str, Any] | None: + if d is None: + return None + return { + "point": d.point_estimate, + "ci_low": d.ci_low, + "ci_high": d.ci_high, + "n_boot_used": d.n_boot_used, + "n_boot_skipped": d.n_boot_skipped, + "separated": d.excludes_zero(), + } + + +def hazard_paired_summary( + left: LandmarkRiskTable, + right: LandmarkRiskTable, + labels: np.ndarray, + *, + n_boot: int, + seed: int, +) -> dict[str, Any]: + """Per event and horizon: paired bootstrap of AUROC and mean-risk differences. + + Both tables must hold the same rows (checked). AUROC differences use + :func:`odyssey.inference.uncertainty.bootstrap_auroc_delta`, the + helper behind ``scripts/alerts_cis.py``; mean-risk differences use + :func:`paired_mean_delta`. Both draw whole subjects. + """ + if not left.same_rows(right): + raise ValueError("paired hazard scoring needs identical landmark rows") + out: dict[str, Any] = {} + for e, name in enumerate(left.event_names): + out[name] = {} + for j in range(len(left.horizons)): + y = labels[:, e, j] + ok = y >= 0 + y_ok = y[ok].astype(np.float64) + p_left = left.risk[ok, e, j].astype(np.float64) + p_right = right.risk[ok, e, j].astype(np.float64) + subjects = left.subject_ids[ok] + auroc = ( + bootstrap_auroc_delta( + y_ok, p_left, p_right, subjects, n_boot=n_boot, seed=seed + ) + if ok.any() + else None + ) + mean = paired_mean_delta( + p_left, p_right, subjects, n_boot=n_boot, seed=seed + ) + out[name][_horizon_key(left, j)] = { + "auroc": _auroc_delta_json(auroc), + # None, like auroc, when the cell has no scoreable row. + "mean_risk": ( + { + "point": mean.point, + "ci_low": mean.ci_low, + "ci_high": mean.ci_high, + "separated": mean.separated, + } + if mean.n_rows + else None + ), + "n_at_risk": int(ok.sum()), + "n_subjects": mean.n_subjects, + } + return out + + +def score_hazard_landmarks( + tables: dict[str, LandmarkRiskTable], + times: dict[str, EventTimes], + *, + n_boot: int = 1000, + seed: int = 0, +) -> tuple[dict[str, dict[str, Any]], dict[str, dict[str, Any]]]: + """Return ``(hazard by mode, hazard_paired by mode)`` for the JSON output. + + Labels are computed once from the first table; every other table must + hold the same rows, which the shared sampler and landmark mask + guarantee. Pairs come from :data:`HAZARD_PAIRS`, limited to the modes + that were run, and land on the left mode's entry. + """ + if not tables: + return {}, {} + first = next(iter(tables.values())) + for mode, table in tables.items(): + if not table.same_rows(first): + raise ValueError(f"mode {mode!r} scored different landmark rows") + labels = landmark_labels(first, times) + per_mode = {mode: hazard_summary(table, labels) for mode, table in tables.items()} + paired: dict[str, dict[str, Any]] = {mode: {} for mode in tables} + for left, right in HAZARD_PAIRS: + if left in tables and right in tables: + paired[left][f"{left}_minus_{right}"] = hazard_paired_summary( + tables[left], tables[right], labels, n_boot=n_boot, seed=seed + ) + return per_mode, paired + def _chunk_intervention( chunk: StreamingChunk, @@ -250,6 +656,7 @@ def run_streaming_intervention( # noqa: PLR0912, PLR0915 -- one linear scoring calibration_gammas: torch.Tensor | None = None, calibrated_tau: float | None = None, concept_names: Sequence[str] | None = None, + hazard_scorer: HazardLandmarkScorer | None = None, ) -> InterventionResult: """Score next-event prediction under one intervention mode. @@ -259,6 +666,11 @@ def run_streaming_intervention( # noqa: PLR0912, PLR0915 -- one linear scoring (scripts/intervention_cis.py), which the aggregate numbers alone cannot support. + ``hazard_scorer``, if given, also reads the per-event hazard heads at + the alert protocol's landmark rows from the intervened features (see + :class:`HazardLandmarkScorer`); it changes nothing in the returned + next-event numbers. + The identical streaming pass as :func:`~odyssey.inference.run_inference.run_streaming_inference` (same sampler, same state carrying), with the bottleneck edited per @@ -375,13 +787,16 @@ def run_streaming_intervention( # noqa: PLR0912, PLR0915 -- one linear scoring base_logits, bottleneck_out, state = model( chunk.batch, state=state_in, reset_mask=chunk.reset_mask ) - flip_logits, _, _ = model( + flip_logits, flip_out, _ = model( chunk.batch, state=state_in, reset_mask=chunk.reset_mask, intervention=intervention, ) logits = base_logits + torch.clamp_max(flip_logits - base_logits, 0.0) + # The gate has no hazard analogue: the heads read the + # flip-intervened features, as in plain flip mode. + hazard_features = flip_out.bottleneck else: logits, bottleneck_out, state = model( chunk.batch, @@ -389,6 +804,10 @@ def run_streaming_intervention( # noqa: PLR0912, PLR0915 -- one linear scoring reset_mask=chunk.reset_mask, intervention=intervention, ) + if mode != "flip_gated": + hazard_features = bottleneck_out.bottleneck + if hazard_scorer is not None: + hazard_scorer.add_chunk(chunk, hazard_features) real = chunk.real_mask if intervention is not None and intervention.probs is not None: own = bottleneck_out.concept_probs @@ -496,6 +915,10 @@ def evaluate_interventions( uncertain_band: float | None = None, per_subject_out: dict[str, dict[int, list[int]]] | None = None, calibrated_tau: float = 1.0, + hazard_heads: bool = False, + hazard_boot: int = 1000, + landmark_hours: float = 4.0, + horizons: Sequence[float] = HORIZONS_HOURS, ) -> list[InterventionResult]: """End-to-end: load a trained run, score every intervention mode. @@ -503,6 +926,14 @@ def evaluate_interventions( [top1_hits, n_predictions]}}`` (see :func:`run_streaming_intervention`). + ``hazard_heads`` adds the hazard-head readout at landmark rows to + every mode (``InterventionResult.hazard``) and the + :data:`HAZARD_PAIRS` paired bootstrap differences with ``hazard_boot`` + resamples (``InterventionResult.hazard_paired``). Events are the + run's landmark alert events that have a hazard head; outcomes and + landmark rows follow :mod:`odyssey.inference.alerts` exactly. The + bootstrap seed is ``seed``. + Data preparation matches :func:`~odyssey.inference.run_inference.evaluate_run` exactly (same normalization, binning, and label scoping from the run's own @@ -518,6 +949,11 @@ def evaluate_interventions( "this evaluation needs a concept bottleneck; the run's model_kind is " f"{getattr(config, 'model_kind', 'bottleneck')!r}" ) + if hazard_heads and model.event_heads is None: + raise ValueError( + "hazard scoring needs a run trained with per-event hazard heads " + "(event_hazards); this run has none" + ) if getattr(config, "backbone", "hybrid") == "transformer": # The TBTT stream below is the same one the transformer trained on: # its ``state`` is an inert sentinel and every chunk is its own @@ -552,55 +988,46 @@ def evaluate_interventions( events_binned = add_value_tokens(raw_events, binner, source=source) supervision: ConceptSupervision = getattr(config, "concept_supervision", "stay") - concept_labels: ConceptLabelDict - concept_mask: ConceptLabelDict - concept_first_times: ConceptLabelDict - if supervision == "visit": - concept_labels, concept_mask = build_visit_concept_label_dicts( - raw_events, concepts - ) - concept_first_times = build_visit_concept_first_times(raw_events, concepts) - else: - concept_labels, concept_mask = build_concept_label_dicts(raw_events, concepts) - concept_first_times = build_concept_first_times(raw_events, concepts) + concept_labels, concept_mask, concept_first_times = _concept_labels( + raw_events, concepts, supervision + ) + hazard_targets = ( + _hazard_targets(model, config, raw_events, source=source) + if hazard_heads + else None + ) del raw_events calibration_gammas: torch.Tensor | None = None gamma_by_name: dict[str, float] | None = None if any(m in CALIBRATED_MODES for m in modes): - # Only a bottleneck whose per-unit displacement is data dependent - # needs the estimation pass. The decomposition's displacement is a - # parameter, so asking for directions there would be a forward - # pass over the whole split to recover something already stored. - directions = ( - mean_concept_directions( - model, - events_binned, - vocab, - num_lanes=num_lanes, - chunk_size=chunk_size, - device=device, - ) - if getattr(model.bottleneck, "needs_calibration_directions", True) - else None - ) - calibration_gammas = calibrated_gammas(model, directions, tau=calibrated_tau) - gamma_by_name = { - c.name: float(g) - for c, g in zip(concepts, calibration_gammas.tolist(), strict=True) - } - logger.info( - "[interventions] calibrated gammas (tau=%.3g): %s", - calibrated_tau, - {k: round(v, 4) for k, v in gamma_by_name.items()}, + calibration_gammas, gamma_by_name = _calibration( + model, + events_binned, + vocab, + concept_names=[c.name for c in concepts], + calibrated_tau=calibrated_tau, + num_lanes=num_lanes, + chunk_size=chunk_size, + device=device, ) results = [] + tables: dict[str, LandmarkRiskTable] = {} for mode in modes: logger.info("[interventions] scoring mode %r", mode) mode_subjects: dict[int, list[int]] | None = None if per_subject_out is not None: mode_subjects = per_subject_out.setdefault(mode, {}) + scorer: HazardLandmarkScorer | None = None + if hazard_targets is not None: + scorer = HazardLandmarkScorer( + hazard_targets.event_heads, + hazard_targets.alerts, + hazard_targets.visit_start, + landmark_hours=landmark_hours, + horizons=horizons, + ) result = run_streaming_intervention( model, events_binned, @@ -619,9 +1046,12 @@ def evaluate_interventions( calibration_gammas=calibration_gammas, calibrated_tau=calibrated_tau, concept_names=[c.name for c in concepts], + hazard_scorer=scorer, ) if mode in CALIBRATED_MODES: result = replace(result, calibration_gamma=gamma_by_name) + if scorer is not None: + tables[mode] = scorer.table() results.append(result) baseline = results[0] latest = results[-1] @@ -632,9 +1062,137 @@ def evaluate_interventions( latest.top1_accuracy - baseline.top1_accuracy, latest.mean_task_loss, ) + if hazard_targets is not None: + logger.info( + "[interventions] paired hazard bootstrap: %d resamples over %d landmark rows", + hazard_boot, + next(iter(tables.values())).n_rows if tables else 0, + ) + per_mode, paired = score_hazard_landmarks( + tables, hazard_targets.times, n_boot=hazard_boot, seed=seed + ) + results = [ + replace(r, hazard=per_mode[r.mode], hazard_paired=paired[r.mode]) + for r in results + ] + _log_hazard(results) return results +def _concept_labels( + raw_events: pl.DataFrame, concepts: Sequence[Any], supervision: ConceptSupervision +) -> tuple[ConceptLabelDict, ConceptLabelDict, ConceptLabelDict]: + """``(labels, mask, first_times)`` under the run's own label scoping.""" + if supervision == "visit": + visit_labels, visit_mask = build_visit_concept_label_dicts(raw_events, concepts) + first = build_visit_concept_first_times(raw_events, concepts) + return visit_labels, visit_mask, first + stay_labels, stay_mask = build_concept_label_dicts(raw_events, concepts) + return stay_labels, stay_mask, build_concept_first_times(raw_events, concepts) + + +def _calibration( + model: ConceptBottleneckSequenceModel, + events_binned: pl.DataFrame, + vocab: Vocabulary, + *, + concept_names: Sequence[str], + calibrated_tau: float, + num_lanes: int, + chunk_size: int, + device: str, +) -> tuple[torch.Tensor, dict[str, float]]: + """Per-concept calibrated steps for the *_calibrated modes, with names.""" + # Only a bottleneck whose per-unit displacement is data dependent + # needs the estimation pass. The decomposition's displacement is a + # parameter, so asking for directions there would be a forward + # pass over the whole split to recover something already stored. + directions = ( + mean_concept_directions( + model, + events_binned, + vocab, + num_lanes=num_lanes, + chunk_size=chunk_size, + device=device, + ) + if getattr(model.bottleneck, "needs_calibration_directions", True) + else None + ) + calibration_gammas = calibrated_gammas(model, directions, tau=calibrated_tau) + gamma_by_name = { + name: float(g) + for name, g in zip(concept_names, calibration_gammas.tolist(), strict=True) + } + logger.info( + "[interventions] calibrated gammas (tau=%.3g): %s", + calibrated_tau, + {k: round(v, 4) for k, v in gamma_by_name.items()}, + ) + return calibration_gammas, gamma_by_name + + +@dataclass(frozen=True) +class _HazardTargets: + """What the hazard readout needs from the held-out split, built once.""" + + event_heads: EventHazardHeads + alerts: list[AlertEvent] + times: dict[str, EventTimes] + visit_start: dict[tuple[int, int], float] + + +def _hazard_targets( + model: ConceptBottleneckSequenceModel, + config: object, + raw_events: pl.DataFrame, + *, + source: str, +) -> _HazardTargets: + """Alert events with a head, their onset tables and the visit envelope. + + The events are the run's landmark alert events as + :func:`odyssey.inference.alerts.evaluate_alerts` scores them, + restricted to those the model has a hazard head for. + """ + if model.event_heads is None: + raise ValueError( + "hazard scoring needs a run trained with per-event hazard heads " + "(event_hazards); this run has none" + ) + task_set = getattr(config, "task_set", "v1") + head_names = set(model.event_heads.event_names) + alerts = [ + a + for a in alert_events_for(task_set, source=source) + if not a.next_visit and a.name in head_names + ] + if not alerts: + raise ValueError( + f"task_set {task_set!r} has no landmark alert event with a hazard head" + ) + logger.info("[interventions] hazard heads for %s", [a.name for a in alerts]) + return _HazardTargets( + event_heads=model.event_heads, + alerts=alerts, + times=all_event_times(raw_events, alerts, source, task_set=task_set), + visit_start=_visit_starts(raw_events), + ) + + +def _log_hazard(results: Sequence[InterventionResult]) -> None: + for r in results: + for event, by_h in (r.hazard or {}).items(): + cells = { + h: ( + None if c["auroc"] is None else round(c["auroc"], 4), + None if c["mean_risk"] is None else round(c["mean_risk"], 4), + ) + for h, c in by_h.items() + } + logger.info("[interventions] %s hazard %s: %s", r.mode, event, cells) + + def _main() -> None: import argparse # noqa: PLC0415 @@ -672,6 +1230,23 @@ def _main() -> None: ) parser.add_argument("--num-lanes", type=int, default=8) parser.add_argument("--chunk-size", type=int, default=256) + parser.add_argument( + "--hazard-heads", + action="store_true", + help=( + "also read the per-event hazard heads at the alert protocol's " + "landmark rows under every mode (AUROC and mean P(event within h) " + "per event and horizon) and attach paired subject-clustered " + "bootstrap differences for truth-none, truth-flip and flip-none. " + "Off by default: without it the output is unchanged." + ), + ) + parser.add_argument( + "--hazard-boot", + type=int, + default=1000, + help="bootstrap resamples for the paired hazard differences.", + ) parser.add_argument( "--dump-per-subject", action="store_true", @@ -712,9 +1287,11 @@ def _main() -> None: uncertain_band=args.uncertain_band, per_subject_out=per_subject, calibrated_tau=args.calibrated_tau, + hazard_heads=args.hazard_heads, + hazard_boot=args.hazard_boot, ) out.parent.mkdir(parents=True, exist_ok=True) - out.write_text(json.dumps([asdict(r) for r in results], indent=2)) + out.write_text(json.dumps([result_to_json(r) for r in results], indent=2)) logger.info("[interventions] wrote %d modes to %s", len(results), out) if per_subject is not None: ps_out = out.with_name(out.stem + "_per_subject.json") @@ -738,8 +1315,19 @@ def _main() -> None: __all__ = [ + "HAZARD_PAIRS", "INTERVENTION_MODES", + "HazardLandmarkScorer", "InterventionResult", - "run_streaming_intervention", + "LandmarkRiskTable", + "PairedMeanDelta", "evaluate_interventions", + "hazard_paired_summary", + "hazard_summary", + "landmark_labels", + "paired_mean_delta", + "result_to_json", + "run_streaming_intervention", + "score_hazard_landmarks", + "subject_bootstrap_means", ] diff --git a/odyssey/inference/leakage.py b/odyssey/inference/leakage.py index 7082f899..c1f31ce6 100644 --- a/odyssey/inference/leakage.py +++ b/odyssey/inference/leakage.py @@ -396,6 +396,36 @@ def _feature_stats(x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: return mean, std +def _probe_device(head: nn.Module) -> torch.device: + """Return the device a probe's parameters live on (where batches must go).""" + return next(head.parameters()).device + + +def _batches(n: int, batch_size: int) -> Iterator[tuple[int, int]]: + for start in range(0, n, batch_size): + yield start, min(n, start + batch_size) + + +@torch.no_grad() +def _predict_in_batches( + head: _StandardizedLinearProbe, x: torch.Tensor, batch_size: int = 65536 +) -> torch.Tensor: + """``head(x)`` on CPU, batch by batch, whatever device ``x`` lives on. + + The bank may sit in host memory (``--bank-on-cpu``) while the probe is + on the GPU; each batch is upcast and moved to the probe's device, and + the logits come back to CPU so downstream numpy/scoring never sees a + CUDA tensor. + """ + head.eval() + device = _probe_device(head) + out = [ + head(x[a:b].float().to(device)).cpu() + for a, b in _batches(x.shape[0], batch_size) + ] + return torch.cat(out) + + def _fit_categorical_probe( in_features: int, num_classes: int, @@ -411,7 +441,13 @@ def _fit_categorical_probe( seed: int = 0, device: str = "cpu", ) -> tuple[_StandardizedLinearProbe, ProbeFitTrace]: - """Multinomial-logistic probe (CTL): Adam on train, early-stopped on tuning CE.""" + """Multinomial-logistic probe (CTL): Adam on train, early-stopped on tuning CE. + + The bank tensors may live on a different device from the probe (a + CPU bank feeding a GPU head): every batch is moved to ``device`` + before it touches the head, and the tuning loss is accumulated batch + by batch rather than by materializing the whole tuning matrix there. + """ torch.manual_seed(seed) mean, std = _feature_stats(train_x.float()) head = _StandardizedLinearProbe(in_features, num_classes, mean, std).to(device) @@ -421,6 +457,7 @@ def _fit_categorical_probe( best = float("inf") bad = 0 n = train_x.shape[0] + n_tune = tune_x.shape[0] t0 = time.time() for epoch in range(epochs): head.train() @@ -428,11 +465,19 @@ def _fit_categorical_probe( for start in range(0, n, batch_size): idx = perm[start : start + batch_size] opt.zero_grad() - loss = F.cross_entropy(head(train_x[idx].float()), train_y[idx]) + x = train_x[idx].float().to(device) + y = train_y[idx].to(device) + loss = F.cross_entropy(head(x), y) loss.backward() # type: ignore[no-untyped-call] opt.step() + head.eval() with torch.no_grad(): - tune_loss = float(F.cross_entropy(head(tune_x.float()), tune_y).item()) + tune_sum = 0.0 + for a, b in _batches(n_tune, batch_size): + x = tune_x[a:b].float().to(device) + y = tune_y[a:b].to(device) + tune_sum += float(F.cross_entropy(head(x), y, reduction="sum").item()) + tune_loss = tune_sum / max(1, n_tune) trace.tuning_loss.append(tune_loss) if tune_loss < best - 1e-6: best, bad = tune_loss, 0 @@ -451,9 +496,9 @@ def _fit_categorical_probe( def _score_categorical_probe( head: _StandardizedLinearProbe, x: torch.Tensor, y: torch.Tensor ) -> tuple[float, float]: - """Return (accuracy, mean cross-entropy) on ``(x, y)``.""" - head.eval() - logits = head(x.float()) + """Return (accuracy, mean cross-entropy) on ``(x, y)``, on any device.""" + logits = _predict_in_batches(head, x) + y = y.cpu() ce = float(F.cross_entropy(logits, y).item()) acc = float((logits.argmax(dim=-1) == y).float().mean().item()) return acc, ce @@ -499,6 +544,7 @@ def masked_bce( best = float("inf") bad = 0 n = train_x.shape[0] + n_tune = tune_x.shape[0] t0 = time.time() for epoch in range(epochs): head.train() @@ -506,15 +552,30 @@ def masked_bce( for start in range(0, n, batch_size): idx = perm[start : start + batch_size] opt.zero_grad() + # Batches move to the probe's device; the bank may be on CPU. loss = masked_bce( - head(train_x[idx].float()), train_y[idx], train_observed[idx] + head(train_x[idx].float().to(device)), + train_y[idx].to(device), + train_observed[idx].to(device), ) loss.backward() # type: ignore[no-untyped-call] opt.step() + head.eval() with torch.no_grad(): - tune_loss = float( - masked_bce(head(tune_x.float()), tune_y, tune_observed).item() - ) + # Masked mean over the whole tuning split, accumulated per batch + # (the ratio of sums, the same number as one unbatched call). + loss_sum = 0.0 + weight_sum = 0.0 + for a, b in _batches(n_tune, batch_size): + logits = head(tune_x[a:b].float().to(device)) + y = tune_y[a:b].to(device) + observed = tune_observed[a:b].to(device) + per_elem = F.binary_cross_entropy_with_logits( + logits, y, reduction="none" + ) + loss_sum += float((per_elem * observed.float()).sum().item()) + weight_sum += float(observed.float().sum().item()) + tune_loss = loss_sum / max(1.0, weight_sum) trace.tuning_loss.append(tune_loss) if tune_loss < best - 1e-6: best, bad = tune_loss, 0 @@ -586,7 +647,9 @@ class CTLResult: n_held_out: int -def _random_projection(in_dim: int, out_dim: int, seed: int) -> torch.Tensor: +def _random_projection( + in_dim: int, out_dim: int, seed: int, *, device: torch.device | str = "cpu" +) -> torch.Tensor: """Build a fixed ``(in_dim, out_dim)`` semi-orthogonal projection, seeded. CTL's capacity control (see the module docstring): a QR decomposition @@ -596,13 +659,18 @@ def _random_projection(in_dim: int, out_dim: int, seed: int) -> torch.Tensor: re-express it at ``out_dim`` width to match ``embeddings_only``'s input dimensionality. Deterministic and self-contained: uses a local :class:`torch.Generator`, never the global RNG, so it has no side - effect on any other seeded call in this module. + effect on any other seeded call in this module. Always drawn on CPU + (so the matrix is identical whatever ``device`` is) and then moved to + ``device``, which must be the device of the bank it multiplies: a + projection on the model's device against a bank kept in host memory + is exactly the mixed-device matmul that lost a 45-minute streaming + pass (2026-09-25). """ gen = torch.Generator().manual_seed(seed) raw = torch.randn(out_dim, in_dim, generator=gen) q, _ = torch.linalg.qr(raw) - out: torch.Tensor = q.T - return out + out: torch.Tensor = q.T.contiguous() + return out.to(device) def compute_ctl( @@ -629,16 +697,25 @@ def compute_ctl( class_names = dict(_CODE_TYPE_NAMES) num_concepts = train_bank.concept_probs.shape[1] embedding_dim = train_bank.concept_embeddings.shape[-1] - projection = _random_projection(num_concepts, num_concepts * embedding_dim, seed) + projection = _random_projection( + num_concepts, + num_concepts * embedding_dim, + seed, + device=train_bank.concept_probs.device, + ) def flat(bank: LeakageBank) -> dict[str, torch.Tensor]: + # Every feature matrix stays on the bank's own device; the probe + # fits move batches to ``device``. The projection follows the bank + # (the three banks are normally co-located, but nothing requires it). + probs = bank.concept_probs.float() return { - "probs_only": bank.concept_probs.float(), + "probs_only": probs, "embeddings_only": bank.concept_embeddings.float().reshape( len(bank), num_concepts * embedding_dim ), "unknown_only": bank.unknown_embedding.float(), - "probs_projected": bank.concept_probs.float() @ projection, + "probs_projected": probs @ projection.to(probs.device), } train_x = flat(train_bank) @@ -790,15 +867,22 @@ def compute_icl( seed=seed, device=device, ) - with torch.no_grad(): - embed_logits = embed_head(held_out_bank.concept_embeddings[:, i, :].float()) - prob_logits = prob_head(held_out_bank.concept_probs[:, i : i + 1].float()) + # Logits and targets on CPU regardless of where the bank or the + # probe lives: roc_auc_score reads numpy. + embed_logits = _predict_in_batches( + embed_head, held_out_bank.concept_embeddings[:, i, :] + ) + prob_logits = _predict_in_batches( + prob_head, held_out_bank.concept_probs[:, i : i + 1] + ) + held_labels = held_out_bank.concept_labels.cpu() + held_observed = held_out_bank.concept_observed.cpu() for j, concept_j in enumerate(names): if j == i: continue - obs = held_out_bank.concept_observed[:, j] + obs = held_observed[:, j] n_obs = int(obs.sum().item()) - y = held_out_bank.concept_labels[obs, j].numpy() + y = held_labels[obs, j].numpy() if n_obs < 2 or y.min() == y.max(): pairs.append( ICLPairScore(concept_i, concept_j, n_obs, None, None, None, None) diff --git a/odyssey/inference/legacy_concept_pins.py b/odyssey/inference/legacy_concept_pins.py index c987b318..a88b39c1 100644 --- a/odyssey/inference/legacy_concept_pins.py +++ b/odyssey/inference/legacy_concept_pins.py @@ -15,25 +15,95 @@ the same answer and a new mismatch fails with a readable message instead of archaeology. -Adding an entry: check out the run's training commit (its -``env_fingerprint.json`` records ``git_commit``) and print -``[c.name for c in concepts_for_source(source, task_set=...)]``. The ORDER is -load-bearing, the bottleneck's slots are positional, so paste the list as -printed rather than sorting it. +Hazard EVENT heads drift the same way: ``alert_events_for`` drops a +concept-backed alert whose concept does not resolve for the source, so the +eICU sepsis3 head appeared the day sepsis3 became resolvable there (PR #222) +and every eICU checkpoint trained before it now has one head fewer than +today's registry builds. :data:`LEGACY_EVENT_PINS` and +:func:`pinned_event_names` are the event-side twins of the concept table. + +Resolution order, for both concepts and events: + +1. ``run_pins.json`` in the run directory (:data:`PIN_FILENAME`), which + training writes since this module gained :func:`write_run_pins`, so any + newer run is self-describing. +2. The legacy tables here, keyed by run directory name, for runs trained + before the pin file existed. +3. Nothing: today's registry, with :func:`check_concept_count` and + :func:`check_event_count` refusing before ``load_state_dict`` when the + checkpoint's own widths disagree with it. + +Adding a legacy entry: check out the run's training commit (its +``env_fingerprint.json`` records ``git_commit``; ``docs/experiments.md`` +names it per run) and print +``[c.name for c in concepts_for_source(source, task_set=...)]`` and +``[a.name for a in hazard_events_for(task_set, source=source)]`` FROM THAT +CHECKOUT (a script run from another directory imports the editable install +instead, and reports today's list). The ORDER is load-bearing, the +bottleneck's slots and the hazard head's rows are positional, so paste the +list as printed rather than sorting it. """ from __future__ import annotations +import json from collections.abc import Sequence +from pathlib import Path from typing import TYPE_CHECKING from odyssey.data.concepts import canonical_concept_name, concepts_for_source +from odyssey.models.time_to_event import DEFAULT_TIME_BIN_EDGES_HOURS if TYPE_CHECKING: from odyssey.data.concepts import AnyConceptDefinition +#: Written into every run directory by training; read first by both +#: :func:`pinned_concept_names` and :func:`pinned_event_names`. +PIN_FILENAME = "run_pins.json" + +#: Bins per hazard head: ``EventHazardHeads.num_bins`` is ``len(edges) + 2`` +#: (a zero-gap bin, one per edge, and the open tail), so a checkpoint's +#: event count is its head's output width divided by this. +HAZARD_NUM_BINS = len(DEFAULT_TIME_BIN_EDGES_HOURS) + 2 + +# eICU-CRD before PR #222 (ca541d8, 2026-09-02): 26 of 29 concepts. SOFA had +# no eICU source config, so sepsis3, hypoxemic_respiratory_failure and +# oliguria dropped out; "shock" is the pre-rename name of +# sustained_hypotension_map (PR #225). Printed from +# concepts_for_source("eicu", task_set="v3") at cdbd4e7, the eICU flagship's +# training commit, and identical at 17d28fc, 1193aa3, a86616a, 69c8cb0 and +# f8633b3 (the other pre-#222 eICU runs' commits). +_EICU_26_CONCEPTS: tuple[str, ...] = ( + "tachycardia", + "bradycardia", + "hypotension", + "hypertension", + "hypoxia", + "fever", + "hypothermia", + "elevated_lactate", + "sustained_tachypnea", + "acute_kidney_injury", + "aki_stage_2", + "aki_stage_3", + "sirs", + "qsofa", + "on_vasopressors", + "hyperkalemia", + "hypokalemia", + "hyponatremia", + "hypernatremia", + "hypoglycemia", + "hyperglycemia", + "anemia", + "thrombocytopenia", + "coagulopathy", + "metabolic_acidosis", + "shock", +) + # run directory basename -> the concept names that run trained with, in slot # order. Verified against each checkpoint's own bottleneck width, not guessed. LEGACY_CONCEPT_PINS: dict[str, tuple[str, ...]] = { @@ -80,48 +150,134 @@ "qsofa", "on_vasopressors", ), - # eICU-CRD, 26 of 29; microbiology mapping landed after this run, so - # sepsis3 (and with it the sepsis3 alert head) resolves today but did - # not then. "shock" is the pre-rename name of sustained_hypotension_map. - "eicu_full_v10": ( - "tachycardia", - "bradycardia", - "hypotension", - "hypertension", - "hypoxia", - "fever", - "hypothermia", - "elevated_lactate", - "sustained_tachypnea", - "acute_kidney_injury", - "aki_stage_2", - "aki_stage_3", - "sirs", - "qsofa", - "on_vasopressors", - "hyperkalemia", - "hypokalemia", - "hyponatremia", - "hypernatremia", - "hypoglycemia", - "hyperglycemia", - "anemia", - "thrombocytopenia", - "coagulopathy", - "metabolic_acidosis", - "shock", - ), + # eICU-CRD, 26 of 29; verified against eicu_full_v10's own bottleneck + # width. sepsis3 (and with it the sepsis3 alert head) resolves today but + # did not then. + "eicu_full_v10": _EICU_26_CONCEPTS, + # The other eICU runs banked under research_journal/figure_data/vm2 that + # trained before PR #222, each on R6's config (docs/experiments.md) at a + # commit whose registry printed the same 26 names: eicu_full_L_v10 + # (17d28fc), eicu_full_ADD_v10 (pre-2026-09-01), eicu_full_DEC_v10 + # (1193aa3), eicu_full_DEC_v11 (a86616a), eicu_full_DEC_v12 (69c8cb0) and + # eicu_full_DEC_v12_steer (f8633b3, one epoch from the v12 checkpoint). + # Recovered from the training commits, not read off the checkpoints. + "eicu_full_L_v10": _EICU_26_CONCEPTS, + "eicu_full_ADD_v10": _EICU_26_CONCEPTS, + "eicu_full_DEC_v10": _EICU_26_CONCEPTS, + "eicu_full_DEC_v11": _EICU_26_CONCEPTS, + "eicu_full_DEC_v12": _EICU_26_CONCEPTS, + "eicu_full_DEC_v12_steer": _EICU_26_CONCEPTS, } +# eICU-CRD hazard heads before PR #222. Printed from +# alert_events_for("v3", source="eicu") at cdbd4e7 (eicu_full_v10's training +# commit) and identical at 17d28fc, 1193aa3, a86616a, 69c8cb0 and f8633b3: +# ('vasopressor_start', 'icu_admission', 'acute_kidney_injury', 'death', +# 'readmission_30d') +# The same call on main today (ca541d8 and later) gives +# ('vasopressor_start', 'icu_admission', 'acute_kidney_injury', 'death', +# 'sepsis3', 'readmission_30d') +# so sepsis3 lands in the MIDDLE of the list: without a pin the checkpoint's +# readmission_30d rows would be read as sepsis3 even if the width matched. +# Each checkpoint's event_heads.proj output width is 80 = 5 x 16 bins. +_EICU_PRE_SOFA_EVENTS: tuple[str, ...] = ( + "vasopressor_start", + "icu_admission", + "acute_kidney_injury", + "death", + "readmission_30d", +) + +# run directory basename -> the hazard event names that run trained heads +# for, in head order. MIMIC-IV runs are absent on purpose: their event set +# has not changed, so today's registry still describes them. +LEGACY_EVENT_PINS: dict[str, tuple[str, ...]] = { + "eicu_full_v10": _EICU_PRE_SOFA_EVENTS, + "eicu_full_L_v10": _EICU_PRE_SOFA_EVENTS, + "eicu_full_ADD_v10": _EICU_PRE_SOFA_EVENTS, + "eicu_full_DEC_v10": _EICU_PRE_SOFA_EVENTS, + "eicu_full_DEC_v11": _EICU_PRE_SOFA_EVENTS, + "eicu_full_DEC_v12": _EICU_PRE_SOFA_EVENTS, + "eicu_full_DEC_v12_steer": _EICU_PRE_SOFA_EVENTS, +} + + +def _run_name(run_dir: str | Path) -> str: + """Return the directory's final component, the key the legacy tables use.""" + return str(run_dir).rstrip("/").rsplit("/", 1)[-1] + + +def write_run_pins( + run_dir: str | Path, + *, + concept_names: Sequence[str], + event_names: Sequence[str] | None, +) -> Path: + """Record what a run trains with, in slot/head order, as ``run_pins.json``. + + Called by training right after the model is built, so the file + describes the checkpoint's actual widths rather than whatever the + registry resolves when the run is later loaded. ``event_names`` is + ``None`` for a run without hazard heads and is stored as ``null``, + which :func:`read_run_pins` reports as "no event pin". + """ + path = Path(run_dir) / PIN_FILENAME + payload = { + "concepts": list(concept_names), + "events": None if event_names is None else list(event_names), + } + path.write_text(json.dumps(payload, indent=2) + "\n") + return path + + +def read_run_pins(run_dir: str | Path) -> dict[str, tuple[str, ...]]: + """Return the pin file's non-null lists, or ``{}`` when there is none. + + A malformed file raises rather than falling through to the legacy + tables: a pin that cannot be read is a broken run directory, not an + unpinned one. + """ + path = Path(run_dir) / PIN_FILENAME + if not path.is_file(): + return {} + raw = json.loads(path.read_text()) + if not isinstance(raw, dict): + raise ValueError(f"{path}: expected a JSON object, got {type(raw).__name__}") + out: dict[str, tuple[str, ...]] = {} + for key in ("concepts", "events"): + value = raw.get(key) + if value is None: + continue + if not isinstance(value, list) or not all(isinstance(v, str) for v in value): + raise ValueError(f"{path}: {key!r} must be a list of names or null") + out[key] = tuple(value) + return out + + def pinned_concept_names(run_dir: str) -> tuple[str, ...] | None: """Return the pinned concept list for ``run_dir``, or ``None`` if unpinned. - Matches on the directory's final component, so an absolute path, a + The run's own ``run_pins.json`` wins; otherwise the legacy table, which + matches on the directory's final component, so an absolute path, a relative one and a trailing slash all resolve the same way. """ - name = str(run_dir).rstrip("/").rsplit("/", 1)[-1] - return LEGACY_CONCEPT_PINS.get(name) + from_file = read_run_pins(run_dir).get("concepts") + if from_file is not None: + return from_file + return LEGACY_CONCEPT_PINS.get(_run_name(run_dir)) + + +def pinned_event_names(run_dir: str) -> tuple[str, ...] | None: + """Return the pinned hazard event list for ``run_dir``, or ``None``. + + Same resolution as :func:`pinned_concept_names`: ``run_pins.json`` + first, then :data:`LEGACY_EVENT_PINS` by directory name. + """ + from_file = read_run_pins(run_dir).get("events") + if from_file is not None: + return from_file + return LEGACY_EVENT_PINS.get(_run_name(run_dir)) def checkpoint_num_concepts(state: dict[str, object]) -> int | None: @@ -163,6 +319,61 @@ def check_concept_count( ) +def checkpoint_num_events(state: dict[str, object]) -> int | None: + """Return how many hazard events a checkpoint's ``event_heads`` carry. + + The output layer is ``event_heads.proj.weight`` for the linear readout + and the highest-indexed ``event_heads.proj.N.weight`` for the MLP one; + its row count is ``num_events * HAZARD_NUM_BINS``. Returns ``None`` for + a checkpoint without hazard heads. + """ + weight = state.get("event_heads.proj.weight") + if weight is None: + layers = [ + k + for k in state + if k.startswith("event_heads.proj.") and k.endswith(".weight") + ] + if not layers: + return None + weight = state[max(layers, key=lambda k: int(k.split(".")[2]))] + rows = int(weight.shape[0]) # type: ignore[attr-defined] + if rows % HAZARD_NUM_BINS: + raise ValueError( + f"event_heads output width {rows} is not a multiple of " + f"{HAZARD_NUM_BINS} hazard bins; the checkpoint's bin edges differ " + "from DEFAULT_TIME_BIN_EDGES_HOURS" + ) + return rows // HAZARD_NUM_BINS + + +def check_event_count( + run_dir: str, state: dict[str, object], event_names: Sequence[str] | None +) -> None: + """Raise if today's event list disagrees with the checkpoint's head width. + + The event-head twin of :func:`check_concept_count`, for the same reason: + one sentence naming both counts and the fix, instead of a + ``load_state_dict`` size-mismatch on ``event_heads.proj.2.weight``. + """ + expected = checkpoint_num_events(state) + resolved = len(event_names or ()) + if expected is None or expected == resolved: + return + name = _run_name(run_dir) + raise ValueError( + f"{name}: checkpoint has hazard heads for {expected} events but this " + f"code resolves {resolved} for its source and task set. The alert " + f"registry has changed since the run. Pin {name}'s training-time " + f"event list: write it as the 'events' list in {name}/{PIN_FILENAME} " + "or add it to LEGACY_EVENT_PINS in " + "odyssey/inference/legacy_concept_pins.py (its env_fingerprint.json " + "records the training commit; print hazard_events_for(task_set, " + "source=source) there, in order). Loading it against today's list " + "would build a head of the wrong width." + ) + + def resolve_concepts_for_run( run_dir: str, source: str, task_set: str ) -> list[AnyConceptDefinition]: diff --git a/odyssey/inference/run_inference.py b/odyssey/inference/run_inference.py index ed69d62e..72e28bfd 100644 --- a/odyssey/inference/run_inference.py +++ b/odyssey/inference/run_inference.py @@ -43,7 +43,9 @@ from odyssey.data.vocabulary import Vocabulary, code_type from odyssey.inference.legacy_concept_pins import ( check_concept_count, + check_event_count, pinned_concept_names, + pinned_event_names, resolve_concepts_for_run, ) from odyssey.models.concept_bottleneck import ConceptBottleneckOutput @@ -88,7 +90,12 @@ compute_observability_metrics, orthogonality_diagnostic, ) -from odyssey.training.train import TrainingConfig, _move_chunk_to_device, build_model +from odyssey.training.train import ( + TrainingConfig, + _move_chunk_to_device, + build_model, + hazard_event_names_for, +) from odyssey.utils.env_fingerprint import verify_run_provenance @@ -421,11 +428,26 @@ def load_run( # exists and refuse loudly where the counts disagree and none does. pinned = pinned_concept_names(str(run_dir)) concept_names = list(pinned) if pinned is not None else [c.name for c in concepts] - model = build_model(config, vocab_size=len(vocab), num_concepts=len(concept_names)) + # Hazard heads drift the same way (alert_events_for drops events whose + # concept does not resolve for the source, and that set grows with the + # code mappings), so the event list is pinned by the same rule. + pinned_events = pinned_event_names(str(run_dir)) + event_names = ( + list(pinned_events) + if pinned_events is not None + else hazard_event_names_for(config) + ) + model = build_model( + config, + vocab_size=len(vocab), + num_concepts=len(concept_names), + event_names=event_names, + ) # After build_model, before load_state_dict: still ahead of the shape # errors this exists to replace, without pre-empting callers that stub # build_model out to inspect the reconstructed config. check_concept_count(str(run_dir), state, concept_names) + check_event_count(str(run_dir), state, event_names) model.load_state_dict(checkpoint["model"]) model = model.to(device) model.eval() diff --git a/odyssey/models/concept_bottleneck.py b/odyssey/models/concept_bottleneck.py index 83dcdd9b..ac92567f 100644 --- a/odyssey/models/concept_bottleneck.py +++ b/odyssey/models/concept_bottleneck.py @@ -67,6 +67,35 @@ class ConceptBottleneckOutput(NamedTuple): """(..., num_concepts) sigmoid(observability_logits).""" +class MixtureParts(NamedTuple): + """The mixture bottleneck's per-position pieces, before mixing. + + Everything the concept probabilities are combined with: the known + poles ``w+``/``w-`` and the unknown slot's pair and logit. Exposed so + a probe can measure what the poles carry on their own + (``scripts/probe_channel.py``); :meth:`ConceptBottleneck.forward` + mixes these same pieces. + """ + + concept_logits: torch.Tensor + """(..., num_concepts) known-concept activation logits, pre-sigmoid.""" + + known_pos: torch.Tensor + """(..., num_concepts, embedding_dim) ``w+`` per known concept.""" + + known_neg: torch.Tensor + """(..., num_concepts, embedding_dim) ``w-`` per known concept.""" + + unknown_pos: torch.Tensor + """(..., unknown_dim) the unknown slot's ``w+``.""" + + unknown_neg: torch.Tensor + """(..., unknown_dim) the unknown slot's ``w-``.""" + + unknown_logit: torch.Tensor + """(...,) the unknown slot's mixing logit, pre-sigmoid.""" + + @dataclass(frozen=True) class BottleneckIntervention: """A do()-style edit applied inside the bottleneck's mixing step. @@ -386,20 +415,17 @@ def unit_displacements(self, directions: torch.Tensor | None) -> torch.Tensor: out[i, i * d : (i + 1) * d] = directions[i] return out - def forward( - self, - hidden_states: torch.Tensor, - intervention: BottleneckIntervention | None = None, - ) -> ConceptBottleneckOutput: - """Project hidden states into known + unknown concept embeddings. + def mixture_parts(self, hidden_states: torch.Tensor) -> MixtureParts: + """Poles, logits and the unknown pair for ``hidden_states``. - ``intervention`` edits the mixing step only (see - :class:`BottleneckIntervention`): the returned - ``concept_logits``/``concept_probs``/observability outputs are - always the model's own predictions. + Call in eval mode: in train mode the dropout draw here would + differ from the forward pass's. """ - batch_shape = hidden_states.shape[:-1] - x = self.dropout(hidden_states) + return self._mixture_parts(self.dropout(hidden_states)) + + def _mixture_parts(self, x: torch.Tensor) -> MixtureParts: + """Compute :meth:`mixture_parts` on an already dropout-applied ``x``.""" + batch_shape = x.shape[:-1] k, d, u = self.num_concepts, self.embedding_dim, self.unknown_dim # Known concepts: (w+, w-) pairs and activation logits. @@ -420,7 +446,6 @@ def forward( torch.einsum("...sd,sd->...s", joint, self.prob_weight[:k]) + self.prob_bias[:k] ) - concept_probs = torch.sigmoid(concept_logits) # Unknown (residual) slot: always a context-dependent pair. u_pos, u_neg = unknown_ctx[..., 0, :], unknown_ctx[..., 1, :] @@ -429,7 +454,36 @@ def forward( else: # shared (num_slots, 2d) weight: the unknown slot is the last row u_weight, u_bias = self.prob_weight[k], self.prob_bias[k] unknown_logit = torch.cat([u_pos, u_neg], dim=-1) @ u_weight + u_bias - unknown_prob = torch.sigmoid(unknown_logit) + return MixtureParts( + concept_logits=concept_logits, + known_pos=k_pos, + known_neg=k_neg, + unknown_pos=u_pos, + unknown_neg=u_neg, + unknown_logit=unknown_logit, + ) + + def forward( + self, + hidden_states: torch.Tensor, + intervention: BottleneckIntervention | None = None, + ) -> ConceptBottleneckOutput: + """Project hidden states into known + unknown concept embeddings. + + ``intervention`` edits the mixing step only (see + :class:`BottleneckIntervention`): the returned + ``concept_logits``/``concept_probs``/observability outputs are + always the model's own predictions. + """ + batch_shape = hidden_states.shape[:-1] + x = self.dropout(hidden_states) + k, d = self.num_concepts, self.embedding_dim + parts = self._mixture_parts(x) + concept_logits = parts.concept_logits + concept_probs = torch.sigmoid(concept_logits) + k_pos, k_neg = parts.known_pos, parts.known_neg + u_pos, u_neg = parts.unknown_pos, parts.unknown_neg + unknown_prob = torch.sigmoid(parts.unknown_logit) mix_probs = concept_probs if intervention is not None and intervention.probs is not None: diff --git a/odyssey/training/train.py b/odyssey/training/train.py index e3744fce..ead2e4b4 100644 --- a/odyssey/training/train.py +++ b/odyssey/training/train.py @@ -69,6 +69,7 @@ from odyssey.data.streaming import PackedLaneSampler, StreamingChunk from odyssey.data.value_binning import CLIP_TAIL, QuantileBinner, add_value_tokens from odyssey.data.vocabulary import PAD_ID, Vocabulary +from odyssey.inference.legacy_concept_pins import write_run_pins from odyssey.models.backbones.base import TimeAwareState from odyssey.models.concept_bottleneck import ( ConceptBottleneckLossWeights, @@ -620,10 +621,40 @@ def _activate_run_sidecars(config: TrainingConfig) -> None: ) +def hazard_event_names_for(config: TrainingConfig) -> list[str] | None: + """Return the hazard-head event names ``config`` implies (``None`` without heads). + + One place for the ``event_hazards`` gate and the registry call, so + :func:`build_model` and ``load_run``'s pre-load width check agree on + what an unpinned run would be built with. + """ + if not getattr(config, "event_hazards", False): + return None + return [ + a.name + for a in hazard_events_for( + getattr(config, "task_set", "v1"), + getattr(config, "auxiliary_event_names", ()), + source=getattr(config, "source", None), + ) + ] + + def build_model( - config: TrainingConfig, *, vocab_size: int, num_concepts: int + config: TrainingConfig, + *, + vocab_size: int, + num_concepts: int, + event_names: Sequence[str] | None = None, ) -> SequenceModel: - """Construct the real backbone + heads from ``config`` (see ``model_kind``).""" + """Construct the real backbone + heads from ``config`` (see ``model_kind``). + + ``event_names`` overrides the hazard-head event list that today's + registry would give (:func:`hazard_event_names_for`); ``load_run`` + passes a checkpoint's pinned list so a run trained before the registry + grew gets heads of its own width, in its own order. It is ignored when + ``config.event_hazards`` is off. Unset, behaviour is unchanged. + """ from odyssey.models.backbones.base import SequenceBackbone # noqa: PLC0415 from odyssey.models.backbones.hybrid import EHRHybridBackbone # noqa: PLC0415 from odyssey.models.backbones.transformer import ( # noqa: PLC0415 @@ -667,18 +698,10 @@ def build_model( if getattr(config, "time_to_event", False) else None ) - event_names = ( - [ - a.name - for a in hazard_events_for( - getattr(config, "task_set", "v1"), - getattr(config, "auxiliary_event_names", ()), - source=getattr(config, "source", None), - ) - ] - if getattr(config, "event_hazards", False) - else None - ) + if event_names is None or not getattr(config, "event_hazards", False): + event_names = hazard_event_names_for(config) + else: + event_names = list(event_names) event_head_hidden = int(getattr(config, "event_head_hidden", 0) or 0) if kind == "baseline": return BaselineSequenceModel( @@ -1425,6 +1448,15 @@ def _run_training( # noqa: PLR0912, PLR0915 model = build_model(config, vocab_size=len(vocab), num_concepts=len(concepts)).to( device ) + # Self-describing run: the concept slots and hazard heads this model was + # built with, in order, so load_run never has to guess them from a + # registry that may have grown since (see legacy_concept_pins). + event_heads = getattr(model, "event_heads", None) + write_run_pins( + output_dir, + concept_names=[c.name for c in concepts], + event_names=None if event_heads is None else event_heads.event_names, + ) if config.init_from is not None: if config.resume_from is not None: raise ValueError( diff --git a/scripts/alerts_cis.py b/scripts/alerts_cis.py index 9828e8f2..d647cc72 100644 --- a/scripts/alerts_cis.py +++ b/scripts/alerts_cis.py @@ -27,6 +27,14 @@ scorer ``baseline_gbm``; both spellings are accepted here and mapped to the dump's column. +The output keeps the nested ``cells`` block (per cell: n, n_positive, +n_subjects, per-scorer AUROC/AUPRC intervals, paired deltas) and adds a +flat ``summary`` list with one record per (event, horizon) -- sample +sizes, every scorer's AUROC point estimate, and the paired AUROC delta +with its interval -- plus the row and subject counts of the dumps before +and after any ``--max-subjects`` subsample. On GEMINI, ``scripts/gemini/ +run.sh alerts-cis `` wraps this script and exports that file. + Runtime notes from the full-held-out runs: this script imports sklearn (for AUPRC), so on the GEMINI node it needs the GPU venv, not the lightweight one. A 1000-draw subject bootstrap over ~30M index rows takes 15-20 h @@ -119,12 +127,19 @@ def score_cell( return None y = sub[f"y@{h}"].to_numpy().astype(np.float64) subj = sub["subject_id"].to_numpy() + n_subjects = int(len(np.unique(subj))) if len(np.unique(y)) < 2: - return {"n": int(len(y)), "n_positive": int(y.sum()), "unscoreable": True} + return { + "n": int(len(y)), + "n_positive": int(y.sum()), + "n_subjects": n_subjects, + "unscoreable": True, + } result: dict[str, Any] = { "n": int(len(y)), "n_positive": int(y.sum()), + "n_subjects": n_subjects, "scorers": {}, "paired_deltas": {}, } @@ -184,6 +199,47 @@ def subsample_subjects( return frame.filter(pl.col("subject_id").is_in(keep.tolist())) +def summary_rows(cells: dict[str, Any], scorers: list[str]) -> list[dict[str, Any]]: + """Flatten ``cells`` into one record per (event, horizon). + + Each record carries the sample sizes, every scorer's AUROC point + estimate, and the paired AUROC delta of the first scorer minus each + other scorer with its 95% interval -- the table-ready view of the + nested ``cells`` block (which keeps AUPRC and the per-scorer + intervals as well). + """ + rows: list[dict[str, Any]] = [] + for key, cell in cells.items(): + event, _, horizon = key.rpartition("@") + row: dict[str, Any] = { + "event": event, + "horizon_hours": float(horizon[:-1]), + "n_at_risk": cell["n"], + "n_positive": cell["n_positive"], + "n_subjects": cell.get("n_subjects"), + "unscoreable": bool(cell.get("unscoreable", False)), + } + for s in scorers: + auroc = cell.get("scorers", {}).get(s, {}).get("auroc") + row[f"{s}_auroc"] = auroc["point"] if auroc else None + ref = scorers[0] + for s in scorers[1:]: + delta = cell.get("paired_deltas", {}).get(f"{ref}_minus_{s}", {}) + auroc = delta.get("auroc") + row[f"{ref}_minus_{s}_auroc"] = ( + { + "point": auroc["point"], + "ci_low": auroc["ci_low"], + "ci_high": auroc["ci_high"], + "separated": auroc["separated"], + } + if auroc + else None + ) + rows.append(row) + return rows + + def main() -> None: """Compute per-cell CIs and paired deltas from alerts row dumps.""" parser = argparse.ArgumentParser(description=__doc__) @@ -208,12 +264,29 @@ def main() -> None: "max_subjects": args.max_subjects, "variance_scope": "finite-sample only (single fitted model); refit " "variance requires seed replicates", + "n_rows_in_dumps": 0, + "n_subjects_in_dumps": 0, + "n_rows_scored": 0, + "n_subjects_scored": 0, "cells": {}, + "summary": [], } for path in args.dump: - frame = subsample_subjects( - pl.read_parquet(path), max_subjects=args.max_subjects, seed=args.seed + full = pl.read_parquet(path) + frame = subsample_subjects(full, max_subjects=args.max_subjects, seed=args.seed) + out["n_rows_in_dumps"] += full.height + out["n_subjects_in_dumps"] += int(full["subject_id"].n_unique()) + out["n_rows_scored"] += frame.height + out["n_subjects_scored"] += int(frame["subject_id"].n_unique()) + logger.info( + "%s: %d rows / %d subjects in the dump, %d rows / %d subjects scored", + path, + full.height, + full["subject_id"].n_unique(), + frame.height, + frame["subject_id"].n_unique(), ) + del full events = ( frame["event"].unique().to_list() if "event" in frame.columns else [None] ) @@ -242,6 +315,7 @@ def main() -> None: if v["auroc"] and v["auroc"]["ci_low"] is not None ), ) + out["summary"] = summary_rows(out["cells"], args.scorers) with open(args.output_json, "w") as f: json.dump(out, f, indent=1) logger.info("wrote %s (%d cells)", args.output_json, len(out["cells"])) diff --git a/scripts/cohort_counts.py b/scripts/cohort_counts.py new file mode 100644 index 00000000..52dcce9c --- /dev/null +++ b/scripts/cohort_counts.py @@ -0,0 +1,598 @@ +"""Aggregate cohort description of one MEDS source, split by split. + +For every split directory (train, tuning, held_out by default) this +reports subjects, admissions, hospitals (when the MEDS metadata carries +a site code), the calendar year range of admissions, length of stay, +sex and age at admission where the source charts them, and the share of +subjects (and admissions) with at least one onset of each hazard event, +using the SAME onset definitions the alerts leg scores +(:func:`odyssey.data.alert_events.alert_events_for` resolved for the +source, :func:`odyssey.data.alert_events.all_event_times`). Nothing +patient-level is kept: every count below :data:`SUPPRESS_BELOW` is +written as ``"<10"`` and the statistics that depend on it are dropped, +so the JSON can leave a closed environment such as GEMINI. + +Definitions: + +- A subject is a distinct ``subject_id``; shards partition subjects, so + per-shard counts add up. +- An admission is a distinct ``(subject_id, hadm_id)`` with at least one + timed event. Its start is the source's admission code + (``HOSPITAL_ADMISSION//`` on MIMIC-IV and eICU, bare ``ADMISSION`` on + GEMINI) when the visit carries one, else the visit's first timed + event; its end is the discharge code, else the last timed event. + Length of stay is end minus start in days; the year range is over + admission starts. +- Sex comes from ``GENDER//`` static rows; age at admission from + ``MEDS_BIRTH``. A source without those rows (GEMINI extracts neither) + reports them as not available rather than as zero. +- Event prevalence per subject is the share of the split's subjects + with at least one onset. For ``readmission_30d`` the onset must fall + within 720 h of the visit's last event, matching the 30-day horizon + the discharge-anchored scorer uses. + +Usage:: + + uv run python scripts/cohort_counts.py --source gemini --task-set v3 \ + --meds-dir /path/to/gemini_meds_v1 \ + [--hadm-hospital-parquet /path/to/metadata/hadm_id_hospital.parquet] \ + --output-json ~/runs//cohort_counts.json + + uv run python scripts/cohort_counts.py --source mimic_iv --task-set v3 \ + --split train=/data/mimic/train --split held_out=/data/mimic/held_out \ + --output-json cohort_counts.json + +``--max-shards`` bounds every split (recorded in the output, so a subset +run can never pass for a full one); ``--no-events`` skips the label pass +and writes only the subject/admission/LOS/year/sex/age block. +""" + +from __future__ import annotations + +import argparse +import json +import logging +from collections import Counter +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any + +import numpy as np +import polars as pl + +from odyssey.data.alert_events import ( + ALERT_TASK_SETS, + AlertEvent, + alert_events_for, + all_event_times, + visit_envelope, +) +from odyssey.data.code_normalization import maybe_normalize +from odyssey.data.sequences import BIRTH_CODE, HOURS_PER_YEAR +from odyssey.data.sidecars import activate_sidecars +from odyssey.training.data import load_meds_shard +from odyssey.training.shard_stream import shard_paths + + +logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") +logger = logging.getLogger("cohort_counts") + +#: Counts below this are written as ``"<10"`` and any statistic over +#: fewer than this many records is dropped. +SUPPRESS_BELOW = 10 +SUPPRESSED = f"<{SUPPRESS_BELOW}" + +DEFAULT_SPLITS: tuple[str, ...] = ("train", "tuning", "held_out") + +#: The source's hospital admission / discharge code families. The +#: admission prefixes are the ones odyssey.data.alert_events resolves the +#: readmission_30d event to per source (the shared MIMIC/eICU definition +#: and GEMINI's bare-token override); discharge has no event of its own. +ADMISSION_PREFIX: dict[str, str] = { + "mimic_iv": "HOSPITAL_ADMISSION//", + "eicu": "HOSPITAL_ADMISSION//", + "gemini": "ADMISSION", +} +DISCHARGE_PREFIX: dict[str, str] = { + "mimic_iv": "HOSPITAL_DISCHARGE//", + "eicu": "HOSPITAL_DISCHARGE//", + "gemini": "DISCHARGE", +} +SEX_PREFIX = "GENDER//" + +#: The 30-day readmission window, in hours: the longest horizon of +#: odyssey.inference.alerts.READMISSION_HORIZONS_HOURS (kept as a literal +#: here so this script does not import torch through that module; a test +#: pins the two together). +READMISSION_WINDOW_HOURS = 720.0 + +LOS_DEFINITION = ( + "per admission (subject, visit id): admission code time else first " + "timed event, to discharge code time else last timed event, in days; " + "negative spans dropped" +) + +#: Top-level keys of the JSON this script writes. run.sh's export +#: whitelist for the cohort-counts step must list exactly these. +OUTPUT_KEYS: tuple[str, ...] = ( + "source", + "task_set", + "normalize_medications", + "suppression_threshold", + "max_shards", + "los_definition", + "hospital_metadata", + "events", + "events_dropped", + "splits", + "all_splits", +) + + +@dataclass +class SplitAccumulator: + """Per-split running aggregates; merges across shards and across splits.""" + + n_shards_read: int = 0 + n_shards_total: int = 0 + n_subjects: int = 0 + n_admissions: int = 0 + hospitals: set[str] = field(default_factory=set) + year_min: int | None = None + year_max: int | None = None + los_days: list[np.ndarray] = field(default_factory=list) + sex: Counter[str] = field(default_factory=Counter) + n_subjects_with_sex: int = 0 + ages: list[np.ndarray] = field(default_factory=list) + event_subjects: Counter[str] = field(default_factory=Counter) + event_admissions: Counter[str] = field(default_factory=Counter) + + def merge(self, other: SplitAccumulator) -> None: + """Fold ``other`` into this accumulator.""" + self.n_shards_read += other.n_shards_read + self.n_shards_total += other.n_shards_total + self.n_subjects += other.n_subjects + self.n_admissions += other.n_admissions + self.hospitals |= other.hospitals + for y in (other.year_min, other.year_max): + if y is None: + continue + self.year_min = y if self.year_min is None else min(self.year_min, y) + self.year_max = y if self.year_max is None else max(self.year_max, y) + self.los_days.extend(other.los_days) + self.sex.update(other.sex) + self.n_subjects_with_sex += other.n_subjects_with_sex + self.ages.extend(other.ages) + self.event_subjects.update(other.event_subjects) + self.event_admissions.update(other.event_admissions) + + +def suppress(n: int) -> int | str: + """Return ``n`` or the suppression marker when it is a small cell.""" + return n if n >= SUPPRESS_BELOW else SUPPRESSED + + +def _quartiles(chunks: list[np.ndarray]) -> dict[str, Any] | None: + values = np.concatenate(chunks) if chunks else np.empty(0) + values = values[np.isfinite(values)] + if len(values) < SUPPRESS_BELOW: + return None + q1, med, q3 = np.percentile(values, [25, 50, 75]) + return { + "n": int(len(values)), + "median": round(float(med), 2), + "q1": round(float(q1), 2), + "q3": round(float(q3), 2), + } + + +def _visits(events: pl.DataFrame, source: str) -> pl.DataFrame: + """One row per admission with its start and end instants.""" + timed = events.filter( + pl.col("time").is_not_null() + & pl.col("hadm_id").is_not_null() + & (pl.col("code") != BIRTH_CODE) + ) + adm, dis = ADMISSION_PREFIX[source], DISCHARGE_PREFIX[source] + visits = timed.group_by("subject_id", "hadm_id").agg( + pl.col("time").min().alias("_first"), + pl.col("time").max().alias("_last"), + pl.col("time").filter(pl.col("code").str.starts_with(adm)).min().alias("_adm"), + pl.col("time").filter(pl.col("code").str.starts_with(dis)).max().alias("_dis"), + ) + return visits.select( + "subject_id", + "hadm_id", + pl.coalesce(pl.col("_adm"), pl.col("_first")).alias("start"), + pl.coalesce(pl.col("_dis"), pl.col("_last")).alias("end"), + ) + + +def _readmission_within_horizon( + onset: dict[tuple[int, int], float], + envelope: dict[tuple[int, int], tuple[float, float]], + horizon: float, +) -> dict[tuple[int, int], float]: + """Keep only next-admission onsets within ``horizon`` of the visit's end.""" + return { + key: t + for key, t in onset.items() + if key in envelope and t - envelope[key][1] <= horizon + } + + +def _event_counts( + events: pl.DataFrame, + alerts: tuple[AlertEvent, ...], + source: str, + task_set: str, +) -> tuple[Counter[str], Counter[str]]: + """Subjects and admissions with at least one onset, per event.""" + times = all_event_times(events, alerts, source, task_set=task_set) + envelope = None + per_subject: Counter[str] = Counter() + per_admission: Counter[str] = Counter() + for alert in alerts: + onset = times[alert.name].onset + if alert.next_visit: + if envelope is None: + envelope = visit_envelope(events) + onset = _readmission_within_horizon( + onset, envelope, READMISSION_WINDOW_HOURS + ) + per_subject[alert.name] = len({s for s, _ in onset}) + if not times[alert.name].subject_scoped: + per_admission[alert.name] = len(onset) + return per_subject, per_admission + + +def summarize_shard( + events: pl.DataFrame, + *, + source: str, + task_set: str, + alerts: tuple[AlertEvent, ...] | None, + normalize_medications: bool, + hospital_by_hadm: pl.DataFrame | None, +) -> SplitAccumulator: + """Aggregate one shard's events into a fresh :class:`SplitAccumulator`.""" + if "hadm_id" not in events.columns: + events = events.with_columns(pl.lit(None, dtype=pl.Int64).alias("hadm_id")) + acc = SplitAccumulator(n_shards_read=1) + acc.n_subjects = int(events["subject_id"].n_unique()) + + visits = _visits(events, source) + acc.n_admissions = visits.height + if visits.height: + years = visits["start"].dt.year() + acc.year_min = int(years.min()) # type: ignore[arg-type] + acc.year_max = int(years.max()) # type: ignore[arg-type] + los = (visits["end"] - visits["start"]).dt.total_seconds().to_numpy() / 86400.0 + acc.los_days.append(los[los >= 0].astype(np.float32)) + if hospital_by_hadm is not None: + hit = visits.select(pl.col("hadm_id").cast(pl.Utf8)).join( + hospital_by_hadm, on="hadm_id", how="inner" + ) + acc.hospitals = set(hit["hospital_num"].unique().to_list()) + + sex_rows = events.filter(pl.col("code").str.starts_with(SEX_PREFIX)) + if sex_rows.height: + first = sex_rows.group_by("subject_id").agg(pl.col("code").first()) + categories = first["code"].str.slice(len(SEX_PREFIX)).to_list() + acc.sex.update(categories) + acc.n_subjects_with_sex = first.height + + births = events.filter(pl.col("code") == BIRTH_CODE) + if births.height and visits.height: + birth_by_subject = births.group_by("subject_id").agg( + pl.col("time").min().alias("_birth") + ) + aged = visits.join(birth_by_subject, on="subject_id", how="inner").filter( + pl.col("_birth").is_not_null() + ) + if aged.height: + hours = (aged["start"] - aged["_birth"]).dt.total_seconds().to_numpy() + acc.ages.append((hours / 3600.0 / HOURS_PER_YEAR).astype(np.float32)) + + if alerts: + prepared = maybe_normalize(events, enabled=normalize_medications, source=source) + acc.event_subjects, acc.event_admissions = _event_counts( + prepared, alerts, source, task_set + ) + return acc + + +def summarize_split( + shard_dir: Path, + *, + source: str, + task_set: str, + alerts: tuple[AlertEvent, ...] | None, + normalize_medications: bool, + hospital_by_hadm: pl.DataFrame | None, + max_shards: int | None, +) -> SplitAccumulator: + """Aggregate every shard of one split directory.""" + paths = shard_paths(shard_dir) + total = SplitAccumulator(n_shards_total=len(paths)) + if max_shards is not None: + paths = paths[:max_shards] + if alerts: + activate_sidecars(shard_dir) + for i, path in enumerate(paths, 1): + acc = summarize_shard( + load_meds_shard(path), + source=source, + task_set=task_set, + alerts=alerts, + normalize_medications=normalize_medications, + hospital_by_hadm=hospital_by_hadm, + ) + total.merge(acc) + if i % 25 == 0 or i == len(paths): + logger.info( + "%s: %d/%d shards, %d subjects, %d admissions", + shard_dir.name, + i, + len(paths), + total.n_subjects, + total.n_admissions, + ) + return total + + +def _rate(numerator: int, denominator: int) -> float | None: + if numerator < SUPPRESS_BELOW or denominator < SUPPRESS_BELOW: + return None + return round(numerator / denominator, 4) + + +def report_split( + acc: SplitAccumulator, + *, + alerts: tuple[AlertEvent, ...] | None, + hospitals_available: bool, +) -> dict[str, Any]: + """Suppressed, JSON-ready view of one accumulator.""" + enough = acc.n_admissions >= SUPPRESS_BELOW + out: dict[str, Any] = { + "n_shards_read": acc.n_shards_read, + "n_shards_total": acc.n_shards_total, + "n_subjects": suppress(acc.n_subjects), + "n_admissions": suppress(acc.n_admissions), + "n_hospitals": len(acc.hospitals) if hospitals_available else None, + "admission_years": ( + {"min": acc.year_min, "max": acc.year_max} if enough else None + ), + "los_days": _quartiles(acc.los_days), + "sex": None, + "age_years": _quartiles(acc.ages), + "events": None, + } + if not hospitals_available: + out["hospitals_note"] = "not available: no site code in the MEDS metadata" + if acc.n_subjects_with_sex: + out["sex"] = { + "n_subjects_with_sex": suppress(acc.n_subjects_with_sex), + "counts": {k: suppress(v) for k, v in sorted(acc.sex.items())}, + } + else: + out["sex_note"] = "not available: no GENDER// rows in this source" + if not acc.ages: + out["age_note"] = "not available: no MEDS_BIRTH rows in this source" + if alerts: + out["events"] = { + a.name: { + "n_subjects_positive": suppress(acc.event_subjects[a.name]), + "prevalence_per_subject": _rate( + acc.event_subjects[a.name], acc.n_subjects + ), + "n_admissions_positive": ( + None if a.subject_scoped else suppress(acc.event_admissions[a.name]) + ), + "prevalence_per_admission": ( + None + if a.subject_scoped + else _rate(acc.event_admissions[a.name], acc.n_admissions) + ), + } + for a in alerts + } + return out + + +def format_table(report: dict[str, Any]) -> str: + """Plain-text summary of the report for the log.""" + lines = [f"source: {report['source']} task_set: {report['task_set']}"] + lines.append( + f"{'split':<12}{'subjects':>12}{'admissions':>12}{'hospitals':>10}" + f"{'years':>12}{'LOS med (IQR) d':>24}" + ) + for name, s in [*report["splits"].items(), ("all", report["all_splits"])]: + years = s["admission_years"] + los = s["los_days"] + hospitals = "n/a" if s["n_hospitals"] is None else str(s["n_hospitals"]) + year_text = f"{years['min']}-{years['max']}" if years else "n/a" + los_text = f"{los['median']} ({los['q1']}-{los['q3']})" if los else "n/a" + lines.append( + f"{name:<12}{s['n_subjects']!s:>12}{s['n_admissions']!s:>12}" + f"{hospitals:>10}{year_text:>12}{los_text:>24}" + ) + for name, s in report["splits"].items(): + if not s["events"]: + continue + lines.append(f"\n{name}: prevalence per subject") + for event, e in s["events"].items(): + rate = e["prevalence_per_subject"] + rate_text = "suppressed" if rate is None else f"{rate:.4f}" + lines.append( + f" {event:<24}{e['n_subjects_positive']!s:>10}{rate_text:>12}" + ) + return "\n".join(lines) + + +def _apply_run_config(args: argparse.Namespace) -> dict[str, Path]: + """Fill unset options from a training run's ``config.json``. + + Returns the run's train/tuning split directories (those that exist) + so the report describes the shards the run actually trained on. + """ + if not args.run_dir: + return {} + config = json.loads((Path(args.run_dir) / "config.json").read_text()) + if args.source is None: + args.source = config.get("source", "mimic_iv") + if args.task_set is None: + args.task_set = config.get("task_set", "v1") + if not args.normalize_medications and config.get("normalize_medications"): + args.normalize_medications = True + splits: dict[str, Path] = {} + for name, key in (("train", "train_shard_dir"), ("tuning", "tuning_shard_dir")): + path = config.get(key) + if path and Path(path).is_dir(): + splits[name] = Path(path) + elif path: + logger.warning("run config's %s %s does not exist; skipping", key, path) + return splits + + +def _parse_splits(args: argparse.Namespace) -> dict[str, Path]: + splits: dict[str, Path] = _apply_run_config(args) + if args.meds_dir: + root = Path(args.meds_dir) / "data" + for name in DEFAULT_SPLITS: + if (root / name).is_dir(): + splits[name] = root / name + for spec in args.split or []: + name, _, path = spec.partition("=") + if not path: + raise SystemExit(f"--split expects NAME=DIR, got {spec!r}") + splits[name] = Path(path) + if not splits: + raise SystemExit("no split directories: pass --meds-dir or --split NAME=DIR") + return splits + + +def build_report( + splits: dict[str, Path], + *, + source: str, + task_set: str, + normalize_medications: bool, + with_events: bool, + hospital_parquet: Path | None, + max_shards: int | None, +) -> dict[str, Any]: + """Compute the full cohort report over ``splits``.""" + alerts = alert_events_for(task_set, source=source) if with_events else None + kept = {a.name for a in alerts} if alerts is not None else set() + dropped = ( + [a.name for a in ALERT_TASK_SETS[task_set] if a.name not in kept] + if alerts is not None + else [] + ) + hospital_by_hadm = None + if hospital_parquet is not None: + hospital_by_hadm = pl.read_parquet(hospital_parquet).select( + pl.col("hadm_id").cast(pl.Utf8), + pl.col("hospital_num").cast(pl.Utf8), + ) + accs: dict[str, SplitAccumulator] = {} + for name, shard_dir in splits.items(): + logger.info("split %s: %s", name, shard_dir) + accs[name] = summarize_split( + shard_dir, + source=source, + task_set=task_set, + alerts=alerts, + normalize_medications=normalize_medications, + hospital_by_hadm=hospital_by_hadm, + max_shards=max_shards, + ) + overall = SplitAccumulator() + for acc in accs.values(): + overall.merge(acc) + available = hospital_by_hadm is not None + return { + "source": source, + "task_set": task_set, + "normalize_medications": normalize_medications, + "suppression_threshold": SUPPRESS_BELOW, + "max_shards": max_shards, + "los_definition": LOS_DEFINITION, + "hospital_metadata": (hospital_parquet.name if hospital_parquet else None), + "events": [a.name for a in alerts] if alerts else [], + "events_dropped": dropped, + "splits": { + name: report_split(acc, alerts=alerts, hospitals_available=available) + for name, acc in accs.items() + }, + "all_splits": report_split( + overall, alerts=alerts, hospitals_available=available + ), + } + + +def main() -> None: + """Compute the suppressed cohort description and write it as JSON.""" + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--run-dir", + default=None, + help="training run directory; its config.json supplies --source, " + "--task-set, --normalize-medications and the train/tuning split " + "directories unless given explicitly", + ) + parser.add_argument("--source", choices=sorted(ADMISSION_PREFIX), default=None) + parser.add_argument("--task-set", default=None, choices=sorted(ALERT_TASK_SETS)) + parser.add_argument( + "--meds-dir", default=None, help="MEDS root with data// directories" + ) + parser.add_argument( + "--split", + action="append", + default=None, + help="NAME=DIR; adds to or overrides the --meds-dir splits", + ) + parser.add_argument( + "--hadm-hospital-parquet", + default=None, + help="hadm_id -> hospital_num table (GEMINI metadata/hadm_id_hospital.parquet)", + ) + parser.add_argument( + "--normalize-medications", + action="store_true", + help="apply the run's medication normalizer before labeling (MIMIC/eICU runs)", + ) + parser.add_argument("--max-shards", type=int, default=None) + parser.add_argument( + "--no-events", action="store_true", help="skip the per-event label pass" + ) + parser.add_argument("--output-json", required=True) + args = parser.parse_args() + splits = _parse_splits(args) + if args.source is None: + raise SystemExit("--source is required without --run-dir") + if args.task_set is None: + args.task_set = "v3" + + report = build_report( + splits, + source=args.source, + task_set=args.task_set, + normalize_medications=args.normalize_medications, + with_events=not args.no_events, + hospital_parquet=( + Path(args.hadm_hospital_parquet) if args.hadm_hospital_parquet else None + ), + max_shards=args.max_shards, + ) + print(format_table(report)) + Path(args.output_json).parent.mkdir(parents=True, exist_ok=True) + with open(args.output_json, "w") as f: + json.dump(report, f, indent=1) + logger.info("wrote %s", args.output_json) + + +if __name__ == "__main__": + main() diff --git a/scripts/compare_runs_paired.py b/scripts/compare_runs_paired.py new file mode 100644 index 00000000..ae2abd83 --- /dev/null +++ b/scripts/compare_runs_paired.py @@ -0,0 +1,549 @@ +"""Paired hazard-head AUROC deltas between two runs on identical landmark rows. + +The WP2 question (docs/ml4h2026_rebuttal_plan.md): what does the concept +bottleneck cost? The answer is a PAIRED difference between two runs of +the same backbone, one with the bottleneck (``full_run_v10``) and one +without (``full_run_baseline_v10``, ``model_kind=baseline``), scored by +the same alert chain on the same held-out shards. Each run's chain +writes a per-index-row dump (``alerts_rows.parquet``); this script +inner-joins the two dumps on the row key, checks that the joined row +set is the row set of BOTH dumps, and hands the aligned columns to +:func:`odyssey.inference.uncertainty.bootstrap_auroc_delta` (subject- +clustered, paired). The refusal on a row-set mismatch is deliberate: +numbers from different row sets in one table is the coverage-mismatch +bug this project has hit three times, so both counts and the unmatched +count are printed and the script stops. + +Row key: ``(event, subject_id, visit_id, time_hours)``, the columns +:func:`odyssey.inference.alerts.index_row_table` writes. Score columns +are ``{scorer}@{h}h`` (``hazard`` for the model's heads, ``gbm`` for the +per-run GBM refit); the label column is ``y@{h}h`` with null meaning +censored before the horizon. + +The GBM is refit inside each run's chain, so the two dumps carry two +different GBMs on the same rows. The ``gbm`` block reports both AUROCs +and their difference WITHOUT a bootstrap: that difference is refit +variance, and it is reported so the reader can put the hazard-head +delta next to it. + +Usage:: + + uv run python scripts/compare_runs_paired.py \ + --dump-a ~/runs/full_run_v10/alerts_rows.parquet \ + --dump-b ~/runs/full_run_baseline_v10/alerts_rows.parquet \ + --label-a bottleneck --label-b baseline \ + [--inference-a ~/runs/full_run_v10/inference_results.json] \ + [--inference-b ~/runs/full_run_baseline_v10/inference_results.json] \ + [--scorer hazard] [--events aki_stage_3 ...] [--horizons 8 24 72] \ + [--n-boot 1000] [--seed 0] [--max-subjects N] \ + --output-json ~/runs/full_run_baseline_v10/paired_vs_v10.json + +Intervals carry finite-sample variance only (one fitted model per arm, +one held-out draw), exactly as in ``scripts/alerts_cis.py``. +""" + +from __future__ import annotations + +import argparse +import json +import logging +import sys +from pathlib import Path +from typing import Any + +import numpy as np +import polars as pl +from sklearn.metrics import roc_auc_score + +from odyssey.inference.uncertainty import bootstrap_auroc_delta +from scripts.alerts_cis import SCORER_ALIASES, horizons_in, subsample_subjects + + +logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s") +logger = logging.getLogger("compare_runs_paired") + +#: Columns that identify one landmark row across two dumps of the same +#: held-out shards (``event`` is added per event when joining). +ROW_KEY: tuple[str, ...] = ("subject_id", "visit_id", "time_hours") + +#: Largest tolerated fraction of either dump's rows (per event) that the +#: inner join may leave unmatched before the script refuses. +MAX_UNMATCHED_FRACTION = 0.001 + +#: ``task_metrics`` keys of ``inference_results.json`` reported side by side. +INFERENCE_KEYS: tuple[str, ...] = ( + "set_top1_accuracy", + "top1_accuracy", + "top5_accuracy", + "cross_entropy", + "n_predictions", + "n_set_predictions", +) + + +class RowSetMismatchError(RuntimeError): + """The two dumps do not describe the same landmark rows.""" + + +def load_dump(path: str | Path) -> pl.DataFrame: + """Read a row dump and normalise the key columns for an exact join.""" + frame = pl.read_parquet(path) + missing = [c for c in (*ROW_KEY, "event") if c not in frame.columns] + if missing: + raise ValueError(f"{path}: dump is missing key columns {missing}") + version = ( + frame["landmark_protocol_version"][0] + if "landmark_protocol_version" in frame.columns and frame.height + else None + ) + logger.info( + "%s: %d rows / %d subjects, landmark protocol %s", + path, + frame.height, + frame["subject_id"].n_unique(), + version, + ) + return frame.with_columns( + pl.col("subject_id").cast(pl.Int64), + pl.col("visit_id").cast(pl.Int64), + pl.col("time_hours").cast(pl.Float64), + pl.col("event").cast(pl.Utf8), + ) + + +def _protocol_version(frame: pl.DataFrame) -> Any: + if "landmark_protocol_version" not in frame.columns or frame.height == 0: + return None + return frame["landmark_protocol_version"][0] + + +def join_event( + a: pl.DataFrame, b: pl.DataFrame, event: str, *, label_a: str, label_b: str +) -> tuple[pl.DataFrame, dict[str, int]]: + """Inner-join one event's rows of both dumps on :data:`ROW_KEY`. + + Refuses (:class:`RowSetMismatchError`) when a key repeats inside either + dump, or when the join leaves more than :data:`MAX_UNMATCHED_FRACTION` + of either dump's rows unmatched. Columns of ``b`` get a ``_b`` + suffix; ``a`` keeps its names. + """ + ev_a = a.filter(pl.col("event") == event) + ev_b = b.filter(pl.col("event") == event) + keys = list(ROW_KEY) + for name, frame in ((label_a, ev_a), (label_b, ev_b)): + n_dup = frame.height - frame.select(keys).n_unique() + if n_dup: + raise RowSetMismatchError( + f"{event}: {name} dump has {n_dup} duplicate row keys " + f"{keys}; the join would multiply rows" + ) + non_key_b = [c for c in ev_b.columns if c not in (*keys, "event")] + joined = ev_a.join( + ev_b.select([*keys, *non_key_b]).rename({c: f"{c}_b" for c in non_key_b}), + on=keys, + how="inner", + ) + counts = { + "n_rows_a": ev_a.height, + "n_rows_b": ev_b.height, + "n_rows_joined": joined.height, + "n_unmatched_a": ev_a.height - joined.height, + "n_unmatched_b": ev_b.height - joined.height, + "n_subjects_joined": int(joined["subject_id"].n_unique()), + } + logger.info( + "%s: %s n=%d, %s n=%d, joined n=%d (unmatched %d / %d), n_subjects=%d", + event, + label_a, + counts["n_rows_a"], + label_b, + counts["n_rows_b"], + counts["n_rows_joined"], + counts["n_unmatched_a"], + counts["n_unmatched_b"], + counts["n_subjects_joined"], + ) + for name, n_rows, n_unmatched in ( + (label_a, ev_a.height, counts["n_unmatched_a"]), + (label_b, ev_b.height, counts["n_unmatched_b"]), + ): + if n_rows == 0 or n_unmatched / n_rows > MAX_UNMATCHED_FRACTION: + raise RowSetMismatchError( + f"{event}: row sets differ. {label_a} has {counts['n_rows_a']} " + f"rows, {label_b} has {counts['n_rows_b']} rows, the join keeps " + f"{counts['n_rows_joined']}; {n_unmatched} of {name}'s {n_rows} " + f"rows are unmatched (limit {MAX_UNMATCHED_FRACTION:.1%}). The " + "two dumps must be scored on identical held-out shards under " + "the same landmark protocol." + ) + return joined, counts + + +def _point_auroc(y: np.ndarray, p: np.ndarray) -> float | None: + if len(np.unique(y)) < 2: + return None + return float(roc_auc_score(y, p)) + + +def score_pair( + joined: pl.DataFrame, + scorer: str, + horizon: float, + *, + n_boot: int, + seed: int, +) -> dict[str, Any] | None: + """Paired AUROC delta (b minus a) for one joined event frame and horizon. + + Rows are those with a label at the horizon in dump A, the same label + in dump B, and a non-null score in both arms; a label disagreement is + refused because the two chains ran the same label rule on the same + events. Returns ``None`` when no row qualifies. + """ + h = f"{horizon:g}h" + y_col, s_col = f"y@{h}", f"{scorer}@{h}" + needed = [y_col, s_col, f"{y_col}_b", f"{s_col}_b"] + missing = [c for c in needed if c not in joined.columns] + if missing: + logger.warning("horizon %s: missing columns %s, skipped", h, missing) + return None + both_labelled = joined.filter( + pl.col(y_col).is_not_null() & pl.col(f"{y_col}_b").is_not_null() + ) + n_disagree = int((both_labelled[y_col] != both_labelled[f"{y_col}_b"]).sum()) + if n_disagree: + raise RowSetMismatchError( + f"horizon {h}: {n_disagree} rows carry different labels in the two " + "dumps; the chains did not run the same label rule" + ) + one_sided = int((joined[y_col].is_null() != joined[f"{y_col}_b"].is_null()).sum()) + sub = both_labelled.filter( + pl.col(s_col).is_not_null() & pl.col(f"{s_col}_b").is_not_null() + ) + if sub.height == 0: + return None + y = sub[y_col].to_numpy().astype(np.float64) + p_a = sub[s_col].to_numpy().astype(np.float64) + p_b = sub[f"{s_col}_b"].to_numpy().astype(np.float64) + subj = sub["subject_id"].to_numpy() + result: dict[str, Any] = { + "n_at_risk": int(len(y)), + "n_positive": int(y.sum()), + "n_subjects": int(len(np.unique(subj))), + "n_label_one_sided": one_sided, + "auroc_a": _point_auroc(y, p_a), + "auroc_b": _point_auroc(y, p_b), + "delta_b_minus_a": None, + } + delta = bootstrap_auroc_delta(y, p_b, p_a, subj, n_boot=n_boot, seed=seed) + if delta is None: + result["unscoreable"] = True + return result + result["delta_b_minus_a"] = { + "point": delta.point_estimate, + "ci_low": delta.ci_low, + "ci_high": delta.ci_high, + "separated": delta.excludes_zero(), + "n_boot_used": delta.n_boot_used, + "n_boot_skipped": delta.n_boot_skipped, + } + return result + + +def gbm_pair(joined: pl.DataFrame, horizon: float) -> dict[str, Any] | None: + """Both dumps' GBM AUROCs on the joined rows, and their difference. + + No bootstrap: the two GBMs are separate refits, so the difference is + refit variance, reported for scale only. + """ + h = f"{horizon:g}h" + y_col, g_col = f"y@{h}", f"gbm@{h}" + if g_col not in joined.columns or f"{g_col}_b" not in joined.columns: + return None + sub = joined.filter( + pl.col(y_col).is_not_null() + & pl.col(g_col).is_not_null() + & pl.col(f"{g_col}_b").is_not_null() + ) + if sub.height == 0: + return None + y = sub[y_col].to_numpy().astype(np.float64) + auroc_a = _point_auroc(y, sub[g_col].to_numpy().astype(np.float64)) + auroc_b = _point_auroc(y, sub[f"{g_col}_b"].to_numpy().astype(np.float64)) + return { + "n_at_risk": int(len(y)), + "auroc_a": auroc_a, + "auroc_b": auroc_b, + "delta_b_minus_a": ( + None if auroc_a is None or auroc_b is None else auroc_b - auroc_a + ), + } + + +def read_inference(path: str | Path | None) -> dict[str, Any] | None: + """Read the ``task_metrics`` keys of an ``inference_results.json``.""" + if path is None: + return None + p = Path(path) + if not p.exists(): + logger.warning("%s: no inference_results.json, skipped", p) + return None + with open(p) as f: + task = json.load(f).get("task_metrics", {}) + return {k: task.get(k) for k in INFERENCE_KEYS} + + +def inference_block( + inf_a: dict[str, Any] | None, inf_b: dict[str, Any] | None +) -> dict[str, Any] | None: + """Side-by-side next-event metrics with b minus a on each numeric key.""" + if inf_a is None and inf_b is None: + return None + out: dict[str, Any] = {"a": inf_a, "b": inf_b, "delta_b_minus_a": {}} + for k in INFERENCE_KEYS: + va = (inf_a or {}).get(k) + vb = (inf_b or {}).get(k) + out["delta_b_minus_a"][k] = ( + None + if va is None or vb is None or k.startswith("n_") + else float(vb) - float(va) + ) + return out + + +def summarise(cells: dict[str, Any]) -> dict[str, list[str]]: + """Cells where b beats a, a beats b, or the paired interval covers 0.""" + out: dict[str, list[str]] = {"b_beats_a": [], "a_beats_b": [], "ties": []} + for key, cell in cells.items(): + delta = cell.get("delta_b_minus_a") + if delta is None or delta["separated"] is None: + continue + if not delta["separated"]: + out["ties"].append(key) + elif delta["point"] > 0: + out["b_beats_a"].append(key) + else: + out["a_beats_b"].append(key) + return out + + +def _fmt(v: float | None, digits: int = 3) -> str: + return "n/a" if v is None else f"{v:.{digits}f}" + + +def markdown_table( + cells: dict[str, Any], gbm: dict[str, Any], *, label_a: str, label_b: str +) -> str: + """One row per (event, horizon): both AUROCs, the paired delta, the GBM refits.""" + lines = [ + f"| event | h | n | n+ | subjects | {label_a} | {label_b} | " + f"{label_b} minus {label_a} [95% CI] | sep | gbm {label_a} | gbm {label_b} |", + "|---|---|---|---|---|---|---|---|---|---|---|", + ] + for key, cell in cells.items(): + event, _, horizon = key.rpartition("@") + delta = cell.get("delta_b_minus_a") + g = gbm.get(key) or {} + delta_txt = ( + "n/a" + if delta is None + else f"{delta['point']:+.3f} [{_fmt(delta['ci_low'])}, " + f"{_fmt(delta['ci_high'])}]" + ) + sep = "" if delta is None else ("yes" if delta["separated"] else "no") + lines.append( + f"| {event} | {horizon} | {cell['n_at_risk']} | {cell['n_positive']} | " + f"{cell['n_subjects']} | {_fmt(cell['auroc_a'])} | {_fmt(cell['auroc_b'])} " + f"| {delta_txt} | {sep} | {_fmt(g.get('auroc_a'))} | " + f"{_fmt(g.get('auroc_b'))} |" + ) + return "\n".join(lines) + + +def compare( + a: pl.DataFrame, + b: pl.DataFrame, + *, + label_a: str, + label_b: str, + scorer: str, + events: list[str] | None, + horizons: list[float] | None, + n_boot: int, + seed: int, +) -> dict[str, Any]: + """Join the two dumps per event and score every (event, horizon) cell.""" + events_a = set(a["event"].unique().to_list()) + events_b = set(b["event"].unique().to_list()) + wanted = sorted(events) if events else sorted(events_a & events_b) + absent = [e for e in wanted if e not in events_a or e not in events_b] + if absent: + raise RowSetMismatchError( + f"events {absent} are not in both dumps ({label_a}: {sorted(events_a)}, " + f"{label_b}: {sorted(events_b)})" + ) + if not events and events_a != events_b: + logger.warning( + "event sets differ, scoring the intersection only: %s only in %s, " + "%s only in %s", + sorted(events_a - events_b), + label_a, + sorted(events_b - events_a), + label_b, + ) + out: dict[str, Any] = {"events": {}, "cells": {}, "gbm": {}} + for event in wanted: + joined, counts = join_event(a, b, event, label_a=label_a, label_b=label_b) + out["events"][event] = counts + hs = horizons if horizons else horizons_in(joined, scorer) + for horizon in hs: + cell = score_pair(joined, scorer, horizon, n_boot=n_boot, seed=seed) + if cell is None: + continue + key = f"{event}@{horizon:g}h" + out["cells"][key] = cell + g = gbm_pair(joined, horizon) + if g is not None: + out["gbm"][key] = g + delta = cell.get("delta_b_minus_a") + logger.info( + "%-28s n=%d (+%d, %d subjects) %s=%s %s=%s delta=%s", + key, + cell["n_at_risk"], + cell["n_positive"], + cell["n_subjects"], + label_a, + _fmt(cell["auroc_a"]), + label_b, + _fmt(cell["auroc_b"]), + "n/a" + if delta is None + else f"{delta['point']:+.4f} [{_fmt(delta['ci_low'], 4)}, " + f"{_fmt(delta['ci_high'], 4)}]", + ) + out["summary"] = summarise(out["cells"]) + return out + + +def main() -> None: + """Paired AUROC deltas between two runs' alert row dumps.""" + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--dump-a", required=True) + parser.add_argument("--dump-b", required=True) + parser.add_argument("--label-a", required=True) + parser.add_argument("--label-b", required=True) + parser.add_argument("--inference-a", default=None) + parser.add_argument("--inference-b", default=None) + parser.add_argument("--scorer", default="hazard") + parser.add_argument("--events", nargs="+", default=None) + parser.add_argument( + "--horizons", + nargs="+", + type=float, + default=None, + help="default: every horizon with score and label columns", + ) + parser.add_argument("--n-boot", type=int, default=1000) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument( + "--max-subjects", + type=int, + default=None, + help="seeded subject-level subsample (same subjects in both dumps)", + ) + parser.add_argument("--output-json", required=True) + args = parser.parse_args() + scorer = SCORER_ALIASES.get(args.scorer, args.scorer) + + a = load_dump(args.dump_a) + b = load_dump(args.dump_b) + version_a, version_b = _protocol_version(a), _protocol_version(b) + if version_a != version_b: + logger.warning( + "landmark protocol versions differ (%s: %s, %s: %s); the join " + "check below decides whether the row sets still agree", + args.label_a, + version_a, + args.label_b, + version_b, + ) + n_before = {"a": a.height, "b": b.height} + a = subsample_subjects(a, max_subjects=args.max_subjects, seed=args.seed) + if args.max_subjects is not None: + # the same subjects in both arms, whatever their order in dump B + b = b.filter(pl.col("subject_id").is_in(a["subject_id"].unique().to_list())) + + try: + result = compare( + a, + b, + label_a=args.label_a, + label_b=args.label_b, + scorer=scorer, + events=args.events, + horizons=args.horizons, + n_boot=args.n_boot, + seed=args.seed, + ) + except RowSetMismatchError as exc: + logger.error("REFUSED: %s", exc) + sys.exit(2) + + out: dict[str, Any] = { + "label_a": args.label_a, + "label_b": args.label_b, + "dump_a": str(args.dump_a), + "dump_b": str(args.dump_b), + "scorer": scorer, + "row_key": ["event", *ROW_KEY], + "n_boot": args.n_boot, + "seed": args.seed, + "max_subjects": args.max_subjects, + "landmark_protocol_version": {"a": version_a, "b": version_b}, + "n_rows_in_dumps": n_before, + "n_rows_scored": {"a": a.height, "b": b.height}, + "variance_scope": "finite-sample only (one fitted model per arm); " + "the gbm block shows refit variance without a bootstrap", + **result, + "inference": inference_block( + read_inference(args.inference_a), read_inference(args.inference_b) + ), + } + Path(args.output_json).parent.mkdir(parents=True, exist_ok=True) + with open(args.output_json, "w") as f: + json.dump(out, f, indent=1) + print( + markdown_table( + out["cells"], out["gbm"], label_a=args.label_a, label_b=args.label_b + ) + ) + if out["inference"] is not None: + print() + print(f"| metric | {args.label_a} | {args.label_b} | b minus a |") + print("|---|---|---|---|") + for k in INFERENCE_KEYS: + va = (out["inference"]["a"] or {}).get(k) + vb = (out["inference"]["b"] or {}).get(k) + d = out["inference"]["delta_b_minus_a"][k] + print( + f"| {k} | {va if va is None or k.startswith('n_') else _fmt(va, 4)} | " + f"{vb if vb is None or k.startswith('n_') else _fmt(vb, 4)} | " + f"{'' if d is None else f'{d:+.4f}'} |" + ) + s = out["summary"] + logger.info( + "wrote %s (%d cells): %s beats %s in %d, %s beats %s in %d, ties %d", + args.output_json, + len(out["cells"]), + args.label_b, + args.label_a, + len(s["b_beats_a"]), + args.label_a, + args.label_b, + len(s["a_beats_b"]), + len(s["ties"]), + ) + + +if __name__ == "__main__": + main() diff --git a/scripts/gemini/run.sh b/scripts/gemini/run.sh index b3e69709..c664968b 100755 --- a/scripts/gemini/run.sh +++ b/scripts/gemini/run.sh @@ -4,7 +4,7 @@ # nobody can log into the node directly, so this is what actually runs there. # # Usage (on the GEMINI node, from the repo root): -# scripts/gemini/run.sh [probe|schema|env-gpu|extract-dry|extract|finalize|export-codes|pipeline|train-smoke|train-smoke-2|train-full|train-smoke-cbm|train-full-cbm|train-smoke-dec|train-full-dec|train-rung2|eval-forecast |interventions |alerts |tabicl |steering |outcome-probes |steering-twosided |atlas |train|eval|all] +# scripts/gemini/run.sh [probe|schema|env-gpu|extract-dry|extract|finalize|export-codes|pipeline|train-smoke|train-smoke-2|train-full|train-smoke-cbm|train-full-cbm|train-smoke-dec|train-full-dec|train-rung2|eval-forecast |interventions |alerts |tabicl |steering |outcome-probes |steering-twosided |atlas |alerts-cis |panel-coverage |cohort-counts |train|eval|all] # # Steps: # probe scripts/gemini/probe_env.sh -> scripts/gemini/out/env_probe.txt @@ -208,6 +208,50 @@ # sidecar) and ~/runs//intervention_cis.json, and # exports only the CI file (aggregate-only) to # scripts/gemini/out/evals/. +# alerts-cis +# the paired subject-clustered bootstrap behind the alerts +# leg (docs/ml4h2026_rebuttal_plan.md WP5(a)): runs +# scripts/alerts_cis.py on the row dump the `alerts` step +# left on this node, hazard heads against the tuned GBM, per +# (event, horizon) cell, 1000 draws. Reads +# ~/runs//alerts_rows_allshards.parquet by default +# (GEMINI_ALERTS_TAG, default _allshards: the current GBM +# refit on all 894 train shards at 10% row thinning, the +# number tab:gemini quotes) and writes ~/runs// +# alerts_cis_allshards.json. Prints the dump's row and +# subject counts before starting; GEMINI_CIS_MAX_SUBJECTS +# (default unset = every subject) maps to --max-subjects +# when the full bootstrap is not affordable (the ~25M-row +# dump takes 15-20 h), GEMINI_CIS_N_BOOT overrides the draw +# count. Re-running with the output already present only +# re-exports it (GEMINI_CIS_FORCE=1 recomputes). Exports +# (aggregate-only) to scripts/gemini/out/evals/ +# _allshards_alerts_cis.json. +# panel-coverage +# how many of the GBM feature panel's 48 signals resolve on +# GEMINI (WP5(b)): scripts/panel_coverage.py against the +# per-source LOINC table in odyssey/data/code_mapping.py, +# and against metadata/codes.parquet when present (which +# resolved prefixes were actually charted). Signal and +# prefix names only, no counts. Writes ~/runs// +# panel_coverage.json and exports it to +# scripts/gemini/out/evals/_panel_coverage.json. +# cohort-counts +# the aggregate cohort description for the GEMINI arm +# (WP5(c)): scripts/cohort_counts.py over the MEDS shards +# the run trained on (train/tuning from the run's +# config.json) plus the held-out split -- subjects, +# admissions, hospitals (metadata/hadm_id_hospital.parquet), +# admission year range, length of stay, sex and age where +# the source charts them (GEMINI extracts neither), and the +# share of subjects with each hazard event under the alerts +# leg's own onset definitions. Every count below 10 is +# written as "<10". GEMINI_COHORT_MAX_SHARDS caps every +# split (recorded in the output), GEMINI_COHORT_NO_EVENTS=1 +# skips the label pass. Re-running with the output present +# only re-exports it (GEMINI_COHORT_FORCE=1 recomputes). +# Writes ~/runs//cohort_counts.json and exports it +# to scripts/gemini/out/evals/_cohort_counts.json. # train not built yet (a general, non-GEMINI-specific full run). # eval not built yet. # all probe, schema, extract-dry, in order (default; deliberately @@ -254,6 +298,16 @@ main() { # given -- the whitelist must match that exactly, not just the keys a # particular invocation cares about, or _export_aggregate_json refuses. STEERING_JSON_KEYS="site control control_seed stratify_by layer_index tau suppress_strength horizons_hours event_names trained_event_names outcome_probes gammas lifted_tokens at_risk_restricted summaries" + # scripts/alerts_cis.py's top-level output keys (its main() builds the + # dict; the per-cell block has n/n_positive/n_subjects/scorers/ + # paired_deltas, the summary block one flat record per cell). + ALERTS_CIS_JSON_KEYS="scorers n_boot seed max_subjects variance_scope n_rows_in_dumps n_subjects_in_dumps n_rows_scored n_subjects_scored cells summary" + # scripts/panel_coverage.py's OUTPUT_KEYS and scripts/cohort_counts.py's + # OUTPUT_KEYS, verbatim -- tests/scripts/gemini/test_run_sh_rebuttal_steps.py + # pins these two lines to the scripts so a new field refuses here in CI, + # not on the node (the #238 incident). + PANEL_COVERAGE_JSON_KEYS="source n_panel_signals n_resolved unresolved resolved prefixes codes_inventory resolved_observed resolved_unobserved" + COHORT_COUNTS_JSON_KEYS="source task_set normalize_medications suppression_threshold max_shards los_definition hospital_metadata events events_dropped splits all_splits" # --- sync with the mirror (fetch + reset, never pull) --------------------- # @@ -2242,6 +2296,209 @@ PY || echo "WARNING: atlas output not exported (see above)." >&2 } + run_alerts_cis() { + # WP5(a) of docs/ml4h2026_rebuttal_plan.md: the interval the paper's + # GEMINI hazard-vs-GBM comparison (tab:gemini) never had. Runs + # scripts/alerts_cis.py on the row dump the `alerts` step left on + # this node -- the dump fixes the rows, labels and every scorer's + # scores, so nothing is re-scored and the CI describes exactly the + # banked point estimates. Defaults to the _allshards dump: the + # current GBM refit (all 894 train shards at 10% row thinning, + # docs/experiments.md "alerts, GBM refit on every train shard"). + local run_name="$1" + if [[ -z "$run_name" ]]; then + echo "alerts-cis needs a run name: scripts/gemini/run.sh alerts-cis " >&2 + echo "e.g.: scripts/gemini/run.sh alerts-cis gemini_full_DEC_v12" >&2 + exit 1 + fi + local alerts_tag="${GEMINI_ALERTS_TAG:-_allshards}" + local n_boot="${GEMINI_CIS_N_BOOT:-1000}" + local max_subjects="${GEMINI_CIS_MAX_SUBJECTS:-}" + local scorers="${GEMINI_CIS_SCORERS:-hazard gbm}" + local seed="${GEMINI_CIS_SEED:-0}" + local force="${GEMINI_CIS_FORCE:-0}" + + echo "=== alerts-cis ($run_name) ===" + echo "Paired subject-clustered bootstrap on the alerts row dump:" + echo "alerts_tag=$alerts_tag (GEMINI_ALERTS_TAG; _allshards = the current GBM refit)," + echo "scorers=$scorers, n_boot=$n_boot, seed=$seed," + echo "max_subjects=${max_subjects:-none (every subject)} (GEMINI_CIS_MAX_SUBJECTS)." + if [[ -z "${TMUX:-}" && -z "${STY:-}" ]]; then + echo "WARNING: this doesn't look like a tmux or screen session --" >&2 + echo "a 1000-draw bootstrap over the full dump runs 15-20 h and dies" >&2 + echo "with a dropped SSH connection. Run it detached instead, e.g.:" >&2 + echo " tmux new -s alerts-cis-$run_name 'scripts/gemini/run.sh $STEP $run_name'" >&2 + fi + _require_run_and_data "$run_name" + + DUMP_ROWS="$RUN_DIR/alerts_rows${alerts_tag}.parquet" + OUTPUT_JSON="$RUN_DIR/alerts_cis${alerts_tag}.json" + if [[ ! -f "$DUMP_ROWS" ]]; then + echo "No row dump at $DUMP_ROWS -- the alerts step writes it. Run" >&2 + echo " GEMINI_ALERT_SHARDS=1000 GEMINI_ALERTS_TAG=$alerts_tag scripts/gemini/run.sh alerts $run_name" >&2 + echo "first, or point GEMINI_ALERTS_TAG at the tag of a dump that exists:" >&2 + ls -1 "$RUN_DIR"/alerts_rows*.parquet 2>/dev/null >&2 || echo " (no alerts_rows*.parquet under $RUN_DIR)" >&2 + exit 1 + fi + echo "Run dir: $RUN_DIR" + echo "Row dump: $DUMP_ROWS (patient-level, stays on this node)" + echo "Output: $OUTPUT_JSON" + + source "$GPU_VENV/bin/activate" + if [[ -f "$OUTPUT_JSON" && "$force" != "1" ]]; then + echo "Output already exists; not recomputing (set GEMINI_CIS_FORCE=1 to)." + else + # Row and subject counts before committing to the bootstrap: the + # operator decides from these whether GEMINI_CIS_MAX_SUBJECTS is + # needed. A lazy scan, so the dump is not materialized twice. + python - "$DUMP_ROWS" <<'PY' +import sys + +import polars as pl + +path = sys.argv[1] +counts = ( + pl.scan_parquet(path) + .select( + pl.len().alias("rows"), + pl.col("subject_id").n_unique().alias("subjects"), + pl.col("event").n_unique().alias("events"), + ) + .collect() +) +rows, subjects, events = counts.row(0) +print(f"dump: {rows:,} rows, {subjects:,} subjects, {events} events") +PY + local extra=() + if [[ -n "$max_subjects" ]]; then + extra+=(--max-subjects "$max_subjects") + fi + # shellcheck disable=SC2086 + python scripts/alerts_cis.py \ + --dump "$DUMP_ROWS" \ + --output-json "$OUTPUT_JSON" \ + --scorers $scorers \ + --n-boot "$n_boot" \ + --seed "$seed" \ + "${extra[@]}" + fi + deactivate + source "$VENV/bin/activate" + + echo "$STEP complete. Results at $OUTPUT_JSON" + _export_aggregate_json \ + "scripts/gemini/out/evals/${run_name}${alerts_tag}_alerts_cis.json" "$OUTPUT_JSON" \ + "$ALERTS_CIS_JSON_KEYS" \ + || echo "WARNING: alerts CI output not exported (see above)." >&2 + } + + run_panel_coverage() { + # WP5(b): which of the GBM feature panel's 48 signals resolve on + # GEMINI, by name. The LOINC table is in the repo, so most of this + # could run anywhere; what only this node can add is + # metadata/codes.parquet, which says whether each resolved prefix + # was actually charted. Names only, never counts. + local run_name="$1" + if [[ -z "$run_name" ]]; then + echo "panel-coverage needs a run name: scripts/gemini/run.sh panel-coverage " >&2 + echo "e.g.: scripts/gemini/run.sh panel-coverage gemini_full_DEC_v12" >&2 + exit 1 + fi + echo "=== panel-coverage ($run_name) ===" + _require_run_and_data "$run_name" + + OUTPUT_JSON="$RUN_DIR/panel_coverage.json" + local codes="$METADATA_DIR/codes.parquet" + local extra=() + if [[ -f "$codes" ]]; then + extra+=(--codes-parquet "$codes") + echo "Code inventory: $codes" + else + echo "No code inventory at $codes -- resolving against the LOINC table only." + fi + echo "Output: $OUTPUT_JSON" + + source "$GPU_VENV/bin/activate" + python scripts/panel_coverage.py \ + --source gemini \ + --output-json "$OUTPUT_JSON" \ + "${extra[@]}" + deactivate + source "$VENV/bin/activate" + + echo "$STEP complete. Results at $OUTPUT_JSON" + _export_aggregate_json \ + "scripts/gemini/out/evals/${run_name}_panel_coverage.json" "$OUTPUT_JSON" \ + "$PANEL_COVERAGE_JSON_KEYS" \ + || echo "WARNING: panel coverage not exported (see above)." >&2 + } + + run_cohort_counts() { + # WP5(c): the aggregate cohort description the paper's GEMINI arm + # never had. Train/tuning come from the run's own config.json (the + # shards it trained on), held-out from the eval dir every other + # step reads; hospitals from metadata/hadm_id_hospital.parquet. + # Every count below 10 leaves as "<10"; the script never writes a + # subject id or a per-patient value. + local run_name="$1" + if [[ -z "$run_name" ]]; then + echo "cohort-counts needs a run name: scripts/gemini/run.sh cohort-counts " >&2 + echo "e.g.: scripts/gemini/run.sh cohort-counts gemini_full_DEC_v12" >&2 + exit 1 + fi + local max_shards="${GEMINI_COHORT_MAX_SHARDS:-}" + local no_events="${GEMINI_COHORT_NO_EVENTS:-0}" + local force="${GEMINI_COHORT_FORCE:-0}" + + echo "=== cohort-counts ($run_name) ===" + echo "max_shards=${max_shards:-none (every shard)} (GEMINI_COHORT_MAX_SHARDS)," + echo "no_events=$no_events (GEMINI_COHORT_NO_EVENTS=1 skips the label pass)." + if [[ -z "${TMUX:-}" && -z "${STY:-}" ]]; then + echo "WARNING: this doesn't look like a tmux or screen session --" >&2 + echo "the label pass over ~1000 shards takes a while; run it detached, e.g.:" >&2 + echo " tmux new -s cohort-$run_name 'scripts/gemini/run.sh $STEP $run_name'" >&2 + fi + _require_run_and_data "$run_name" + + OUTPUT_JSON="$RUN_DIR/cohort_counts.json" + local hospitals="$METADATA_DIR/hadm_id_hospital.parquet" + local extra=() + if [[ -f "$hospitals" ]]; then + extra+=(--hadm-hospital-parquet "$hospitals") + echo "Hospital table: $hospitals" + else + echo "No hospital table at $hospitals -- hospitals will read 'not available'." + fi + if [[ -n "$max_shards" ]]; then + extra+=(--max-shards "$max_shards") + fi + if [[ "$no_events" == "1" ]]; then + extra+=(--no-events) + fi + echo "Run dir: $RUN_DIR (train/tuning split dirs from its config.json)" + echo "Held-out: $HELD_OUT_SHARD_DIR" + echo "Output: $OUTPUT_JSON" + + source "$GPU_VENV/bin/activate" + if [[ -f "$OUTPUT_JSON" && "$force" != "1" ]]; then + echo "Output already exists; not recomputing (set GEMINI_COHORT_FORCE=1 to)." + else + python scripts/cohort_counts.py \ + --run-dir "$RUN_DIR" \ + --split "held_out=$HELD_OUT_SHARD_DIR" \ + --output-json "$OUTPUT_JSON" \ + "${extra[@]}" + fi + deactivate + source "$VENV/bin/activate" + + echo "$STEP complete. Results at $OUTPUT_JSON" + _export_aggregate_json \ + "scripts/gemini/out/evals/${run_name}_cohort_counts.json" "$OUTPUT_JSON" \ + "$COHORT_COUNTS_JSON_KEYS" \ + || echo "WARNING: cohort counts not exported (see above)." >&2 + } + case "$STEP" in probe) run_probe ;; schema) run_schema ;; @@ -2268,11 +2525,14 @@ PY outcome-probes) run_outcome_probes "${2:-}" ;; steering-twosided) run_steering_twosided "${2:-}" ;; atlas) run_atlas "${2:-}" ;; + alerts-cis) run_alerts_cis "${2:-}" ;; + panel-coverage) run_panel_coverage "${2:-}" ;; + cohort-counts) run_cohort_counts "${2:-}" ;; train) run_pending_stub train ;; eval) run_pending_stub eval ;; all) run_probe; run_schema; run_extract_dry ;; *) - echo "unknown step: $STEP (expected probe, schema, env-gpu, extract-dry, extract, finalize, export-codes, pipeline, train-smoke, train-smoke-2, train-full, train-smoke-cbm, train-full-cbm, train-smoke-dec, train-full-dec, train-rung2, eval-forecast, interventions, alerts, tabicl, steering, outcome-probes, steering-twosided, atlas, train, eval, or all)" >&2 + echo "unknown step: $STEP (expected probe, schema, env-gpu, extract-dry, extract, finalize, export-codes, pipeline, train-smoke, train-smoke-2, train-full, train-smoke-cbm, train-full-cbm, train-smoke-dec, train-full-dec, train-rung2, eval-forecast, interventions, alerts, tabicl, steering, outcome-probes, steering-twosided, atlas, alerts-cis, panel-coverage, cohort-counts, train, eval, or all)" >&2 exit 1 ;; esac diff --git a/scripts/make_hparams_table.py b/scripts/make_hparams_table.py new file mode 100644 index 00000000..0070e84b --- /dev/null +++ b/scripts/make_hparams_table.py @@ -0,0 +1,300 @@ +"""Model size and training settings of the flagship runs, from banked configs. + +The ML4H submission gives no hyperparameters: no hidden size, no +optimizer, no learning rate, no loss weights, no parameter count, and +no description of the GBM comparator's search. This script makes that +appendix table a build product of the banked ``config.json`` each +flagship training run wrote, so the numbers cannot drift from what was +actually trained. + +One column per database. A run whose ``config.json`` was never exported +(GEMINI's ``gemini_full_v10_15c`` only banked its evaluation files) +gets a column of ``--`` under a "not exported" marker rather than being +silently dropped or, worse, filled in from another database's config. + +The GBM block below the model rows comes from the code, not from a +config: the estimator class, the four configurations searched, the +round budget and the validation scheme are read from +``odyssey.inference.alerts`` (``GBM_GRID``, ``GBM_MAX_ITER``, +``_tune_gbm``), and the panel size and its count-feature share are +computed from ``odyssey.inference.baseline_features.feature_names`` and +the ablation's ``feature_groups`` partition. + +Parameter counts are not in the configs; pass ``--params`` with a JSON +file mapping run name (the basename of the config's ``output_dir``, or +the column label) to an integer. Missing runs print ``tbd``. + +Usage:: + + uv run python scripts/make_hparams_table.py \\ + --run MIMIC-IV research_journal/figure_data/vm1/full_run_v10/config.json \\ + --run eICU-CRD research_journal/figure_data/vm2/eicu_full_v10/config.json \\ + --missing GEMINI \\ + --params research_journal/figure_data/param_counts.json \\ + --output paper/ml4h/tables/hparams.tex +""" + +from __future__ import annotations + +import argparse +import json +import logging +from collections.abc import Callable +from pathlib import Path +from typing import Any + +from odyssey.inference.alerts import GBM_GRID, GBM_MAX_ITER, GBM_TUNE_MAX_ROWS +from odyssey.inference.baseline_features import feature_names +from scripts.gbm_feature_ablation import feature_groups + + +logger = logging.getLogger(__name__) + +#: Marker for a column whose config was never exported. +NOT_EXPORTED = "not exported" +#: Marker for a parameter count that has not been computed yet. +TBD = "tbd" + +_BACKBONES = { + "hybrid": "hybrid Mamba-2 + chunk attention", + "mamba": "Mamba-2", + "transformer": "transformer", +} + + +def _num(value: Any) -> str: + """Numbers the way the other tables print them: no trailing zeros.""" + if isinstance(value, bool): + return "yes" if value else "no" + if isinstance(value, int): + return f"{value:,}".replace(",", "{,}") + if isinstance(value, float): + if value != 0 and abs(value) < 1e-2: + mantissa, exponent = f"{value:.0e}".split("e") + return f"${mantissa}\\times10^{{{int(exponent)}}}$" + return f"{value:g}" + return str(value) + + +def _key(name: str) -> Callable[[dict[str, Any]], str]: + def get(config: dict[str, Any]) -> str: + return _num(config[name]) + + return get + + +def _backbone(config: dict[str, Any]) -> str: + backbone = str(config["backbone"]) + return _BACKBONES.get(backbone, backbone) + + +def _lanes(config: dict[str, Any]) -> str: + return f"{_num(config['num_lanes'])} $\\times$ {_num(config['chunk_size'])}" + + +def _lr(config: dict[str, Any]) -> str: + return _num(float(config["learning_rate"])) + + +#: Model rows: (label, reader). Every reader takes the config dict. +MODEL_ROWS: tuple[tuple[str, Callable[[dict[str, Any]], str]], ...] = ( + ("Backbone", _backbone), + ("Hidden size", _key("hidden_size")), + ("Layers", _key("num_hidden_layers")), + ("Attention heads", _key("attn_num_heads")), + ("Mamba state size", _key("mamba_state_size")), + ("Mamba head dim", _key("mamba_headdim")), + ("Mamba chunk size", _key("mamba_chunk_size")), + ("Concept embedding dim", _key("embedding_dim")), + ("Lanes $\\times$ chunk (tokens)", _lanes), + ("Max context (tokens)", _key("max_context")), +) + +#: Training rows. AdamW with torch defaults and a constant learning rate: +#: odyssey/training/train.py builds ``torch.optim.AdamW`` from +#: ``optimizer_param_groups`` and no scheduler. +TRAINING_ROWS: tuple[tuple[str, Callable[[dict[str, Any]], str]], ...] = ( + ("Optimizer", lambda _: "AdamW (constant LR)"), + ("Learning rate", _lr), + ("Weight decay", _key("weight_decay")), + ("Gradient clip (norm)", _key("grad_clip_norm")), + ("Epochs", _key("num_epochs")), + ("Early-stopping patience (evals)", _key("early_stopping_patience")), + ("Checkpoint every (steps)", _key("checkpoint_every")), + ("Seed", _key("seed")), + ("RandInt probability", _key("randint_prob")), +) + +LOSS_ROWS: tuple[tuple[str, Callable[[dict[str, Any]], str]], ...] = ( + ("Concept", _key("concept_weight")), + ("Orthogonality", _key("orthogonality_weight")), + ("Observability", _key("observability_weight")), + ("Task (next token)", _key("task_weight")), + ("Time to event", _key("time_weight")), + ("Event hazard", _key("event_hazard_weight")), +) + +VOCAB_ROWS: tuple[tuple[str, Callable[[dict[str, Any]], str]], ...] = ( + ("Vocabulary min count", _key("vocab_min_count")), + ("Vocabulary max size", _key("vocab_max_size")), + ("Vocabulary backoff", _key("vocab_backoff")), + ("Quantile bins per lab", _key("quantile_n_bins")), + ("Quantile min count", _key("quantile_min_count")), +) + + +def run_name(config: dict[str, Any] | None, label: str) -> str: + """Return the run's directory name, which is how ``--params`` keys it.""" + if config is None or not config.get("output_dir"): + return label + return Path(str(config["output_dir"])).name + + +def _param_cell( + config: dict[str, Any] | None, label: str, params: dict[str, int] +) -> str: + for key in (run_name(config, label), label): + if key in params: + return _num(int(params[key])) + return TBD + + +def gbm_rows() -> list[tuple[str, str]]: + """Read the GBM comparator's settings from the code that fits it.""" + names = feature_names() + n_counts = len(feature_groups(names)["counts_occurrence"]) + grid = "; ".join( + f"({_num(p['learning_rate'])}, {_num(int(p['max_leaf_nodes']))}, " + f"{_num(int(p['min_samples_leaf']))})" + for p in GBM_GRID + ) + return [ + ("Estimator", "scikit-learn \\texttt{HistGradientBoostingClassifier}"), + ( + "Search grid (LR, max leaves, min leaf)", + grid, + ), + ( + "Boosting rounds", + f"up to {_num(GBM_MAX_ITER)}; best round by validation log loss, " + "then refit at that count", + ), + ( + "Validation", + "10\\% of training subjects held out (subject-grouped), " + f"tuned on at most {_num(GBM_TUNE_MAX_ROWS)} rows", + ), + ("Feature panel", f"{_num(len(names))} features, {_num(n_counts)} counts"), + ("Missing values", "native (columns observed in $<$200 rows filled with 0)"), + ] + + +def render( + runs: list[tuple[str, dict[str, Any] | None]], + params: dict[str, int] | None = None, +) -> str: + """One row per setting, one column per database, GBM block at the end.""" + params = params or {} + n = len(runs) + labels = [label for label, _ in runs] + + def line(label: str, cells: list[str]) -> str: + return f"{label} & " + " & ".join(cells) + " \\\\" + + def row( + label: str, reader: Callable[[dict[str, Any]], str], missing: str = "--" + ) -> str: + return line(label, [missing if cfg is None else reader(cfg) for _, cfg in runs]) + + def block( + title: str, rows: tuple[tuple[str, Callable[..., str]], ...] + ) -> list[str]: + return [ + f"\\multicolumn{{{n + 1}}}{{@{{}}l}}{{\\emph{{{title}}}}} \\\\", + *(row(label, reader) for label, reader in rows), + ] + + lines = [ + "% GENERATED by scripts/make_hparams_table.py -- do not hand-edit.", + "% Model and training settings from each flagship run's banked", + "% config.json. '--': that run's config.json was not exported.", + "% Parameter counts come from --params; 'tbd' means not computed yet.", + "% GBM block: odyssey/inference/alerts.py (GBM_GRID, GBM_MAX_ITER,", + "% GBM_TUNE_MAX_ROWS, _tune_gbm, BaselineModel) and the panel from", + "% odyssey/inference/baseline_features.py (feature_names) partitioned", + "% by scripts/gbm_feature_ablation.py (feature_groups).", + "\\begin{tabular}{@{}l" + "r" * n + "@{}}", + "\\toprule", + "Setting & " + " & ".join(labels) + " \\\\", + "\\midrule", + row("Banked config", lambda _: "yes", missing=NOT_EXPORTED), + *block("Model", MODEL_ROWS), + line("Parameters", [_param_cell(cfg, label, params) for label, cfg in runs]), + *block("Training", TRAINING_ROWS), + *block("Loss weights", LOSS_ROWS), + *block("Vocabulary and value bins", VOCAB_ROWS), + "\\midrule", + f"\\multicolumn{{{n + 1}}}{{@{{}}l}}{{\\emph{{GBM comparator " + "(same on every database)}} \\\\", + ] + for label, value in gbm_rows(): + lines.append(f"{label} & \\multicolumn{{{n}}}{{l}}{{{value}}} \\\\") + lines += ["\\bottomrule", "\\end{tabular}"] + return "\n".join(lines) + "\n" + + +def load_config(path: Path) -> dict[str, Any]: + """Load the training config a run banked, as a plain dict.""" + config = json.loads(path.read_text()) + if not isinstance(config, dict): + raise SystemExit(f"{path} is not a JSON object") + return config + + +def main() -> None: + """Write the hyperparameter table.""" + parser = argparse.ArgumentParser( + description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter + ) + parser.add_argument( + "--run", + nargs=2, + action="append", + default=[], + metavar=("LABEL", "CONFIG_JSON"), + help="column label and the run's banked config.json (repeatable)", + ) + parser.add_argument( + "--missing", + action="append", + default=[], + metavar="LABEL", + help="column for a run whose config.json was not exported (repeatable)", + ) + parser.add_argument( + "--params", + type=Path, + help="JSON mapping run name (output_dir basename or label) to parameter count", + ) + parser.add_argument("--output", type=Path, required=True) + args = parser.parse_args() + + if not args.run and not args.missing: + raise SystemExit("give at least one --run or --missing column") + runs: list[tuple[str, dict[str, Any] | None]] = [ + (label, load_config(Path(path))) for label, path in args.run + ] + runs += [(label, None) for label in args.missing] + params: dict[str, int] = {} + if args.params is not None: + params = {k: int(v) for k, v in json.loads(args.params.read_text()).items()} + + args.output.parent.mkdir(parents=True, exist_ok=True) + args.output.write_text(render(runs, params)) + logger.info("wrote %s (%d columns)", args.output, len(runs)) + print(f"wrote {args.output}") + + +if __name__ == "__main__": + logging.basicConfig(level=logging.INFO) + main() diff --git a/scripts/panel_coverage.py b/scripts/panel_coverage.py new file mode 100644 index 00000000..b1d9d1cf --- /dev/null +++ b/scripts/panel_coverage.py @@ -0,0 +1,162 @@ +"""How many of the GBM feature panel's signals resolve on one source. + +The tuned tabular baseline reads :data:`odyssey.data.signal_panel.SIGNAL_PANEL` +(48 vitals and labs, keyed by LOINC). A signal resolves on a source when +that source's LOINC table in :mod:`odyssey.data.code_mapping` maps at +least one code prefix to the signal's LOINC; a signal with no prefix +never produces a feature there, so the baseline on that source runs with +a smaller panel. This script reports the resolved and unresolved names +for one source, which is the number the ML4H rebuttal needs for GEMINI +(review W3: "how many of the 48 signals does the GEMINI GBM actually +see?"). + +With ``--codes-parquet`` (a MEDS ``metadata/codes.parquet`` with a +``code`` column) the resolved signals are further split by whether any +code in that inventory starts with one of the signal's prefixes, so a +prefix that is in the table but never charted is reported too. Only +signal and prefix NAMES are written, never counts, so the output is safe +to export from a closed environment. + +Usage:: + + uv run python scripts/panel_coverage.py --source gemini \ + [--codes-parquet /metadata/codes.parquet] \ + --output-json ~/runs//panel_coverage.json +""" + +from __future__ import annotations + +import argparse +import json +from collections.abc import Iterable +from pathlib import Path +from typing import Any + +import polars as pl + +from odyssey.data.code_mapping import prefixes_for_loinc +from odyssey.data.signal_panel import N_PANEL_SIGNALS, SIGNAL_PANEL + + +SOURCES: tuple[str, ...] = ("mimic_iv", "eicu", "gemini") + +#: Top-level keys of the JSON this script writes. run.sh's export +#: whitelist for the panel-coverage step must list exactly these. +OUTPUT_KEYS: tuple[str, ...] = ( + "source", + "n_panel_signals", + "n_resolved", + "unresolved", + "resolved", + "prefixes", + "codes_inventory", + "resolved_observed", + "resolved_unobserved", +) + + +def panel_coverage( + source: str, *, observed_codes: Iterable[str] | None = None +) -> dict[str, Any]: + """Resolve every panel signal against ``source``'s LOINC table. + + Returns the names that resolve and the names that do not, plus the + prefixes each resolved signal matches. When ``observed_codes`` is + given (the distinct MEDS codes of that source), the resolved signals + are further split into those with at least one observed code and + those whose prefixes never occur. Codes carrying a ``::`` suffix + are matched on their un-binned form, the rule + :class:`~odyssey.data.signal_panel.SignalPanelResolver` uses. + """ + resolved: list[str] = [] + unresolved: list[str] = [] + prefixes: dict[str, list[str]] = {} + for name, loinc in SIGNAL_PANEL: + hits = sorted(prefixes_for_loinc(loinc, source=source)) + if hits: + resolved.append(name) + prefixes[name] = hits + else: + unresolved.append(name) + + out: dict[str, Any] = { + "source": source, + "n_panel_signals": N_PANEL_SIGNALS, + "n_resolved": len(resolved), + "unresolved": unresolved, + "resolved": resolved, + "prefixes": prefixes, + "codes_inventory": False, + "resolved_observed": None, + "resolved_unobserved": None, + } + if observed_codes is None: + return out + + bases = {c.rsplit("::", 1)[0] if "::" in c else c for c in observed_codes} + observed: list[str] = [] + unobserved: list[str] = [] + for name in resolved: + if any(base.startswith(p) for p in prefixes[name] for base in bases): + observed.append(name) + else: + unobserved.append(name) + out["codes_inventory"] = True + out["resolved_observed"] = observed + out["resolved_unobserved"] = unobserved + return out + + +def load_codes(path: str | Path) -> list[str]: + """Distinct ``code`` strings from a MEDS ``codes.parquet`` (or any parquet).""" + frame = pl.read_parquet(path, columns=["code"]) + return frame["code"].drop_nulls().unique().to_list() + + +def format_table(report: dict[str, Any]) -> str: + """Render the report as a plain-text table for the log.""" + lines = [ + f"source: {report['source']}", + f"panel signals: {report['n_panel_signals']}", + f"resolved: {report['n_resolved']}", + "", + f"{'signal':<24}{'status':<12}prefixes", + ] + observed = report.get("resolved_observed") + unobserved = set(report.get("resolved_unobserved") or []) + for name, _ in SIGNAL_PANEL: + if name in report["prefixes"]: + status = "resolved" + if observed is not None and name in unobserved: + status = "no codes" + lines.append(f"{name:<24}{status:<12}{', '.join(report['prefixes'][name])}") + else: + lines.append(f"{name:<24}{'unresolved':<12}") + return "\n".join(lines) + + +def main() -> None: + """Report panel coverage for one source and write it as JSON.""" + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--source", choices=SOURCES, required=True) + parser.add_argument( + "--codes-parquet", + default=None, + help="MEDS metadata/codes.parquet; splits resolved signals by whether " + "any charted code matches (names only are written)", + ) + parser.add_argument("--output-json", default=None) + args = parser.parse_args() + + codes = load_codes(args.codes_parquet) if args.codes_parquet else None + report = panel_coverage(args.source, observed_codes=codes) + print(format_table(report)) + if args.output_json: + Path(args.output_json).parent.mkdir(parents=True, exist_ok=True) + with open(args.output_json, "w") as f: + json.dump(report, f, indent=1) + print(f"wrote {args.output_json}") + + +if __name__ == "__main__": + main() diff --git a/scripts/probe_channel.py b/scripts/probe_channel.py new file mode 100644 index 00000000..c4febf7a --- /dev/null +++ b/scripts/probe_channel.py @@ -0,0 +1,1446 @@ +"""Channel probes: how much of the forecast runs through k versus the poles. + +The mixture bottleneck builds each named slot as ``z_i = k_i w+_i + +(1 - k_i) w-_i``, where ``k_i`` is the concept probability shown to a +clinician and the poles ``w+_i``/``w-_i`` are free functions of the +hidden state. Zeroing every ``z_i`` collapses next-event top-1 to a few +percent, but that does not say whether the information lives in the +probabilities or in the poles. This script measures it directly on a +frozen checkpoint: one streaming pass over the held-out shards dumps +``k``, the unnamed slot's probability ``u``, the poles, and the unnamed +embedding ``z_u`` at every scored position (or a seeded, shard- +stratified subsample), then fresh readouts are fit on a SUBJECT-level +split and scored on held-out subjects for exact next-event top-1, the +paper's definition (argmax over the full vocabulary equals the next +token). + +Readouts (each a fresh head, trained from scratch on frozen features): + +- ``k_only``: the concept probabilities alone (Yeh et al. completeness). +- ``k_plus_u``: ``k`` plus the unnamed slot's probability ``u``. +- ``z_named``: the concatenated named embeddings, what the LM head reads + minus the unnamed slot. +- ``poles_mean_k``: ``z`` recomputed with every ``k_i`` held at its + training-split mean, so only the poles vary across positions. +- ``poles_both``: the raw ``[w+, w-]`` concatenation, the ceiling of what + the poles carry regardless of ``k``. +- ``h_bar``: the full bottleneck ``[z, z_u]``, which should recover the + model's own accuracy and serves as the ceiling. +- ``model_head``: the model's own LM head, no refit, the reference. + +The readout is linear by default: a standardized ``nn.Linear`` over the +full next-event vocabulary trained with Adam and early stopping on a +tuning slice of the training subjects, the same probe family +:mod:`odyssey.inference.leakage` uses, chosen because sklearn's +multinomial logistic regression does not scale to a vocabulary of +thousands of classes over millions of rows. ``--mlp`` swaps in a one- +hidden-layer MLP. Every accuracy carries a subject-clustered bootstrap +interval (1000 resamples, seeded) and its ratio to the model's own +accuracy on the same rows ("retained"), with a paired interval on the +ratio. The Yeh-style completeness score is also reported, with the +majority next-event rate as the random baseline (the binary +:func:`odyssey.training.metrics.compute_completeness` does not apply to a +vocabulary-sized target, but its formula does). + +Reading the result: if ``k_only`` recovers most of ``h_bar``, the forecast +runs through the named STATES. If ``poles_mean_k`` does, it runs through +the poles, and the concept probabilities are a small part of the channel. + +The CTL leakage probes of :mod:`odyssey.inference.leakage` (probs-only, +known-embeddings, residual, random-projected probs, on the next-token +code family) are fit on the same subject split from the same dump, via +``compute_ctl``. Hazard-head readouts at landmark positions are NOT +dumped: the landmark protocol (v4, one information boundary per landmark) +lives inside :mod:`odyssey.inference.alerts` and reproducing it here +would risk a number that disagrees with the paper's alert tables. + +Crash recovery: the streaming pass is the expensive part, so the bank is +saved right after it to ``.bank.pt`` (``--bank-path`` +overrides) and ``--from-bank PATH`` reloads it without touching the +checkpoint. The readout block is written to ``--output-json`` as soon as +the readouts finish, before the CTL probes run, and the ``ctl`` block is +added at the end. The JSON is append-only: an existing file with a +``ctl`` block is refused without ``--overwrite``; one without a ``ctl`` +block (a run that died in the CTL probes) has only that block filled in. +With ``--bank-on-cpu`` the CTL banks stay in host memory as well and +:func:`odyssey.inference.leakage.compute_ctl` moves each batch to the GPU. + +Usage:: + + uv run python scripts/probe_channel.py \\ + --run-dir ~/runs/full_run_v10 \\ + --held-out-shard-dir ~/data/mimic/held_out \\ + --output-json ~/runs/full_run_v10/channel_probes.json + + # after a crash past the streaming pass + uv run python scripts/probe_channel.py ... \\ + --from-bank ~/runs/full_run_v10/channel_probes.bank.pt +""" + +from __future__ import annotations + +import argparse +import json +import logging +import math +import time +from collections.abc import Callable, Iterator, Sequence +from dataclasses import asdict, dataclass, field +from pathlib import Path +from typing import Any, cast + +import numpy as np +import polars as pl +import torch +import torch.nn.functional as F # noqa: N812 +from torch import nn + +from odyssey.data.code_normalization import maybe_normalize +from odyssey.data.history_recap import maybe_history_recap +from odyssey.data.sequences import PatientSequence +from odyssey.data.sidecars import activate_sidecars +from odyssey.data.streaming import PackedLaneSampler +from odyssey.data.value_binning import QuantileBinner, add_value_tokens +from odyssey.data.vocabulary import PAD_ID, Vocabulary +from odyssey.inference.leakage import ( + CTLResult, + LeakageBank, + ProbeFitTrace, + _StandardizedLinearProbe, + compute_ctl, +) +from odyssey.inference.legacy_concept_pins import resolve_concepts_for_run +from odyssey.inference.run_inference import _build_type_lookup, load_run +from odyssey.models.concept_bottleneck import ConceptBottleneck +from odyssey.models.sequence_model import ( + ConceptBottleneckSequenceModel, + ConceptLabelDict, + ConceptSupervision, +) +from odyssey.training.data import ( + build_concept_first_times, + build_concept_label_dicts, + build_visit_concept_first_times, + build_visit_concept_label_dicts, + iter_patient_sequences, + load_meds_shard, +) +from odyssey.training.running_labels import position_running_labels +from odyssey.training.shard_stream import shard_paths +from odyssey.training.train import TrainingConfig, _move_chunk_to_device + + +logger = logging.getLogger("probe_channel") + +READOUT_NAMES = ( + "k_only", + "k_plus_u", + "z_named", + "poles_mean_k", + "poles_both", + "h_bar", +) +MODEL_READOUT = "model_head" + + +# --------------------------------------------------------------------------- +# Feature bank +# --------------------------------------------------------------------------- + + +@dataclass +class ChannelBank: + """Frozen bottleneck pieces at scored positions, one row per position. + + Float tensors are float16 at rest; readouts upcast per batch. ``z`` is + not stored: it is ``k w+ + (1 - k) w-`` and is recomputed from the + poles, which is what lets the ``poles_mean_k`` readout swap ``k`` out. + """ + + concept_probs: torch.Tensor + """``(N, k)`` the concept probabilities.""" + unknown_prob: torch.Tensor + """``(N,)`` the unnamed slot's mixing probability.""" + known_pos: torch.Tensor + """``(N, k, d)`` the ``w+`` poles.""" + known_neg: torch.Tensor + """``(N, k, d)`` the ``w-`` poles.""" + unknown_embedding: torch.Tensor + """``(N, u)`` the unnamed slot's mixed embedding ``z_u``.""" + targets: torch.Tensor + """``(N,)`` long: the next event token id.""" + subject_ids: torch.Tensor + """``(N,)`` long.""" + model_pred: torch.Tensor + """``(N,)`` long: the model's own top-1 next event.""" + family_labels: torch.Tensor + """``(N,)`` long: the next token's code family (the CTL target).""" + concept_labels: torch.Tensor + """``(N, k)`` float32 running concept labels (for the CTL bank).""" + concept_observed: torch.Tensor + """``(N, k)`` bool.""" + shard_index: torch.Tensor + """``(N,)`` long: which held-out shard the row came from.""" + concept_names: tuple[str, ...] + n_positions_seen: int + + _TENSORS = ( + "concept_probs", + "unknown_prob", + "known_pos", + "known_neg", + "unknown_embedding", + "targets", + "subject_ids", + "model_pred", + "family_labels", + "concept_labels", + "concept_observed", + "shard_index", + ) + + def __len__(self) -> int: + """Return the number of banked positions.""" + return int(self.targets.numel()) + + def _map(self, fn: Callable[[torch.Tensor], torch.Tensor]) -> ChannelBank: + return _build_bank( + {name: fn(getattr(self, name)) for name in self._TENSORS}, + concept_names=self.concept_names, + n_positions_seen=self.n_positions_seen, + ) + + def to(self, device: str) -> ChannelBank: + """Copy every tensor to ``device``.""" + return self._map(lambda t: t.to(device)) + + def subset(self, idx: torch.Tensor) -> ChannelBank: + """Rows ``idx`` (a long index tensor on this bank's device).""" + return self._map(lambda t: t[idx]) + + @staticmethod + def concat(banks: Sequence[ChannelBank]) -> ChannelBank: + """Stack per-shard banks with the same concept names.""" + if not banks: + raise ValueError("no banks to concatenate") + names = banks[0].concept_names + if any(b.concept_names != names for b in banks): + raise ValueError("banks disagree on concept_names") + return _build_bank( + { + name: torch.cat([getattr(b, name) for b in banks]) + for name in ChannelBank._TENSORS + }, + concept_names=names, + n_positions_seen=sum(b.n_positions_seen for b in banks), + ) + + def mixed(self, idx: torch.Tensor, probs: torch.Tensor) -> torch.Tensor: + """``(n, k * d)`` named embeddings for rows ``idx`` mixed with ``probs``.""" + p = probs.float().unsqueeze(-1) + w_pos = self.known_pos[idx].float() + w_neg = self.known_neg[idx].float() + mixed: torch.Tensor = (p * w_pos + (1.0 - p) * w_neg).flatten(1) + return mixed + + +def _build_bank( + t: dict[str, torch.Tensor], *, concept_names: tuple[str, ...], n_positions_seen: int +) -> ChannelBank: + """Construct a bank from a name-keyed tensor dict (typed field by field).""" + return ChannelBank( + concept_probs=t["concept_probs"], + unknown_prob=t["unknown_prob"], + known_pos=t["known_pos"], + known_neg=t["known_neg"], + unknown_embedding=t["unknown_embedding"], + targets=t["targets"], + subject_ids=t["subject_ids"], + model_pred=t["model_pred"], + family_labels=t["family_labels"], + concept_labels=t["concept_labels"], + concept_observed=t["concept_observed"], + shard_index=t["shard_index"], + concept_names=concept_names, + n_positions_seen=n_positions_seen, + ) + + +def collect_channel_bank( # noqa: PLR0915 -- one linear streaming pass + model: ConceptBottleneckSequenceModel, + events_binned: pl.DataFrame, + vocab: Vocabulary, + *, + concept_labels: ConceptLabelDict, + concept_mask: ConceptLabelDict, + concept_first_times: ConceptLabelDict, + concept_names: Sequence[str], + supervision: ConceptSupervision, + sample_rate: float = 1.0, + seed: int = 0, + num_lanes: int = 8, + chunk_size: int = 256, + device: str = "cpu", + max_positions: int | None = None, + shard_index: int = 0, +) -> ChannelBank: + """One frozen streaming pass; bank the bottleneck pieces per position. + + Scored positions are those with a real, non-padding next-token + target, the same rows the model's own top-1 is computed on. Each is + kept with probability ``sample_rate`` (seeded); if more than + ``max_positions`` survive, a seeded random subset of that size is + kept, so the subsample is random within the shard rather than the + first rows streamed. + """ + model.eval() + bottleneck = model.bottleneck + if not isinstance(bottleneck, ConceptBottleneck): + raise ValueError( + "channel probes need the mixture bottleneck (poles are functions " + f"of the hidden state); got {type(bottleneck).__name__}" + ) + num_concepts = bottleneck.num_concepts + if len(concept_names) != num_concepts: + raise ValueError( + f"{len(concept_names)} concept names but the bottleneck has " + f"{num_concepts} concepts" + ) + gen = torch.Generator().manual_seed(seed) + type_lookup = _build_type_lookup(vocab, device) + patients: Iterator[PatientSequence] = iter_patient_sequences(events_binned, vocab) + sampler = PackedLaneSampler( + patients, num_lanes=num_lanes, chunk_size=chunk_size, reset_prob=0.0 + ) + parts: dict[str, list[torch.Tensor]] = {name: [] for name in ChannelBank._TENSORS} + seen = 0 + checked = False + state = None + with torch.no_grad(): + for chunk in sampler: + chunk = _move_chunk_to_device(chunk, device) # noqa: PLW2901 + hidden, state = model.backbone( + chunk.batch, state=state, reset_mask=chunk.reset_mask + ) + out = bottleneck(hidden) + mix = bottleneck.mixture_parts(hidden) + if not checked: + # The poles must reproduce the forward pass's own mixing. + p = out.concept_probs.unsqueeze(-1) + rebuilt = p * mix.known_pos + (1.0 - p) * mix.known_neg + if not torch.allclose(rebuilt, out.concept_embeddings, atol=1e-4): + raise RuntimeError( + "mixture_parts does not reproduce concept_embeddings" + ) + checked = True + pred = model.lm_head(out.bottleneck).argmax(dim=-1) + valid = chunk.real_mask & (chunk.targets != PAD_ID) + n_valid = int(valid.sum().item()) + seen += n_valid + if n_valid == 0: + continue + if sample_rate < 1.0: + draw = torch.rand(valid.shape, generator=gen) < sample_rate + valid = valid & draw.to(valid.device) + if not bool(valid.any()): + continue + labels, observed = position_running_labels( + chunk, + concept_labels, + concept_mask, + concept_first_times, + supervision=supervision, + num_concepts=num_concepts, + ) + valid_cpu = valid.to(labels.device) + parts["concept_probs"].append(out.concept_probs[valid].half().cpu()) + parts["unknown_prob"].append( + torch.sigmoid(mix.unknown_logit)[valid].half().cpu() + ) + parts["known_pos"].append(mix.known_pos[valid].half().cpu()) + parts["known_neg"].append(mix.known_neg[valid].half().cpu()) + parts["unknown_embedding"].append(out.unknown_embedding[valid].half().cpu()) + parts["targets"].append(chunk.targets[valid].long().cpu()) + parts["subject_ids"].append(chunk.subject_ids[valid].long().cpu()) + parts["model_pred"].append(pred[valid].long().cpu()) + parts["family_labels"].append( + type_lookup[chunk.targets][valid].long().cpu() + ) + parts["concept_labels"].append(labels[valid_cpu].float()) + parts["concept_observed"].append(observed[valid_cpu].bool()) + parts["shard_index"].append( + torch.full((int(valid.sum().item()),), shard_index, dtype=torch.long) + ) + if not parts["targets"]: + raise ValueError("no scored positions collected; empty split?") + bank = _build_bank( + {name: torch.cat(parts[name]) for name in ChannelBank._TENSORS}, + concept_names=tuple(concept_names), + n_positions_seen=seen, + ) + if max_positions is not None and len(bank) > max_positions: + keep = torch.randperm(len(bank), generator=gen)[:max_positions] + bank = bank.subset(keep) + logger.info( + "[probe_channel] shard %d: banked %d of %d scored positions", + shard_index, + len(bank), + seen, + ) + return bank + + +def bank_from_shards( + model: ConceptBottleneckSequenceModel, + vocab: Vocabulary, + binner: QuantileBinner, + config: TrainingConfig, + shard_dir: str | Path, + *, + run_dir: str | Path, + max_shards: int | None, + sample_rate: float, + seed: int, + num_lanes: int, + chunk_size: int, + device: str, + max_positions: int | None, +) -> ChannelBank: + """Bank every held-out shard, stratified: an equal position cap per shard. + + Data preparation matches ``evaluate_interventions`` (normalization, + history recap, sidecars, the run's own concept set, the train-fit + binner) so ``model_head`` reproduces the paper's evaluation rows. + """ + source = getattr(config, "source", "mimic_iv") + task_set = getattr(config, "task_set", "v1") + supervision: ConceptSupervision = getattr(config, "concept_supervision", "visit") + activate_sidecars(shard_dir) + concepts = resolve_concepts_for_run(str(run_dir), source, task_set) + concept_names = [c.name for c in concepts] + paths = shard_paths(shard_dir, max_shards=max_shards) + if not paths: + raise ValueError(f"no shards under {shard_dir}") + per_shard_cap = ( + int(math.ceil(max_positions / len(paths))) if max_positions else None + ) + banks: list[ChannelBank] = [] + for k, path in enumerate(paths): + raw = load_meds_shard(path) + raw = maybe_normalize( + raw, + enabled=getattr(config, "normalize_medications", False), + source=source, + ) + raw = maybe_history_recap(raw, enabled=getattr(config, "history_recap", False)) + concept_labels: ConceptLabelDict + concept_mask: ConceptLabelDict + concept_first_times: ConceptLabelDict + if supervision == "visit": + concept_labels, concept_mask = build_visit_concept_label_dicts( + raw, concepts + ) + concept_first_times = build_visit_concept_first_times(raw, concepts) + else: + concept_labels, concept_mask = build_concept_label_dicts(raw, concepts) + concept_first_times = build_concept_first_times(raw, concepts) + binned = add_value_tokens(raw, binner, source=source) + del raw + banks.append( + collect_channel_bank( + model, + binned, + vocab, + concept_labels=concept_labels, + concept_mask=concept_mask, + concept_first_times=concept_first_times, + concept_names=concept_names, + supervision=supervision, + sample_rate=sample_rate, + seed=seed * 7919 + k, + num_lanes=num_lanes, + chunk_size=chunk_size, + device=device, + max_positions=per_shard_cap, + shard_index=k, + ) + ) + del binned + return ChannelBank.concat(banks) + + +# --------------------------------------------------------------------------- +# Subject-level split +# --------------------------------------------------------------------------- + + +@dataclass(frozen=True) +class SubjectSplit: + """Row indices of one bank, partitioned by subject.""" + + train: torch.Tensor + tune: torch.Tensor + test: torch.Tensor + n_subjects: dict[str, int] + + +def split_by_subject( + subject_ids: torch.Tensor, + *, + seed: int, + train_frac: float = 0.7, + tune_frac: float = 0.1, +) -> SubjectSplit: + """Partition rows so no subject appears in two parts (seeded).""" + if train_frac <= 0 or tune_frac < 0 or train_frac + tune_frac >= 1.0: + raise ValueError("need 0 < train_frac, 0 <= tune_frac, sum < 1") + subjects = torch.unique(subject_ids.cpu()) + if subjects.numel() < 3: + raise ValueError(f"need at least 3 subjects to split, got {subjects.numel()}") + gen = torch.Generator().manual_seed(seed) + order = subjects[torch.randperm(subjects.numel(), generator=gen)] + n = order.numel() + n_train = max(1, int(round(n * train_frac))) + n_tune = max(1, int(round(n * tune_frac))) if tune_frac > 0 else 0 + n_train = min(n_train, n - n_tune - 1) + parts = { + "train": order[:n_train], + "tune": order[n_train : n_train + n_tune], + "test": order[n_train + n_tune :], + } + sids = subject_ids.cpu() + idx = { + name: torch.nonzero(torch.isin(sids, members)).flatten() + for name, members in parts.items() + } + return SubjectSplit( + train=idx["train"], + tune=idx["tune"], + test=idx["test"], + n_subjects={name: int(members.numel()) for name, members in parts.items()}, + ) + + +# --------------------------------------------------------------------------- +# Readouts +# --------------------------------------------------------------------------- + + +FeatureFn = Callable[[ChannelBank, torch.Tensor], torch.Tensor] + + +def readout_features( + bank: ChannelBank, train_idx: torch.Tensor +) -> dict[str, FeatureFn]: + """Return the feature map of every refit readout, keyed by name. + + ``poles_mean_k`` needs the training-split mean of ``k`` per concept, + computed here once so the test split never informs it. + """ + k_mean = bank.concept_probs[train_idx].float().mean(dim=0) + + def k_only(b: ChannelBank, idx: torch.Tensor) -> torch.Tensor: + return b.concept_probs[idx].float() + + def k_plus_u(b: ChannelBank, idx: torch.Tensor) -> torch.Tensor: + return torch.cat( + [b.concept_probs[idx].float(), b.unknown_prob[idx].float().unsqueeze(-1)], + dim=-1, + ) + + def z_named(b: ChannelBank, idx: torch.Tensor) -> torch.Tensor: + return b.mixed(idx, b.concept_probs[idx]) + + def poles_mean_k(b: ChannelBank, idx: torch.Tensor) -> torch.Tensor: + return b.mixed(idx, k_mean.to(b.concept_probs.device).expand(idx.numel(), -1)) + + def poles_both(b: ChannelBank, idx: torch.Tensor) -> torch.Tensor: + return torch.cat( + [b.known_pos[idx].float().flatten(1), b.known_neg[idx].float().flatten(1)], + dim=-1, + ) + + def h_bar(b: ChannelBank, idx: torch.Tensor) -> torch.Tensor: + return torch.cat( + [b.mixed(idx, b.concept_probs[idx]), b.unknown_embedding[idx].float()], + dim=-1, + ) + + return { + "k_only": k_only, + "k_plus_u": k_plus_u, + "z_named": z_named, + "poles_mean_k": poles_mean_k, + "poles_both": poles_both, + "h_bar": h_bar, + } + + +class _StandardizedMLPProbe(nn.Module): + """One hidden layer behind the same fixed standardization as the linear probe.""" + + def __init__( + self, + in_features: int, + out_features: int, + mean: torch.Tensor, + std: torch.Tensor, + hidden: int = 512, + ) -> None: + """Standardize, then ``Linear -> GELU -> Linear``.""" + super().__init__() + self.register_buffer("mean", mean) + self.register_buffer("std", std) + self.net = nn.Sequential( + nn.Linear(in_features, hidden), nn.GELU(), nn.Linear(hidden, out_features) + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Standardize then apply the MLP.""" + mean = cast(torch.Tensor, self.mean) + std = cast(torch.Tensor, self.std) + out: torch.Tensor = self.net((x - mean) / std) + return out + + +def _batches(idx: torch.Tensor, batch_size: int) -> Iterator[torch.Tensor]: + for start in range(0, idx.numel(), batch_size): + yield idx[start : start + batch_size] + + +def _running_stats( + bank: ChannelBank, idx: torch.Tensor, feature_fn: FeatureFn, batch_size: int +) -> tuple[torch.Tensor, torch.Tensor]: + """Per-feature (mean, std) over ``idx`` without materializing every row.""" + total = 0 + s1: torch.Tensor | None = None + s2: torch.Tensor | None = None + for b in _batches(idx, batch_size): + x = feature_fn(bank, b).double() + s1 = x.sum(dim=0) if s1 is None else s1 + x.sum(dim=0) + s2 = (x * x).sum(dim=0) if s2 is None else s2 + (x * x).sum(dim=0) + total += x.shape[0] + assert s1 is not None and s2 is not None # noqa: S101 -- idx is non-empty + mean = s1 / total + var = (s2 / total - mean * mean).clamp_min(0.0) + if total > 1: + var = var * total / (total - 1) + std = var.sqrt().clamp_min(1e-6) + return mean.float().unsqueeze(0), std.float().unsqueeze(0) + + +def fit_readout( + bank: ChannelBank, + train_idx: torch.Tensor, + tune_idx: torch.Tensor, + feature_fn: FeatureFn, + num_classes: int, + *, + mlp: bool = False, + epochs: int = 5, + batch_size: int = 4096, + lr: float = 1e-3, + patience: int = 2, + seed: int = 0, + device: str = "cpu", +) -> tuple[nn.Module, ProbeFitTrace]: + """Adam on the training rows, early-stopped on tuning cross-entropy. + + Features are built per batch from the bank, so the poles are read + once per epoch and never expanded to a full ``(N, k * d)`` matrix. + """ + torch.manual_seed(seed) + if train_idx.numel() < batch_size: + # A small bank needs more optimizer steps per epoch to converge. + batch_size = max(32, train_idx.numel() // 8) + mean, std = _running_stats(bank, train_idx, feature_fn, batch_size) + in_features = int(mean.shape[1]) + head: nn.Module = ( + _StandardizedMLPProbe(in_features, num_classes, mean, std) + if mlp + else _StandardizedLinearProbe(in_features, num_classes, mean, std) + ).to(device) + opt = torch.optim.Adam(head.parameters(), lr=lr) + trace = ProbeFitTrace() + best_state = {k: v.detach().clone() for k, v in head.state_dict().items()} + best = float("inf") + bad = 0 + gen = torch.Generator(device=train_idx.device).manual_seed(seed) + t0 = time.time() + for epoch in range(epochs): + head.train() + perm = train_idx[ + torch.randperm(train_idx.numel(), generator=gen, device=train_idx.device) + ] + for b in _batches(perm, batch_size): + opt.zero_grad() + x = feature_fn(bank, b).to(device) + y = bank.targets[b].to(device) + loss = F.cross_entropy(head(x), y) + loss.backward() # type: ignore[no-untyped-call] + opt.step() + head.eval() + with torch.no_grad(): + tune_loss_sum = 0.0 + for b in _batches(tune_idx, batch_size): + x = feature_fn(bank, b).to(device) + y = bank.targets[b].to(device) + tune_loss_sum += float( + F.cross_entropy(head(x), y, reduction="sum").item() + ) + tune_loss = tune_loss_sum / max(1, tune_idx.numel()) + trace.tuning_loss.append(tune_loss) + if tune_loss < best - 1e-6: + best, bad = tune_loss, 0 + trace.best_epoch = epoch + best_state = {k: v.detach().clone() for k, v in head.state_dict().items()} + else: + bad += 1 + if bad >= patience: + break + head.load_state_dict(best_state) + trace.seconds = time.time() - t0 + return head, trace + + +@torch.no_grad() +def score_readout( + head: nn.Module, + bank: ChannelBank, + idx: torch.Tensor, + feature_fn: FeatureFn, + *, + batch_size: int = 4096, + device: str = "cpu", +) -> torch.Tensor: + """``(n,)`` bool on CPU: exact next-event top-1 hit per row of ``idx``.""" + head.eval() + hits = [] + for b in _batches(idx, batch_size): + x = feature_fn(bank, b).to(device) + y = bank.targets[b].to(device) + hits.append((head(x).argmax(dim=-1) == y).cpu()) + return torch.cat(hits) + + +# --------------------------------------------------------------------------- +# Bootstrap and scoring +# --------------------------------------------------------------------------- + + +@dataclass +class ReadoutScore: + """One readout's held-out accuracy and its intervals.""" + + readout: str + n_positions: int + n_subjects: int + top1_accuracy: float + ci95: tuple[float, float] + retained: float + """``top1_accuracy / model top1_accuracy`` on the same rows.""" + retained_ci95: tuple[float, float] + """Paired subject-clustered interval on ``retained``.""" + completeness_score: float + """Yeh-style: ``(acc - majority) / (model_acc - majority)``.""" + in_features: int | None = None + parameters: int | None = None + fit: ProbeFitTrace | None = None + + +def subject_bootstrap( + subject_ids: np.ndarray, + hits: dict[str, np.ndarray], + reference: str, + *, + n_boot: int = 1000, + seed: int = 0, +) -> dict[str, dict[str, tuple[float, float]]]: + """Subject-clustered percentile intervals on accuracy and on the ratio. + + One subject multiset is drawn per resample and reused for every + readout, so the interval on ``readout / reference`` is paired. + """ + _, inverse = np.unique(subject_ids, return_inverse=True) + n_subj = int(inverse.max()) + 1 + counts = np.bincount(inverse, minlength=n_subj).astype(np.float64) + per_subject = { + name: np.bincount(inverse, weights=h.astype(np.float64), minlength=n_subj) + for name, h in hits.items() + } + rng = np.random.default_rng(seed) + weights = rng.multinomial(n_subj, np.full(n_subj, 1.0 / n_subj), size=n_boot) + weights = weights.astype(np.float64) + denom = weights @ counts + accs = {name: (weights @ per_subject[name]) / denom for name in hits} + ref = accs[reference] + out: dict[str, dict[str, tuple[float, float]]] = {} + for name, acc in accs.items(): + with np.errstate(divide="ignore", invalid="ignore"): + ratio = np.where(ref > 0, acc / ref, np.nan) + out[name] = { + "accuracy": ( + float(np.percentile(acc, 2.5)), + float(np.percentile(acc, 97.5)), + ), + "retained": ( + float(np.nanpercentile(ratio, 2.5)), + float(np.nanpercentile(ratio, 97.5)), + ), + } + return out + + +def _majority_rate(targets: torch.Tensor) -> float: + counts = torch.bincount(targets.cpu()) + return float(counts.max().item() / targets.numel()) + + +def _completeness(acc: float, model_acc: float, majority: float) -> float: + denom = model_acc - majority + return (acc - majority) / denom if denom > 1e-9 else float("nan") + + +# --------------------------------------------------------------------------- +# CTL bridge +# --------------------------------------------------------------------------- + + +def leakage_bank(bank: ChannelBank, idx: torch.Tensor) -> LeakageBank: + """Rows ``idx`` as a :class:`LeakageBank` (named embeddings materialized).""" + sub = bank.subset(idx) + k, d = sub.known_pos.shape[1], sub.known_pos.shape[2] + z = sub.mixed(torch.arange(len(sub), device=sub.targets.device), sub.concept_probs) + return LeakageBank( + sub.concept_probs.half().cpu(), + z.reshape(len(sub), k, d).half().cpu(), + sub.unknown_embedding.half().cpu(), + sub.family_labels.cpu(), + sub.concept_labels.float().cpu(), + sub.concept_observed.bool().cpu(), + sub.concept_names, + n_positions_seen=len(sub), + sample_rate=1.0, + ) + + +def _cap(idx: torch.Tensor, cap: int | None, seed: int) -> torch.Tensor: + if cap is None or idx.numel() <= cap: + return idx + gen = torch.Generator().manual_seed(seed) + keep = torch.randperm(idx.numel(), generator=gen)[:cap].to(idx.device) + return idx[keep] + + +# --------------------------------------------------------------------------- +# Bank persistence +# --------------------------------------------------------------------------- + + +BANK_FORMAT = 1 + + +@dataclass(frozen=True) +class BankMeta: + """What a saved bank carries besides its tensors, so a reload needs no model.""" + + vocab_size: int + notes: tuple[str, ...] + + +def bank_nbytes(bank: ChannelBank) -> int: + """Bytes held by every tensor of ``bank`` (fp16 at rest).""" + return sum( + getattr(bank, name).numel() * getattr(bank, name).element_size() + for name in ChannelBank._TENSORS + ) + + +def default_bank_path(output_json: Path) -> Path: + """``.bank.pt`` next to the JSON.""" + return output_json.with_name(output_json.stem + ".bank.pt") + + +def save_bank( + bank: ChannelBank, path: Path, *, vocab_size: int, notes: Sequence[str] +) -> None: + """Persist the streamed bank (CPU, fp16 at rest) with ``torch.save``. + + The streaming pass is the expensive part of this script (about 45 + minutes for 37 shards on an A100); everything after it is minutes. + Saving right after the pass means a crash in the fits costs a + ``--from-bank`` rerun, not the pass. + """ + path.parent.mkdir(parents=True, exist_ok=True) + torch.save( + { + "format": BANK_FORMAT, + "tensors": {n: getattr(bank, n).cpu() for n in ChannelBank._TENSORS}, + "concept_names": list(bank.concept_names), + "n_positions_seen": int(bank.n_positions_seen), + "vocab_size": int(vocab_size), + "notes": list(notes), + }, + path, + ) + logger.info( + "[probe_channel] saved bank to %s (%.2f GB in memory, %.2f GB on disk)", + path, + bank_nbytes(bank) / 1e9, + path.stat().st_size / 1e9, + ) + + +def load_bank(path: Path) -> tuple[ChannelBank, BankMeta]: + """Load a bank written by :func:`save_bank` onto CPU.""" + blob = torch.load(path, map_location="cpu", weights_only=True) + if not isinstance(blob, dict) or blob.get("format") != BANK_FORMAT: + raise ValueError(f"{path} is not a probe_channel bank (format {BANK_FORMAT})") + tensors = blob["tensors"] + missing = [n for n in ChannelBank._TENSORS if n not in tensors] + if missing: + raise ValueError(f"{path} lacks bank tensors {missing}") + bank = _build_bank( + {n: tensors[n] for n in ChannelBank._TENSORS}, + concept_names=tuple(blob["concept_names"]), + n_positions_seen=int(blob["n_positions_seen"]), + ) + meta = BankMeta(vocab_size=int(blob["vocab_size"]), notes=tuple(blob["notes"])) + logger.info( + "[probe_channel] loaded bank from %s (%d positions, %.2f GB)", + path, + len(bank), + bank_nbytes(bank) / 1e9, + ) + return bank, meta + + +# --------------------------------------------------------------------------- +# Driver +# --------------------------------------------------------------------------- + + +@dataclass +class ProbeOptions: + """Everything that shapes the dump and the fits.""" + + max_positions: int | None = 2_000_000 + max_shards: int | None = None + sample_rate: float = 1.0 + seed: int = 0 + num_lanes: int = 16 + chunk_size: int = 512 + epochs: int = 5 + batch_size: int = 4096 + lr: float = 1e-3 + patience: int = 2 + mlp: bool = False + n_boot: int = 1000 + train_frac: float = 0.7 + tune_frac: float = 0.1 + skip_ctl: bool = False + ctl_max_positions: int | None = 500_000 + ctl_epochs: int = 20 + bank_on_cpu: bool = False + checkpoint: str | None = None + readouts: tuple[str, ...] = READOUT_NAMES + notes: list[str] = field(default_factory=list) + bank_path: Path | None = None + """Where to save the streamed bank (``None``: do not save).""" + from_bank: Path | None = None + """Load this saved bank instead of streaming (the model is not loaded).""" + + +SplitIdx = tuple[torch.Tensor, torch.Tensor, torch.Tensor] + + +def _fit_all_readouts( + bank: ChannelBank, + *, + split_idx: SplitIdx, + hits: dict[str, torch.Tensor], + num_classes: int, + opts: ProbeOptions, + device: str, +) -> dict[str, dict[str, Any]]: + """Fit and score every requested readout; fill ``hits``, return fit metadata.""" + train_idx, tune_idx, test_idx = split_idx + features = readout_features(bank, train_idx) + meta: dict[str, dict[str, Any]] = {} + for name in opts.readouts: + if name not in features: + raise ValueError(f"unknown readout {name!r}; known: {READOUT_NAMES}") + fn = features[name] + head, trace = fit_readout( + bank, + train_idx, + tune_idx, + fn, + num_classes, + mlp=opts.mlp, + epochs=opts.epochs, + batch_size=opts.batch_size, + lr=opts.lr, + patience=opts.patience, + seed=opts.seed, + device=device, + ) + hits[name] = score_readout( + head, bank, test_idx, fn, batch_size=opts.batch_size, device=device + ) + meta[name] = { + "in_features": int(cast(torch.Tensor, head.mean).shape[1]), + "parameters": sum(p.numel() for p in head.parameters()), + "fit": trace, + } + logger.info( + "[probe_channel] %-13s top-1 %.4f in %.1fs", + name, + float(hits[name].float().mean().item()), + trace.seconds, + ) + del head + return meta + + +def _load_model( + run_dir: Path, checkpoint_path: Path, device: str +) -> tuple[ + ConceptBottleneckSequenceModel, + Vocabulary, + QuantileBinner, + TrainingConfig, + list[str], +]: + """Load the run and check it has a mixture bottleneck; return it plus notes.""" + model, vocab, binner, config = load_run( + run_dir, device=device, checkpoint_path=checkpoint_path + ) + if not isinstance(model, ConceptBottleneckSequenceModel): + raise ValueError( + "this probe needs a concept bottleneck; the run's model_kind is " + f"{getattr(config, 'model_kind', 'bottleneck')!r}" + ) + if not isinstance(model.bottleneck, ConceptBottleneck): + raise ValueError( + "channel probes are defined for the mixture bottleneck, whose poles " + "are functions of the hidden state; the run's bottleneck_kind is " + f"{getattr(config, 'bottleneck_kind', 'mixture')!r}" + ) + notes: list[str] = [] + if model.bottleneck.global_pairs: + notes.append( + "global_pairs=True: the poles are input-independent parameters, so " + "poles_mean_k is a constant feature and reads the majority class" + ) + notes.append( + "hazard-head readouts at landmarks are not dumped; see the module docstring" + ) + return model, vocab, binner, config, notes + + +def _score_readouts( + bank: ChannelBank, + *, + split_idx: SplitIdx, + num_classes: int, + opts: ProbeOptions, + device: str, +) -> tuple[dict[str, ReadoutScore], float]: + """Fit every readout and bootstrap; return scores and the majority rate.""" + train_idx, tune_idx, test_idx = split_idx + test_targets = bank.targets[test_idx] + model_hits = (bank.model_pred[test_idx] == test_targets).cpu() + model_acc = float(model_hits.float().mean().item()) + majority = _majority_rate(test_targets) + hits: dict[str, torch.Tensor] = {MODEL_READOUT: model_hits} + meta = _fit_all_readouts( + bank, + split_idx=(train_idx, tune_idx, test_idx), + hits=hits, + num_classes=num_classes, + opts=opts, + device=device, + ) + subject_np = bank.subject_ids[test_idx].cpu().numpy() + hits_np = {name: h.numpy() for name, h in hits.items()} + intervals = subject_bootstrap( + subject_np, hits_np, MODEL_READOUT, n_boot=opts.n_boot, seed=opts.seed + ) + n_test_subjects = int(np.unique(subject_np).size) + scores: dict[str, ReadoutScore] = {} + for name, h in hits_np.items(): + acc = float(h.mean()) + extra = meta.get(name, {}) + scores[name] = ReadoutScore( + readout=name, + n_positions=int(h.size), + n_subjects=n_test_subjects, + top1_accuracy=acc, + ci95=intervals[name]["accuracy"], + retained=acc / model_acc if model_acc > 0 else float("nan"), + retained_ci95=intervals[name]["retained"], + completeness_score=_completeness(acc, model_acc, majority), + in_features=extra.get("in_features"), + parameters=extra.get("parameters"), + fit=extra.get("fit"), + ) + return scores, majority + + +def _ctl_block( + bank: ChannelBank, + *, + split_idx: SplitIdx, + opts: ProbeOptions, + device: str, + bank_device: str, +) -> CTLResult: + """Run the CTL probes on the same subject split; banks stay on ``bank_device``. + + With ``--bank-on-cpu`` the three leakage banks stay in host memory too + and :func:`compute_ctl` moves each batch to ``device``; moving them to + the model's device here is what the flag exists to avoid. + """ + train_idx, tune_idx, test_idx = split_idx + cap = opts.ctl_max_positions + return compute_ctl( + leakage_bank(bank, _cap(train_idx, cap, opts.seed + 11)).to(bank_device), + leakage_bank(bank, _cap(tune_idx, cap, opts.seed + 12)).to(bank_device), + leakage_bank(bank, _cap(test_idx, cap, opts.seed + 13)).to(bank_device), + epochs=opts.ctl_epochs, + batch_size=opts.batch_size, + patience=opts.patience, + seed=opts.seed, + device=device, + ) + + +def _write_payload(path: Path | None, payload: dict[str, Any]) -> None: + """Write ``payload`` atomically (temp file, then rename); no-op without a path.""" + if path is None: + return + path.parent.mkdir(parents=True, exist_ok=True) + tmp = path.with_name(path.name + ".tmp") + tmp.write_text(json.dumps(payload, indent=2)) + tmp.replace(path) + logger.info("[probe_channel] wrote %s", path) + + +def _check_resumable( + existing: dict[str, Any], bank: ChannelBank, split_idx: SplitIdx, opts: ProbeOptions +) -> None: + """Refuse to fill ``ctl`` into a payload fit on a different bank or split.""" + if existing.get("ctl") is not None: + raise ValueError("payload already holds a ctl block") + train_idx, tune_idx, test_idx = split_idx + expected = { + "n_positions.bank": (existing["n_positions"]["bank"], len(bank)), + "n_positions.train": (existing["n_positions"]["train"], train_idx.numel()), + "n_positions.tune": (existing["n_positions"]["tune"], tune_idx.numel()), + "n_positions.test": (existing["n_positions"]["test"], test_idx.numel()), + "protocol.seed": (existing["protocol"]["seed"], opts.seed), + "protocol.train_frac": (existing["protocol"]["train_frac"], opts.train_frac), + "protocol.tune_frac": (existing["protocol"]["tune_frac"], opts.tune_frac), + "concept_names": (tuple(existing["concept_names"]), bank.concept_names), + } + bad = {k: v for k, v in expected.items() if v[0] != v[1]} + if bad: + raise ValueError( + "existing payload does not match this bank/split, refusing to fill " + f"its ctl block: {bad} (json value, this run)" + ) + + +def run_channel_probes( + run_dir: str | Path, + held_out_shard_dir: str | Path, + *, + options: ProbeOptions | None = None, + device: str | None = None, + output_json: Path | None = None, + existing: dict[str, Any] | None = None, +) -> dict[str, Any]: + """Load the run, dump the bank, fit every readout, return the JSON payload. + + Staged so that a crash late in the run costs as little as possible: + the bank is saved to ``options.bank_path`` right after streaming + (``options.from_bank`` reloads it and skips the model entirely); with + ``output_json`` the readout payload is written as soon as the readouts + finish, before the CTL probes run, and updated with the ``ctl`` block + at the end. ``existing`` is a payload written earlier without a + ``ctl`` block: the readouts are not refit, only ``ctl`` is filled in. + """ + opts = options or ProbeOptions() + device = device or ("cuda" if torch.cuda.is_available() else "cpu") + run_dir = Path(run_dir) + checkpoint_path = run_dir / (opts.checkpoint or "checkpoint_best.pt") + + if opts.from_bank is not None: + bank, meta = load_bank(Path(opts.from_bank)) + vocab_size = meta.vocab_size + notes = [*opts.notes, *meta.notes] + else: + model, vocab, binner, config, model_notes = _load_model( + run_dir, checkpoint_path, device + ) + vocab_size = len(vocab) + notes = [*opts.notes, *model_notes] + bank = bank_from_shards( + model, + vocab, + binner, + config, + held_out_shard_dir, + run_dir=run_dir, + max_shards=opts.max_shards, + sample_rate=opts.sample_rate, + seed=opts.seed, + num_lanes=opts.num_lanes, + chunk_size=opts.chunk_size, + device=device, + max_positions=opts.max_positions, + ) + del model + if opts.bank_path is not None: + save_bank( + bank, Path(opts.bank_path), vocab_size=vocab_size, notes=model_notes + ) + logger.info( + "[probe_channel] bank: %d positions of %d seen, %d subjects, %d shards, %.2f GB", + len(bank), + bank.n_positions_seen, + int(torch.unique(bank.subject_ids).numel()), + int(torch.unique(bank.shard_index).numel()), + bank_nbytes(bank) / 1e9, + ) + split = split_by_subject( + bank.subject_ids, + seed=opts.seed, + train_frac=opts.train_frac, + tune_frac=opts.tune_frac, + ) + bank_device = "cpu" if opts.bank_on_cpu else device + bank = bank.to(bank_device) + split_idx: SplitIdx = ( + split.train.to(bank_device), + split.tune.to(bank_device), + split.test.to(bank_device), + ) + train_idx, tune_idx, test_idx = split_idx + + if existing is not None: + _check_resumable(existing, bank, split_idx, opts) + payload = existing + logger.info("[probe_channel] readouts already fit; filling the ctl block only") + else: + scores, majority = _score_readouts( + bank, split_idx=split_idx, num_classes=vocab_size, opts=opts, device=device + ) + payload = { + "run_dir": str(run_dir), + "checkpoint": str(checkpoint_path), + "held_out_shard_dir": str(held_out_shard_dir), + "concept_names": list(bank.concept_names), + "protocol": { + "readout": "mlp" if opts.mlp else "linear", + "target": "exact next event over the full vocabulary (top-1)", + "split": "by subject", + "train_frac": opts.train_frac, + "tune_frac": opts.tune_frac, + "seed": opts.seed, + "epochs": opts.epochs, + "batch_size": opts.batch_size, + "lr": opts.lr, + "n_boot": opts.n_boot, + "max_positions": opts.max_positions, + "sample_rate": opts.sample_rate, + "max_shards": opts.max_shards, + "ctl_max_positions": opts.ctl_max_positions, + }, + "n_positions": { + "seen": bank.n_positions_seen, + "bank": len(bank), + "train": int(train_idx.numel()), + "tune": int(tune_idx.numel()), + "test": int(test_idx.numel()), + }, + "n_subjects": { + "bank": int(torch.unique(bank.subject_ids).numel()), + **split.n_subjects, + }, + "n_shards": int(torch.unique(bank.shard_index).numel()), + "vocab_size": vocab_size, + "majority_class_accuracy": majority, + "model_top1_accuracy_all_banked": float( + (bank.model_pred == bank.targets).float().mean().item() + ), + "model": asdict(scores[MODEL_READOUT]), + "readouts": {name: asdict(scores[name]) for name in opts.readouts}, + "ctl": None, + "hazards": None, + "notes": notes, + } + _write_payload(output_json, payload) + + if not opts.skip_ctl: + ctl = _ctl_block( + bank, split_idx=split_idx, opts=opts, device=device, bank_device=bank_device + ) + payload["ctl"] = asdict(ctl) + _write_payload(output_json, payload) + return payload + + +def markdown_table(payload: dict[str, Any]) -> str: + """Render a small table of every readout for the terminal.""" + rows = [ + "| readout | top-1 | 95% CI | retained | completeness | features |", + "|---|---|---|---|---|---|", + ] + entries = [payload["model"], *payload["readouts"].values()] + for s in entries: + lo, hi = s["ci95"] + feats = "" if s["in_features"] is None else str(s["in_features"]) + rows.append( + f"| {s['readout']} | {100 * s['top1_accuracy']:.2f} | " + f"[{100 * lo:.2f}, {100 * hi:.2f}] | {s['retained']:.3f} | " + f"{s['completeness_score']:.3f} | {feats} |" + ) + return "\n".join(rows) + + +def existing_payload(path: Path, *, overwrite: bool) -> dict[str, Any] | None: + """Append-only gate on the output JSON. + + Returns ``None`` when there is nothing to build on (no file, or + ``--overwrite``). A file with a ``ctl`` block is complete and is + refused. A file without one (readouts written, CTL not yet) is + returned so the run fills only its ``ctl`` block. + """ + if overwrite or not path.exists(): + return None + loaded: dict[str, Any] = json.loads(path.read_text()) + if loaded.get("ctl") is not None: + raise FileExistsError( + f"{path} already holds readouts and a ctl block; pass --overwrite " + "to replace it (channel probes are append-only by default)" + ) + logger.info("[probe_channel] %s exists without a ctl block; will fill it", path) + return loaded + + +def _parse_args(argv: Sequence[str] | None = None) -> argparse.Namespace: + parser = argparse.ArgumentParser(description=(__doc__ or "").split("Usage::")[0]) + parser.add_argument("--run-dir", required=True) + parser.add_argument("--held-out-shard-dir", required=True) + parser.add_argument("--output-json", required=True) + parser.add_argument("--checkpoint", default=None) + parser.add_argument("--max-positions", type=int, default=2_000_000) + parser.add_argument("--max-shards", type=int, default=None) + parser.add_argument("--sample-rate", type=float, default=1.0) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--num-lanes", type=int, default=16) + parser.add_argument("--chunk-size", type=int, default=512) + parser.add_argument("--epochs", type=int, default=5) + parser.add_argument("--batch-size", type=int, default=4096) + parser.add_argument("--lr", type=float, default=1e-3) + parser.add_argument("--patience", type=int, default=2) + parser.add_argument("--mlp", action="store_true", help="one-hidden-layer readout") + parser.add_argument("--n-boot", type=int, default=1000) + parser.add_argument("--train-frac", type=float, default=0.7) + parser.add_argument("--tune-frac", type=float, default=0.1) + parser.add_argument("--skip-ctl", action="store_true") + parser.add_argument("--ctl-max-positions", type=int, default=500_000) + parser.add_argument("--ctl-epochs", type=int, default=20) + parser.add_argument( + "--bank-on-cpu", + action="store_true", + help="keep the dump in host memory and move batches to the GPU", + ) + parser.add_argument( + "--bank-path", + default=None, + help=( + "where to save the streamed bank (torch.save, fp16); default " + ".bank.pt next to the JSON" + ), + ) + parser.add_argument( + "--from-bank", + default=None, + help=( + "load a saved bank instead of streaming the shards (the checkpoint " + "is not loaded); --run-dir and --held-out-shard-dir are only recorded" + ), + ) + parser.add_argument("--readouts", nargs="*", default=list(READOUT_NAMES)) + parser.add_argument( + "--overwrite", + action="store_true", + help=( + "replace an existing --output-json (and bank file) instead of " + "refusing; without it a JSON lacking a ctl block is filled in place" + ), + ) + return parser.parse_args(argv) + + +def main(argv: Sequence[str] | None = None) -> None: + """Command-line entry point.""" + args = _parse_args(argv) + out = Path(args.output_json) + existing = existing_payload(out, overwrite=args.overwrite) + if existing is not None and args.skip_ctl: + raise ValueError( + f"{out} exists without a ctl block and --skip-ctl leaves nothing to " + "fill; drop --skip-ctl or pass --overwrite to refit the readouts" + ) + from_bank = Path(args.from_bank) if args.from_bank else None + bank_path: Path | None = None + if from_bank is None: + bank_path = Path(args.bank_path) if args.bank_path else default_bank_path(out) + if bank_path.exists() and not args.overwrite: + raise FileExistsError( + f"{bank_path} exists; pass --from-bank {bank_path} to reuse it " + "or --overwrite to stream the shards again" + ) + options = ProbeOptions( + max_positions=args.max_positions or None, + max_shards=args.max_shards, + sample_rate=args.sample_rate, + seed=args.seed, + num_lanes=args.num_lanes, + chunk_size=args.chunk_size, + epochs=args.epochs, + batch_size=args.batch_size, + lr=args.lr, + patience=args.patience, + mlp=args.mlp, + n_boot=args.n_boot, + train_frac=args.train_frac, + tune_frac=args.tune_frac, + skip_ctl=args.skip_ctl, + ctl_max_positions=args.ctl_max_positions or None, + ctl_epochs=args.ctl_epochs, + bank_on_cpu=args.bank_on_cpu, + checkpoint=args.checkpoint, + readouts=tuple(args.readouts), + bank_path=bank_path, + from_bank=from_bank, + ) + payload = run_channel_probes( + args.run_dir, + args.held_out_shard_dir, + options=options, + output_json=out, + existing=existing, + ) + print(markdown_table(payload)) + + +if __name__ == "__main__": + logging.basicConfig( + level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s" + ) + main() diff --git a/tests/odyssey/inference/test_interventions_hazard.py b/tests/odyssey/inference/test_interventions_hazard.py new file mode 100644 index 00000000..a27f37f6 --- /dev/null +++ b/tests/odyssey/inference/test_interventions_hazard.py @@ -0,0 +1,477 @@ +"""Hazard-head scoring of the label-override test (CPU, tiny fixtures).""" + +import json +from datetime import datetime, timedelta +from pathlib import Path + +import numpy as np +import polars as pl +import pytest +import torch + +import odyssey.inference.interventions as interventions_module +from odyssey.data.alert_events import ALERT_EVENTS, all_event_times +from odyssey.data.concepts import concepts_for_source +from odyssey.data.value_binning import add_value_tokens +from odyssey.data.vocabulary import Vocabulary +from odyssey.inference.alerts import ( + HORIZONS_HOURS, + _visit_starts, + collect_model_scores, + score_alerts, +) +from odyssey.inference.interventions import ( + HazardLandmarkScorer, + InterventionResult, + LandmarkRiskTable, + evaluate_interventions, + hazard_paired_summary, + hazard_summary, + landmark_labels, + paired_mean_delta, + result_to_json, + run_streaming_intervention, + subject_bootstrap_means, +) +from odyssey.models.backbones.tiny_gru import TinyGRUBackbone +from odyssey.models.sequence_model import ( + BaselineSequenceModel, + ConceptBottleneckSequenceModel, +) +from odyssey.models.time_to_event import DEFAULT_TIME_BIN_EDGES_HOURS +from odyssey.training.data import ( + build_concept_first_times, + build_concept_label_dicts, +) +from odyssey.training.train import TrainingConfig + + +T0 = datetime(2024, 1, 1) +CONCEPTS = concepts_for_source("mimic_iv", task_set="v1") +EVENT_NAMES = [a.name for a in ALERT_EVENTS] + +# The JSON keys interventions.py wrote before hazard scoring existed. +OLD_KEYS = [ + "mode", + "n_predictions", + "top1_accuracy", + "mean_task_loss", + "top1_by_code_type", + "n_by_code_type", + "n_intervened_positions", + "uncertain_band", + "mean_abs_displacement", + "calibrated_tau", + "n_replaced_by_concept", + "mean_abs_displacement_by_concept", + "calibration_gamma", +] + + +def _events(n_subjects: int = 16) -> pl.DataFrame: + """Hourly heart rates; even subjects start norepinephrine at hour 14. + + Every fourth subject also gets an ICU admission at hour 6, so the + vasopressor and ICU cells have both outcome classes at 8 h and 24 h. + """ + rows: list[tuple[int, str, datetime, float | None, int]] = [] + for sid in range(1, n_subjects + 1): + hadm = 1000 + sid + for h in range(24): + hr = 130.0 if sid % 2 == 0 and h >= 12 else 80.0 + rows.append((sid, "LAB//220045//bpm", T0 + timedelta(hours=h), hr, hadm)) + if sid % 2 == 0: + rows.append( + ( + sid, + "MEDICATION//norepinephrine//Administered", + T0 + timedelta(hours=14), + None, + hadm, + ) + ) + if sid % 4 == 0: + rows.append( + (sid, "ICU_ADMISSION//MICU", T0 + timedelta(hours=6), None, hadm) + ) + return pl.DataFrame( + rows, + schema={ + "subject_id": pl.Int64, + "code": pl.Utf8, + "time": pl.Datetime, + "numeric_value": pl.Float32, + "hadm_id": pl.Int64, + }, + orient="row", + ) + + +def _model(vocab_size: int, with_heads: bool = True) -> ConceptBottleneckSequenceModel: + torch.manual_seed(0) + return ConceptBottleneckSequenceModel( + backbone=TinyGRUBackbone( + vocab_size=vocab_size, hidden_size=8, num_layers=1, padding_idx=0 + ), + vocab_size=vocab_size, + num_concepts=len(CONCEPTS), + embedding_dim=4, + padding_idx=0, + time_bin_edges=DEFAULT_TIME_BIN_EDGES_HOURS, + event_names=EVENT_NAMES if with_heads else None, + ) + + +class _Fixture: + """One synthetic held-out split, its model, labels and alert targets.""" + + def __init__(self) -> None: + self.raw = _events() + self.binned = add_value_tokens(self.raw) + self.vocab = Vocabulary.build(self.binned["code"].to_list(), min_count=1) + self.model = _model(len(self.vocab)) + self.labels, self.mask = build_concept_label_dicts(self.raw, CONCEPTS) + self.first_times = build_concept_first_times(self.raw, CONCEPTS) + self.times = all_event_times(self.raw, ALERT_EVENTS, "mimic_iv") + self.visit_start = _visit_starts(self.raw) + + def scorer(self) -> HazardLandmarkScorer: + assert self.model.event_heads is not None + return HazardLandmarkScorer( + self.model.event_heads, ALERT_EVENTS, self.visit_start + ) + + def run( + self, mode: str, *, hazard: bool + ) -> tuple[InterventionResult, LandmarkRiskTable | None]: + scorer = self.scorer() if hazard else None + result = run_streaming_intervention( + self.model, + self.binned, + self.vocab, + self.labels, + self.mask, + mode=mode, + concept_first_times=self.first_times, + supervision="stay", + num_lanes=2, + chunk_size=16, + device="cpu", + seed=0, + hazard_scorer=scorer, + ) + return result, (scorer.table() if scorer is not None else None) + + +@pytest.fixture(scope="module") +def fx() -> _Fixture: + return _Fixture() + + +# --------------------------------------------------------------------------- +# (i) hazard scoring off leaves the old output untouched +# --------------------------------------------------------------------------- + + +def test_hazard_off_keeps_the_old_schema_and_on_does_not_move_the_numbers( + fx: _Fixture, +) -> None: + plain, _ = fx.run("truth", hazard=False) + scored, table = fx.run("truth", hazard=True) + assert list(result_to_json(plain)) == OLD_KEYS + assert plain.hazard is None and plain.hazard_paired is None + # Reading the heads is a side readout: the serialised next-event + # numbers are byte for byte the same (NaN entries compare as text). + assert json.dumps(result_to_json(scored)) == json.dumps(result_to_json(plain)) + assert table is not None and table.n_rows > 0 + + +def test_cli_writes_the_old_json_without_the_flag_and_passes_it_through( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path +) -> None: + seen: dict[str, object] = {} + hazard = {"death": {"8h": {"auroc": None, "mean_risk": 0.1}}} + + def fake_evaluate(*_args: object, **kwargs: object) -> list[InterventionResult]: + seen.update(kwargs) + base = InterventionResult("none", 1, 1.0, 0.5) + if kwargs["hazard_heads"]: + base = InterventionResult( + "none", 1, 1.0, 0.5, hazard=hazard, hazard_paired={} + ) + return [base] + + monkeypatch.setattr(interventions_module, "evaluate_interventions", fake_evaluate) + argv = [ + "prog", + "--run-dir", + "/runs/x", + "--held-out-shard-dir", + "/data/held_out", + "--output-json", + ] + off = tmp_path / "off.json" + monkeypatch.setattr("sys.argv", [*argv, str(off)]) + interventions_module._main() + written = json.loads(off.read_text()) + assert list(written[0]) == OLD_KEYS + assert seen["hazard_heads"] is False and seen["hazard_boot"] == 1000 + + on = tmp_path / "on.json" + monkeypatch.setattr( + "sys.argv", [*argv, str(on), "--hazard-heads", "--hazard-boot", "7"] + ) + interventions_module._main() + written = json.loads(on.read_text()) + assert list(written[0]) == [*OLD_KEYS, "hazard", "hazard_paired"] + assert written[0]["hazard"] == hazard + assert seen["hazard_heads"] is True and seen["hazard_boot"] == 7 + + +# --------------------------------------------------------------------------- +# (ii) the landmark rows are alerts.py's rows +# --------------------------------------------------------------------------- + + +def test_landmark_rows_and_counts_match_the_alert_protocol(fx: _Fixture) -> None: + rows = collect_model_scores( + fx.model, + fx.binned, + fx.vocab, + [c.name for c in CONCEPTS], + ALERT_EVENTS, + visit_start=fx.visit_start, + landmark_hours=4.0, + num_lanes=2, + chunk_size=16, + device="cpu", + horizons=HORIZONS_HOURS, + ) + alert_metrics = { + (m.event, m.horizon_hours): m + for m in score_alerts(rows, fx.times, horizons=HORIZONS_HOURS) + if m.scorer == "hazard" + } + assert alert_metrics, "the alert protocol scored no hazard cell" + + _, table = fx.run("none", hazard=True) + assert table is not None + mine = set( + zip( + table.subject_ids.tolist(), + table.visit_ids.tolist(), + table.time_hours.tolist(), + ) + ) + for event in EVENT_NAMES: + theirs = {(r.subject_id, r.visit_id, r.time_hours) for r in rows[event]} + assert mine == theirs, event + + summary = hazard_summary(table, landmark_labels(table, fx.times)) + for event in EVENT_NAMES: + for h in HORIZONS_HOURS: + cell = summary[event][f"{h:g}h"] + metric = alert_metrics.get((event, h)) + if metric is None: + # alerts.py skips a single-class cell; we report it unscoreable + assert cell["auroc"] is None + continue + assert cell["n_at_risk"] == metric.n_at_risk, (event, h) + assert cell["n_positive"] == metric.n_positive, (event, h) + assert cell["n_censored"] == metric.n_censored, (event, h) + assert cell["auroc"] == pytest.approx(metric.auroc, abs=1e-6), (event, h) + + +# --------------------------------------------------------------------------- +# (iii) paired bootstrap: interval covers the point, collapses when identical +# --------------------------------------------------------------------------- + + +def test_paired_interval_covers_the_point_and_is_zero_width_when_identical( + fx: _Fixture, +) -> None: + _, none_table = fx.run("none", hazard=True) + _, truth_table = fx.run("truth", hazard=True) + assert none_table is not None and truth_table is not None + labels = landmark_labels(none_table, fx.times) + + paired = hazard_paired_summary(truth_table, none_table, labels, n_boot=200, seed=0) + scored = 0 + for event in EVENT_NAMES: + for h in HORIZONS_HOURS: + cell = paired[event][f"{h:g}h"] + mean = cell["mean_risk"] + if mean is None: + assert cell["n_at_risk"] == 0, (event, h) + continue + assert mean["ci_low"] <= mean["point"] <= mean["ci_high"], (event, h) + if cell["auroc"] is not None and cell["auroc"]["ci_low"] is not None: + auc = cell["auroc"] + assert auc["ci_low"] <= auc["point"] <= auc["ci_high"], (event, h) + scored += 1 + assert scored > 0, "no two-class cell to check the AUROC interval on" + + same = hazard_paired_summary(none_table, none_table, labels, n_boot=200, seed=0) + for event in EVENT_NAMES: + for h in HORIZONS_HOURS: + cell = same[event][f"{h:g}h"] + if cell["mean_risk"] is None: + continue + assert cell["mean_risk"] == { + "point": 0.0, + "ci_low": 0.0, + "ci_high": 0.0, + "separated": False, + } + if cell["auroc"] is not None and cell["auroc"]["ci_low"] is not None: + assert cell["auroc"]["point"] == 0.0 + assert cell["auroc"]["ci_low"] == 0.0 == cell["auroc"]["ci_high"] + assert cell["auroc"]["separated"] is False + + +# --------------------------------------------------------------------------- +# (iv) resampling moves whole subjects +# --------------------------------------------------------------------------- + + +def test_bootstrap_resamples_whole_subjects() -> None: + # Subject 1 has four rows of 1.0, subject 2 two rows of 10.0. Drawing + # two subjects with replacement gives {1,1}, {1,2}, {2,2}: means of + # 1, 4 (pooled: 24/6) or 10. Row resampling would produce other values. + diff = np.array([1.0, 1.0, 1.0, 1.0, 10.0, 10.0]) + subjects = np.array([1, 1, 1, 1, 2, 2]) + boots = subject_bootstrap_means(diff, subjects, n_boot=300, seed=0, block=7) + seen = set(np.round(boots, 12).tolist()) + assert seen <= {1.0, 4.0, 10.0} + assert seen == {1.0, 4.0, 10.0} + + delta = paired_mean_delta(diff, np.zeros_like(diff), subjects, n_boot=300, seed=0) + assert delta.point == pytest.approx(4.0) + assert delta.n_rows == 6 and delta.n_subjects == 2 + assert delta.ci_low == 1.0 and delta.ci_high == 10.0 + assert delta.separated + + +def test_subject_bootstrap_is_deterministic_per_seed() -> None: + diff = np.arange(20, dtype=float) + subjects = np.repeat(np.arange(5), 4) + a = subject_bootstrap_means(diff, subjects, n_boot=50, seed=3) + b = subject_bootstrap_means(diff, subjects, n_boot=50, seed=3) + c = subject_bootstrap_means(diff, subjects, n_boot=50, seed=4) + assert np.array_equal(a, b) + assert not np.array_equal(a, c) + + +# --------------------------------------------------------------------------- +# orchestration: refusals and the attached blocks +# --------------------------------------------------------------------------- + + +def test_hazard_scoring_refuses_a_run_without_heads_before_reading_shards( + monkeypatch: pytest.MonkeyPatch, +) -> None: + model = _model(10, with_heads=False) + config = TrainingConfig( + train_shard_dir="/train", tuning_shard_dir="/tuning", output_dir="/out" + ) + monkeypatch.setattr( + interventions_module, "load_run", lambda *a, **k: (model, None, None, config) + ) + + def _boom(*_args: object, **_kwargs: object) -> None: + raise AssertionError("must not read shards before the hazard gate fires") + + monkeypatch.setattr(interventions_module, "load_meds_shards", _boom) + with pytest.raises(ValueError, match="hazard heads"): + evaluate_interventions("/runs/x", "/data/held_out", hazard_heads=True) + + +def test_hazard_scoring_refuses_a_baseline_model( + monkeypatch: pytest.MonkeyPatch, +) -> None: + model = BaselineSequenceModel( + TinyGRUBackbone(vocab_size=10, hidden_size=4), vocab_size=10 + ) + config = TrainingConfig( + train_shard_dir="/train", + tuning_shard_dir="/tuning", + output_dir="/out", + model_kind="baseline", + ) + monkeypatch.setattr( + interventions_module, "load_run", lambda *a, **k: (model, None, None, config) + ) + with pytest.raises(ValueError, match="needs a concept bottleneck"): + evaluate_interventions("/runs/x", "/data/held_out", hazard_heads=True) + + +def test_evaluate_interventions_attaches_hazard_and_paired_blocks( + fx: _Fixture, monkeypatch: pytest.MonkeyPatch, tmp_path: Path +) -> None: + config = TrainingConfig( + train_shard_dir="/train", + tuning_shard_dir="/tuning", + output_dir="/out", + source="mimic_iv", + task_set="v1", + concept_supervision="stay", + ) + monkeypatch.setattr( + interventions_module, + "load_run", + lambda *a, **k: (fx.model, fx.vocab, None, config), + ) + monkeypatch.setattr( + interventions_module, "load_meds_shards", lambda *a, **k: fx.raw.clone() + ) + held_out = tmp_path / "held_out" + held_out.mkdir() + results = evaluate_interventions( + tmp_path, + held_out, + modes=["none", "truth", "flip"], + num_lanes=2, + chunk_size=16, + device="cpu", + hazard_heads=True, + hazard_boot=20, + ) + by_mode = {r.mode: r for r in results} + assert set(by_mode) == {"none", "truth", "flip"} + for r in results: + assert r.hazard is not None and set(r.hazard) == set(EVENT_NAMES) + for event in EVENT_NAMES: + assert set(r.hazard[event]) == {f"{h:g}h" for h in HORIZONS_HOURS} + assert by_mode["none"].hazard_paired == {} + assert set(by_mode["truth"].hazard_paired or {}) == { + "truth_minus_none", + "truth_minus_flip", + } + assert set(by_mode["flip"].hazard_paired or {}) == {"flip_minus_none"} + cell = (by_mode["truth"].hazard_paired or {})["truth_minus_none"] + vaso = cell["vasopressor_start"]["8h"] + assert ( + vaso["n_at_risk"] + == by_mode["truth"].hazard["vasopressor_start"]["8h"]["n_at_risk"] + ) + assert vaso["mean_risk"]["point"] == pytest.approx( + by_mode["truth"].hazard["vasopressor_start"]["8h"]["mean_risk"] + - by_mode["none"].hazard["vasopressor_start"]["8h"]["mean_risk"], + abs=1e-6, + ) + # The whole thing round-trips through the CLI serialiser. + dumped = json.loads(json.dumps([result_to_json(r) for r in results])) + assert list(dumped[0]) == [*OLD_KEYS, "hazard", "hazard_paired"] + + # Without the flag on the same fixture: no hazard keys, same numbers. + plain = evaluate_interventions( + tmp_path, + held_out, + modes=["none"], + num_lanes=2, + chunk_size=16, + device="cpu", + ) + assert list(result_to_json(plain[0])) == OLD_KEYS + assert plain[0].top1_accuracy == by_mode["none"].top1_accuracy + assert plain[0].mean_task_loss == by_mode["none"].mean_task_loss diff --git a/tests/odyssey/inference/test_leakage.py b/tests/odyssey/inference/test_leakage.py index b971115d..8245dbca 100644 --- a/tests/odyssey/inference/test_leakage.py +++ b/tests/odyssey/inference/test_leakage.py @@ -11,6 +11,8 @@ from odyssey.inference.leakage import ( LeakageBank, _fit_categorical_probe, + _predict_in_batches, + _random_projection, _score_categorical_probe, compute_ctl, compute_icl, @@ -204,6 +206,77 @@ def test_compute_ctl_capacity_control_on_pure_noise() -> None: assert result.ctl_vs_projected_cross_entropy >= -0.08 +# --------------------------------------------------------------------------- +# Device discipline: bank on CPU, probe on a named device +# --------------------------------------------------------------------------- +# +# The 2026-09-25 incident: probe_channel.py streamed 37 shards with the +# bank kept on CPU (--bank-on-cpu) and the model on cuda:0, then compute_ctl +# multiplied the CPU bank by a projection on the other device and lost the +# pass. A CPU-only machine cannot reproduce two devices, so these tests pin +# the contract the fix relies on: the projection is created on the bank's +# device by the device-safe helper, and every fit/score path accepts a bank +# on CPU together with an explicit ``device`` argument. + + +def test_random_projection_lands_on_the_requested_device_and_is_orthonormal() -> None: + proj = _random_projection(NUM_CONCEPTS, 24, seed=0, device="cpu") + assert proj.device.type == "cpu" + assert proj.shape == (NUM_CONCEPTS, 24) + assert torch.allclose(proj @ proj.T, torch.eye(NUM_CONCEPTS), atol=1e-5) + # Drawn on CPU before the move, so the matrix does not depend on the + # device spelling (str or torch.device) and is reproducible per seed. + again = _random_projection(NUM_CONCEPTS, 24, seed=0, device=torch.device("cpu")) + assert torch.equal(proj, again) + other = _random_projection(NUM_CONCEPTS, 24, seed=1) + assert not torch.equal(proj, other) + + +def test_compute_ctl_accepts_a_cpu_bank_with_an_explicit_probe_device() -> None: + train, tune, held = _three_splits(leak_slot=0) + for bank in (train, tune, held): + assert bank.concept_probs.device.type == "cpu" + result = compute_ctl(train, tune, held, seed=0, device="cpu", **FIT_KW) + + for score in ( + result.probs_only, + result.embeddings_only, + result.unknown_only, + result.probs_projected, + ): + assert score.n == len(held) + assert 0.0 <= score.accuracy <= 1.0 + assert score.cross_entropy == score.cross_entropy # not NaN + assert result.probs_projected.parameters == result.embeddings_only.parameters + assert result.embeddings_only.accuracy > 0.9 + + +def test_predict_in_batches_matches_one_shot_and_returns_cpu() -> None: + g = torch.Generator().manual_seed(0) + n = 500 + y = torch.randint(0, NUM_CLASSES, (n,), generator=g) + x = (y.float().unsqueeze(-1) - 2.0) * 4.0 + 0.1 * torch.randn(n, 3, generator=g) + head, _ = _fit_categorical_probe( + 3, + NUM_CLASSES, + train_x=x.half(), + train_y=y, + tune_x=x.half(), + tune_y=y, + device="cpu", + epochs=5, + batch_size=64, + ) + head.eval() + with torch.no_grad(): + one_shot = head(x.half().float()) + batched = _predict_in_batches(head, x.half(), batch_size=64) + assert batched.device.type == "cpu" + assert torch.allclose(batched, one_shot, atol=1e-5) + acc, ce = _score_categorical_probe(head, x.half(), y) + assert 0.0 <= acc <= 1.0 and ce > 0.0 + + # --------------------------------------------------------------------------- # ICL # --------------------------------------------------------------------------- diff --git a/tests/odyssey/inference/test_legacy_concept_pins.py b/tests/odyssey/inference/test_legacy_concept_pins.py index 51d08cb7..cbf7238a 100644 --- a/tests/odyssey/inference/test_legacy_concept_pins.py +++ b/tests/odyssey/inference/test_legacy_concept_pins.py @@ -7,13 +7,22 @@ import pytest import torch +from odyssey.data.alert_events import hazard_events_for from odyssey.data.concepts import concepts_for_source from odyssey.inference.legacy_concept_pins import ( + HAZARD_NUM_BINS, LEGACY_CONCEPT_PINS, + LEGACY_EVENT_PINS, + PIN_FILENAME, check_concept_count, + check_event_count, checkpoint_num_concepts, + checkpoint_num_events, pinned_concept_names, + pinned_event_names, + read_run_pins, resolve_concepts_for_run, + write_run_pins, ) @@ -156,3 +165,148 @@ def test_no_checkpoint_caller_pairs_the_registry_with_a_loaded_model() -> None: "registry; use resolve_concepts_for_run(run_dir, source, task_set) " f"instead: {offenders}" ) + + +# --------------------------------------------------------------------------- +# Hazard event pins: the run's own file, then the legacy table, then a +# readable refusal. +# --------------------------------------------------------------------------- + +EICU_PRE_SOFA_EVENTS = ( + "vasopressor_start", + "icu_admission", + "acute_kidney_injury", + "death", + "readmission_30d", +) + + +def _event_state(n_events: int, *, mlp: bool) -> dict[str, object]: + rows = n_events * HAZARD_NUM_BINS + if mlp: + return { + "event_heads.proj.0.weight": torch.zeros(256, 64), + "event_heads.proj.0.bias": torch.zeros(256), + "event_heads.proj.2.weight": torch.zeros(rows, 256), + "event_heads.proj.2.bias": torch.zeros(rows), + } + return { + "event_heads.proj.weight": torch.zeros(rows, 64), + "event_heads.proj.bias": torch.zeros(rows), + } + + +def test_hazard_num_bins_matches_the_eicu_flagship_checkpoint() -> None: + # eicu_full_v10's event_heads.proj.2.weight is (80, 256): 5 events x 16. + assert HAZARD_NUM_BINS == 16 + assert 5 * HAZARD_NUM_BINS == 80 + + +def test_pin_file_round_trip(tmp_path: Path) -> None: + concepts = ("tachycardia", "shock") + events = ("death", "readmission_30d") + path = write_run_pins(tmp_path, concept_names=concepts, event_names=events) + + assert path == tmp_path / PIN_FILENAME + assert read_run_pins(tmp_path) == {"concepts": concepts, "events": events} + assert pinned_concept_names(str(tmp_path)) == concepts + assert pinned_event_names(str(tmp_path)) == events + # Order is what the file says, not what the registry says. + assert pinned_event_names(str(tmp_path) + "/") == events + + +def test_pin_file_without_hazard_heads_stores_null_events(tmp_path: Path) -> None: + write_run_pins(tmp_path, concept_names=["a"], event_names=None) + assert read_run_pins(tmp_path) == {"concepts": ("a",)} + assert pinned_event_names(str(tmp_path)) is None + assert pinned_concept_names(str(tmp_path)) == ("a",) + + +def test_pin_file_beats_the_legacy_table(tmp_path: Path) -> None: + run_dir = tmp_path / "eicu_full_v10" + run_dir.mkdir() + assert pinned_event_names(str(run_dir)) == EICU_PRE_SOFA_EVENTS + write_run_pins(run_dir, concept_names=["x"], event_names=["death"]) + assert pinned_event_names(str(run_dir)) == ("death",) + assert pinned_concept_names(str(run_dir)) == ("x",) + + +def test_missing_pin_file_is_unpinned_and_malformed_one_raises(tmp_path: Path) -> None: + assert read_run_pins(tmp_path) == {} + assert read_run_pins(tmp_path / "never_written") == {} + (tmp_path / PIN_FILENAME).write_text('{"events": "death"}') + with pytest.raises(ValueError, match="list of names"): + read_run_pins(tmp_path) + + +def test_legacy_event_pin_for_the_eicu_flagship_is_the_historical_order() -> None: + # Printed from alert_events_for("v3", source="eicu") at cdbd4e7, the + # training commit: five heads, no sepsis3, readmission_30d last. + assert pinned_event_names("eicu_full_v10") == EICU_PRE_SOFA_EVENTS + assert pinned_event_names("/home/x/runs/eicu_full_v10/") == EICU_PRE_SOFA_EVENTS + assert pinned_event_names("runs/full_run_v10") is None + assert pinned_event_names("runs/never_trained") is None + + +def test_legacy_event_pins_cover_every_pre_sofa_eicu_run() -> None: + for name in ( + "eicu_full_v10", + "eicu_full_L_v10", + "eicu_full_ADD_v10", + "eicu_full_DEC_v10", + "eicu_full_DEC_v11", + "eicu_full_DEC_v12", + "eicu_full_DEC_v12_steer", + ): + assert LEGACY_EVENT_PINS[name] == EICU_PRE_SOFA_EVENTS, name + # Their concept sets froze at the same commits, so they are pinned + # to the same 26 names as the flagship. + assert LEGACY_CONCEPT_PINS[name] == LEGACY_CONCEPT_PINS["eicu_full_v10"] + + +def test_legacy_event_pins_are_a_strict_subset_of_todays_eicu_events() -> None: + # Today's list has sepsis3 spliced in before readmission_30d; the pin + # keeps the old order, which is why a width match alone is not enough. + today = [a.name for a in hazard_events_for("v3", source="eicu")] + assert "sepsis3" in today + assert [n for n in today if n != "sepsis3"] == list(EICU_PRE_SOFA_EVENTS) + + +def test_checkpoint_num_events_reads_the_output_layer() -> None: + assert checkpoint_num_events(_event_state(5, mlp=True)) == 5 + assert checkpoint_num_events(_event_state(6, mlp=False)) == 6 + assert checkpoint_num_events({}) is None + assert checkpoint_num_events(_state(16)) is None + + +def test_checkpoint_num_events_refuses_a_foreign_bin_count() -> None: + state = {"event_heads.proj.weight": torch.zeros(5 * HAZARD_NUM_BINS + 1, 64)} + with pytest.raises(ValueError, match="not a multiple"): + checkpoint_num_events(state) + + +def test_check_event_count_passes_when_counts_agree() -> None: + check_event_count("runs/x", _event_state(5, mlp=True), ["e"] * 5) + check_event_count("runs/x", _event_state(6, mlp=False), ["e"] * 6) + # No hazard heads at all: nothing to compare against. + check_event_count("runs/x", {}, ["e"] * 6) + check_event_count("runs/x", _state(16), None) + + +def test_check_event_count_names_the_run_both_counts_and_the_fix() -> None: + # The eICU flagship against today's registry: 80 rows = 5 events, 6 today. + state = _event_state(5, mlp=True) + with pytest.raises(ValueError) as excinfo: + check_event_count("/home/x/runs/eicu_full_v10", state, ["e"] * 6) + msg = str(excinfo.value) + assert msg.startswith("eicu_full_v10: ") + assert "5 events" in msg + assert "resolves 6" in msg + assert PIN_FILENAME in msg + assert "LEGACY_EVENT_PINS" in msg + assert "legacy_concept_pins.py" in msg + + +def test_check_event_count_treats_no_names_as_zero() -> None: + with pytest.raises(ValueError, match="resolves 0"): + check_event_count("runs/x", _event_state(5, mlp=True), None) diff --git a/tests/odyssey/inference/test_run_inference.py b/tests/odyssey/inference/test_run_inference.py index 1410950d..47cafe37 100644 --- a/tests/odyssey/inference/test_run_inference.py +++ b/tests/odyssey/inference/test_run_inference.py @@ -20,6 +20,12 @@ from odyssey.data.streaming import PackedLaneSampler from odyssey.data.value_binning import QuantileBinner from odyssey.data.vocabulary import Vocabulary +from odyssey.inference.legacy_concept_pins import ( + HAZARD_NUM_BINS, + PIN_FILENAME, + pinned_event_names, + write_run_pins, +) from odyssey.inference.run_inference import ( InferenceResults, _block_set_hits, @@ -877,7 +883,7 @@ def test_value_embeddings_flag_ignores_the_backbone_merge_attention( torch.save({"model": keys}, run_dir / "checkpoint_final.pt") seen = {} - def fake_build_model(cfg, *, vocab_size, num_concepts): # noqa: ARG001 + def fake_build_model(cfg, *, vocab_size, num_concepts, event_names=None): # noqa: ARG001 seen["value_embeddings"] = cfg.value_embeddings seen["event_head_hidden"] = cfg.event_head_hidden seen["concept_global_pairs"] = cfg.concept_global_pairs @@ -966,7 +972,7 @@ def test_unknown_dim_round_trips_for_the_non_global_pairs_unequal_width_case( seen = {} - def fake_build_model(cfg, *, vocab_size, num_concepts): # noqa: ARG001 + def fake_build_model(cfg, *, vocab_size, num_concepts, event_names=None): # noqa: ARG001 seen["unknown_dim"] = cfg.unknown_dim seen["concept_global_pairs"] = cfg.concept_global_pairs raise RuntimeError("stop here") @@ -1008,6 +1014,7 @@ def _write_transformer_run( event_hazards: bool = False, task_set: str = "v1", auxiliary_event_names: tuple[str, ...] = (), + event_names: list[str] | None = None, ) -> Path: """Build a real (CPU-only) transformer BaselineSequenceModel run dir. @@ -1015,13 +1022,17 @@ def _write_transformer_run( caller's behavior exactly (no event_heads.* keys at all) -- only set it to exercise hazard-head reconstruction (see test_load_run_reconstructs_widened_event_heads below). + ``event_names`` builds the hazard heads for an explicit list instead of + the task set's, standing in for a checkpoint whose registry has since + grown (the event-pin tests below). """ torch.manual_seed(0) - event_names = ( - [a.name for a in hazard_events_for(task_set, auxiliary_event_names)] - if event_hazards - else None - ) + if event_names is None: + event_names = ( + [a.name for a in hazard_events_for(task_set, auxiliary_event_names)] + if event_hazards + else None + ) model = BaselineSequenceModel( backbone=TransformerBackbone( vocab_size=len(vocab), @@ -1112,6 +1123,68 @@ def test_load_run_auxiliary_event_names_empty_matches_pre_existing_behavior( assert model.event_heads.event_names == [a.name for a in hazard_events_for("v1")] +# --------------------------------------------------------------------------- +# Hazard event pins: a checkpoint trained before the registry grew an event +# (the eICU sepsis3 head after PR #222) must rebuild its heads at its own +# width and order, or refuse with a readable message. +# --------------------------------------------------------------------------- + + +def test_load_run_honours_the_event_pin_file(tmp_path: Path) -> None: + """A run_pins.json event list rebuilds the heads at the checkpoint's width.""" + vocab = _vocab() + today = [a.name for a in hazard_events_for("v1")] + # Drop one from the middle, as sepsis3 sits in the middle of today's + # eICU list: the surviving heads keep their historical order. + historical = [n for i, n in enumerate(today) if i != 1] + run_dir = _write_transformer_run( + tmp_path, vocab, value_head=False, event_hazards=True, event_names=historical + ) + write_run_pins(run_dir, concept_names=[], event_names=historical) + + model, _, _, _ = load_run(run_dir, device="cpu") + + assert model.event_heads is not None + assert model.event_heads.event_names == historical + assert model.event_heads.proj.out_features == len(historical) * HAZARD_NUM_BINS + + +def test_load_run_refuses_an_unpinned_event_width_mismatch(tmp_path: Path) -> None: + """Without a pin, the mismatch is one ValueError, not a load_state_dict wall.""" + vocab = _vocab() + today = [a.name for a in hazard_events_for("v1")] + (tmp_path / "eicu_like_v10").mkdir() + run_dir = _write_transformer_run( + tmp_path / "eicu_like_v10", + vocab, + value_head=False, + event_hazards=True, + event_names=today[:-1], + ) + assert not (run_dir / PIN_FILENAME).exists() + + with pytest.raises(ValueError, match="eicu_like_v10: checkpoint has hazard"): + load_run(run_dir, device="cpu") + + +def test_load_run_unpinned_run_with_agreeing_counts_is_unchanged( + tmp_path: Path, +) -> None: + """The MIMIC path: no pin file, counts agree, today's list is used as before.""" + vocab = _vocab() + (tmp_path / "full_run_v10").mkdir() + run_dir = _write_transformer_run( + tmp_path / "full_run_v10", vocab, value_head=False, event_hazards=True + ) + assert not (run_dir / PIN_FILENAME).exists() + assert pinned_event_names(str(run_dir)) is None + + model, _, _, _ = load_run(run_dir, device="cpu") + + assert model.event_heads is not None + assert model.event_heads.event_names == [a.name for a in hazard_events_for("v1")] + + def test_load_run_reconstructs_value_head(tmp_path: Path) -> None: vocab = _vocab() _write_transformer_run(tmp_path, vocab, value_head=True, value_fourier=True) diff --git a/tests/odyssey/training/test_train.py b/tests/odyssey/training/test_train.py index 8fbf354e..f08614da 100644 --- a/tests/odyssey/training/test_train.py +++ b/tests/odyssey/training/test_train.py @@ -18,6 +18,7 @@ from odyssey.data.streaming import PackedLaneSampler from odyssey.data.types import AuxiliaryInputs, ClinicalSequenceBatch from odyssey.data.vocabulary import Vocabulary +from odyssey.inference.legacy_concept_pins import HAZARD_NUM_BINS from odyssey.models.backbones.base import TimeAwareState from odyssey.models.backbones.hybrid import HybridState from odyssey.models.backbones.tiny_gru import TinyGRUBackbone @@ -40,6 +41,7 @@ build_model, build_objective, evaluate_streaming, + hazard_event_names_for, train, ) @@ -413,6 +415,71 @@ def test_build_model_widens_event_heads_with_auxiliary_events() -> None: ] +def test_build_model_event_names_override_sets_head_width_and_order() -> None: + """A pinned list builds heads at its own width, in its own order.""" + config = TrainingConfig( + train_shard_dir="/train", + tuning_shard_dir="/tuning", + output_dir="/out", + backbone="transformer", + hidden_size=16, + num_hidden_layers=1, + attn_num_heads=4, + task_set="v3", + source="eicu", + event_hazards=True, + ) + pinned = [ + "vasopressor_start", + "icu_admission", + "acute_kidney_injury", + "death", + "readmission_30d", + ] + assert len(hazard_events_for("v3", source="eicu")) == 6 + + model = build_model(config, vocab_size=50, num_concepts=5, event_names=pinned) + + assert model.event_heads is not None + assert model.event_heads.event_names == pinned + assert model.event_heads.proj.out_features == 5 * HAZARD_NUM_BINS + assert model.event_heads.num_bins == HAZARD_NUM_BINS + + +def test_build_model_event_names_override_is_ignored_without_hazard_heads() -> None: + config = TrainingConfig( + train_shard_dir="/train", + tuning_shard_dir="/tuning", + output_dir="/out", + backbone="transformer", + hidden_size=16, + num_hidden_layers=1, + attn_num_heads=4, + event_hazards=False, + ) + assert hazard_event_names_for(config) is None + model = build_model(config, vocab_size=50, num_concepts=5, event_names=["death"]) + assert model.event_heads is None + + +def test_hazard_event_names_for_matches_the_unpinned_build() -> None: + """load_run's pre-load check and build_model must see the same list.""" + config = TrainingConfig( + train_shard_dir="/train", + tuning_shard_dir="/tuning", + output_dir="/out", + backbone="transformer", + hidden_size=16, + num_hidden_layers=1, + attn_num_heads=4, + task_set="v1", + auxiliary_event_names=("vasopressor",), + ) + model = build_model(config, vocab_size=50, num_concepts=5) + assert model.event_heads is not None + assert hazard_event_names_for(config) == model.event_heads.event_names + + def test_build_model_auxiliary_event_names_default_matches_pre_existing_behavior() -> ( None ): diff --git a/tests/scripts/gemini/test_run_sh_rebuttal_steps.py b/tests/scripts/gemini/test_run_sh_rebuttal_steps.py new file mode 100644 index 00000000..0aff9183 --- /dev/null +++ b/tests/scripts/gemini/test_run_sh_rebuttal_steps.py @@ -0,0 +1,137 @@ +"""The rebuttal steps' export whitelists match what their scripts write. + +run.sh exports aggregate JSON through a top-level key whitelist that +refuses anything unknown. PR #238 was the second time a script gained +a field and the export refused on the node, where a retry costs a +session inside the secure environment. These tests read the shipped +run.sh and pin each step's key list to the script that produces the +file, so the mismatch fails here. +""" + +import json +import re +import sys +from pathlib import Path + +import numpy as np +import polars as pl +import pytest + +from scripts import alerts_cis, cohort_counts, panel_coverage + + +REPO = Path(__file__).resolve().parents[3] +RUN_SH = REPO / "scripts" / "gemini" / "run.sh" +STEPS = ("alerts-cis", "panel-coverage", "cohort-counts") + + +def _keys(variable: str) -> set[str]: + source = RUN_SH.read_text() + match = re.search(rf'^\s*{variable}="([^"]+)"', source, re.MULTILINE) + assert match, f"{variable} not found in run.sh" + return set(match.group(1).split()) + + +def _export_keys_used(export_name: str) -> str: + """Return the key-list variable the export call for ``export_name`` passes.""" + source = RUN_SH.read_text() + match = re.search( + rf'_export_aggregate_json \\\n\s*"scripts/gemini/out/evals/\$\{{run_name\}}[^"]*{export_name}\.json" "\$OUTPUT_JSON" \\\n\s*"\$(\w+)"', + source, + ) + assert match, f"no _export_aggregate_json call for {export_name}" + return match.group(1) + + +def test_steps_are_documented_dispatched_and_listed_as_unknown_step_hints() -> None: + source = RUN_SH.read_text() + usage = re.search( + r"^# Usage.*\n#\s+scripts/gemini/run.sh \[(.*)\]$", source, re.MULTILINE + ) + assert usage + for step in STEPS: + assert f"{step} " in usage.group(1) + assert re.search( + rf'^\s+{re.escape(step)}\) run_\w+ "\$\{{2:-\}}" ;;$', source, re.MULTILINE + ) + assert re.search(rf"^# {re.escape(step)} $", source, re.MULTILINE) + assert re.search(rf"unknown step: .*\b{re.escape(step)}\b", source) + + +def test_panel_coverage_whitelist_matches_the_script() -> None: + assert _export_keys_used("panel_coverage") == "PANEL_COVERAGE_JSON_KEYS" + assert _keys("PANEL_COVERAGE_JSON_KEYS") == set(panel_coverage.OUTPUT_KEYS) + assert set(panel_coverage.panel_coverage("gemini")) == set( + panel_coverage.OUTPUT_KEYS + ) + + +def test_cohort_counts_whitelist_matches_the_script() -> None: + assert _export_keys_used("cohort_counts") == "COHORT_COUNTS_JSON_KEYS" + assert _keys("COHORT_COUNTS_JSON_KEYS") == set(cohort_counts.OUTPUT_KEYS) + + +def test_alerts_cis_whitelist_matches_what_main_writes( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + assert _export_keys_used("alerts_cis") == "ALERTS_CIS_JSON_KEYS" + rng = np.random.default_rng(0) + sids = np.repeat(np.arange(30), 3) + y = np.repeat((np.arange(30) % 3 == 0).astype(float), 3) + noise = rng.uniform(0, 1, len(sids)) + dump = tmp_path / "alerts_rows_allshards.parquet" + pl.DataFrame( + { + "event": ["death"] * len(sids), + "subject_id": sids, + "visit_id": sids, + "time_hours": np.arange(len(sids), dtype=float), + "y@8h": y, + "hazard@8h": np.where(y == 1, 0.4 + 0.5 * noise, 0.1 + 0.5 * noise), + "gbm@8h": np.where(y == 1, 0.3 + 0.6 * noise, 0.1 + 0.6 * noise), + } + ).write_parquet(dump) + out = tmp_path / "alerts_cis_allshards.json" + monkeypatch.setattr( + sys, + "argv", + [ + "alerts_cis", + "--dump", + str(dump), + "--output-json", + str(out), + "--scorers", + "hazard", + "gbm", + "--n-boot", + "20", + ], + ) + alerts_cis.main() + written = json.loads(out.read_text()) + assert set(written) == _keys("ALERTS_CIS_JSON_KEYS") + assert written["n_rows_in_dumps"] == written["n_rows_scored"] == 90 + assert written["n_subjects_in_dumps"] == written["n_subjects_scored"] == 30 + (row,) = written["summary"] + assert row["event"] == "death" and row["horizon_hours"] == 8.0 + assert ( + row["n_at_risk"] == 90 and row["n_positive"] == 30 and row["n_subjects"] == 30 + ) + assert row["hazard_auroc"] is not None and row["gbm_auroc"] is not None + delta = row["hazard_minus_gbm_auroc"] + assert set(delta) == {"point", "ci_low", "ci_high", "separated"} + assert delta["ci_low"] <= delta["point"] <= delta["ci_high"] + # The export validator's forbidden patient-level keys never appear. + forbidden = {"subject_id", "subject_ids", "rows", "per_subject"} + + def scan(node: object) -> None: + if isinstance(node, dict): + assert not forbidden & set(node) + for value in node.values(): + scan(value) + elif isinstance(node, list): + for value in node: + scan(value) + + scan(written) diff --git a/tests/scripts/test_cohort_counts.py b/tests/scripts/test_cohort_counts.py new file mode 100644 index 00000000..19ba1c38 --- /dev/null +++ b/tests/scripts/test_cohort_counts.py @@ -0,0 +1,302 @@ +"""Tests for the aggregate cohort description (scripts/cohort_counts.py). + +A tiny MIMIC-shaped MEDS split: 24 subjects, one three-day admission +each, half of them readmitted ten days later, half of them dying, one +ICU admission. Every count the test reads back is either at or above +the suppression threshold by construction, or deliberately below it to +check the marker. +""" + +import json +import sys +from datetime import datetime, timedelta +from pathlib import Path + +import polars as pl +import pytest + +from odyssey.inference.alerts import READMISSION_HORIZONS_HOURS +from scripts import cohort_counts +from scripts.cohort_counts import ( + OUTPUT_KEYS, + READMISSION_WINDOW_HOURS, + SUPPRESSED, + build_report, + suppress, +) + + +T0 = datetime(2150, 3, 1) +BIRTH = datetime(2100, 1, 1) +N_SUBJECTS = 24 +SCHEMA = { + "subject_id": pl.Int64, + "time": pl.Datetime("us"), + "code": pl.Utf8, + "numeric_value": pl.Float32, + "hadm_id": pl.Int64, +} + + +def _rows( + subject: int, +) -> list[tuple[int, datetime | None, str, float | None, int | None]]: + start = T0 + timedelta(days=subject) + hadm = 1000 + subject + rows: list[tuple[int, datetime | None, str, float | None, int | None]] = [ + (subject, None, "GENDER//F" if subject % 2 == 0 else "GENDER//M", None, None), + (subject, BIRTH, "MEDS_BIRTH", None, None), + (subject, start, "HOSPITAL_ADMISSION//EW EMER.", None, hadm), + (subject, start + timedelta(hours=2), "LAB//50912//mg/dL", 1.0, hadm), + (subject, start + timedelta(days=3), "HOSPITAL_DISCHARGE//HOME", None, hadm), + ] + if subject == 5: + rows.append( + (subject, start + timedelta(hours=6), "ICU_ADMISSION//MICU", None, hadm) + ) + if subject % 2 == 0: + # A second admission ten days after the first one's discharge. + again = start + timedelta(days=13) + rows += [ + (subject, again, "HOSPITAL_ADMISSION//EW EMER.", None, hadm + 500), + ( + subject, + again + timedelta(days=1), + "HOSPITAL_DISCHARGE//HOME", + None, + hadm + 500, + ), + ] + if subject % 2 == 1: + rows.append((subject, start + timedelta(days=20), "MEDS_DEATH", None, None)) + return rows + + +def _shard(subjects: range) -> pl.DataFrame: + rows = [r for s in subjects for r in _rows(s)] + return pl.DataFrame(rows, schema=SCHEMA, orient="row") + + +@pytest.fixture +def split_dir(tmp_path: Path) -> Path: + shard_dir = tmp_path / "held_out" + shard_dir.mkdir() + _shard(range(0, 12)).write_parquet(shard_dir / "0.parquet") + _shard(range(12, N_SUBJECTS)).write_parquet(shard_dir / "1.parquet") + return shard_dir + + +def _keys_in(node: object) -> set[str]: + """Every dict key anywhere in a JSON-like tree (the export validator's scan).""" + keys: set[str] = set() + if isinstance(node, dict): + keys |= set(node) + for value in node.values(): + keys |= _keys_in(value) + elif isinstance(node, list): + for value in node: + keys |= _keys_in(value) + return keys + + +def test_readmission_window_matches_the_alerts_horizon() -> None: + assert READMISSION_HORIZONS_HOURS[-1] == READMISSION_WINDOW_HOURS + + +def test_suppress_marks_small_cells() -> None: + assert suppress(10) == 10 + assert suppress(9) == SUPPRESSED == "<10" + assert suppress(0) == "<10" + + +def test_report_counts_admissions_los_years_sex_age_and_events(split_dir: Path) -> None: + report = build_report( + {"held_out": split_dir}, + source="mimic_iv", + task_set="v2", + normalize_medications=False, + with_events=True, + hospital_parquet=None, + max_shards=None, + ) + assert set(report) == set(OUTPUT_KEYS) + assert report["events_dropped"] == [] + split = report["splits"]["held_out"] + assert split["n_shards_read"] == split["n_shards_total"] == 2 + assert split["n_subjects"] == N_SUBJECTS + assert split["n_admissions"] == N_SUBJECTS + N_SUBJECTS // 2 + assert split["n_hospitals"] is None + assert "not available" in split["hospitals_note"] + assert split["admission_years"] == {"min": 2150, "max": 2150} + # 24 three-day stays and 12 one-day stays: median 3, IQR 1 to 3. + assert split["los_days"]["n"] == 36 + assert split["los_days"]["median"] == 3.0 + assert split["los_days"]["q1"] == 1.0 + assert split["los_days"]["q3"] == 3.0 + assert split["sex"]["n_subjects_with_sex"] == N_SUBJECTS + assert split["sex"]["counts"] == {"F": 12, "M": 12} + assert 50.0 <= split["age_years"]["median"] <= 50.2 + events = split["events"] + assert set(events) == { + "vasopressor_start", + "icu_admission", + "acute_kidney_injury", + "death", + "sepsis3", + "readmission_30d", + } + assert events["death"]["n_subjects_positive"] == 12 + assert events["death"]["prevalence_per_subject"] == 0.5 + assert events["death"]["n_admissions_positive"] is None # subject-scoped + assert events["readmission_30d"]["n_subjects_positive"] == 12 + assert events["readmission_30d"]["n_admissions_positive"] == 12 + assert events["readmission_30d"]["prevalence_per_admission"] == round(12 / 36, 4) + # One ICU admission: the count and its rate are both suppressed. + assert events["icu_admission"]["n_subjects_positive"] == "<10" + assert events["icu_admission"]["prevalence_per_subject"] is None + assert events["acute_kidney_injury"]["n_subjects_positive"] == "<10" + # The pooled block equals the single split here. + assert report["all_splits"]["n_subjects"] == N_SUBJECTS + assert not _keys_in(report) & {"subject_id", "subject_ids", "rows", "per_subject"} + + +def test_hospitals_come_from_the_metadata_table( + split_dir: Path, tmp_path: Path +) -> None: + table = tmp_path / "hadm_id_hospital.parquet" + pl.DataFrame( + { + "hadm_id": [1000 + s for s in range(N_SUBJECTS)], + "hospital_num": [101 if s < 12 else 202 for s in range(N_SUBJECTS)], + } + ).write_parquet(table) + report = build_report( + {"held_out": split_dir}, + source="mimic_iv", + task_set="v1", + normalize_medications=False, + with_events=False, + hospital_parquet=table, + max_shards=1, + ) + assert report["hospital_metadata"] == "hadm_id_hospital.parquet" + assert report["max_shards"] == 1 + split = report["splits"]["held_out"] + assert split["n_shards_read"] == 1 and split["n_shards_total"] == 2 + assert split["n_hospitals"] == 1 # shard 0 holds subjects 0-11: only hospital 101 + assert split["events"] is None + assert report["events"] == [] + + +def test_small_split_is_suppressed_end_to_end(tmp_path: Path) -> None: + shard_dir = tmp_path / "tiny" + shard_dir.mkdir() + _shard(range(0, 4)).write_parquet(shard_dir / "0.parquet") + report = build_report( + {"tuning": shard_dir}, + source="mimic_iv", + task_set="v1", + normalize_medications=False, + with_events=True, + hospital_parquet=None, + max_shards=None, + ) + split = report["splits"]["tuning"] + assert split["n_subjects"] == "<10" + assert split["n_admissions"] == "<10" + assert split["admission_years"] is None + assert split["los_days"] is None + assert split["age_years"] is None + assert split["sex"]["counts"] == {"F": "<10", "M": "<10"} + assert split["events"]["death"]["n_subjects_positive"] == "<10" + assert split["events"]["death"]["prevalence_per_subject"] is None + + +def test_gemini_source_without_sex_or_birth_reports_not_available( + tmp_path: Path, +) -> None: + shard_dir = tmp_path / "gemini" + shard_dir.mkdir() + start = T0 + rows = [] + for s in range(12): + hadm = 7000 + s + rows += [ + (s, start, "ADMISSION", None, hadm), + (s, start + timedelta(hours=1), "VITALS//3027018//", 80.0, hadm), + (s, start + timedelta(days=2), "DISCHARGE", None, hadm), + ] + pl.DataFrame(rows, schema=SCHEMA, orient="row").write_parquet( + shard_dir / "shard_0000.parquet" + ) + report = build_report( + {"train": shard_dir}, + source="gemini", + task_set="v3", + normalize_medications=False, + with_events=True, + hospital_parquet=None, + max_shards=None, + ) + assert report["events_dropped"] == ["sepsis3"] + split = report["splits"]["train"] + assert split["n_subjects"] == 12 + assert split["sex"] is None and "not available" in split["sex_note"] + assert split["age_years"] is None and "not available" in split["age_note"] + assert split["los_days"]["median"] == 2.0 + + +def test_main_reads_the_run_config_for_source_and_splits( + split_dir: Path, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + run_dir = tmp_path / "run" + run_dir.mkdir() + (run_dir / "config.json").write_text( + json.dumps( + { + "source": "mimic_iv", + "task_set": "v1", + "normalize_medications": False, + "train_shard_dir": str(split_dir), + "tuning_shard_dir": str(tmp_path / "missing"), + } + ) + ) + out = tmp_path / "cohort_counts.json" + monkeypatch.setattr( + sys, + "argv", + [ + "cohort_counts", + "--run-dir", + str(run_dir), + "--split", + f"held_out={split_dir}", + "--no-events", + "--output-json", + str(out), + ], + ) + cohort_counts.main() + report = json.loads(out.read_text()) + assert report["source"] == "mimic_iv" and report["task_set"] == "v1" + assert set(report["splits"]) == {"train", "held_out"} # missing tuning dir skipped + assert report["all_splits"]["n_subjects"] == 2 * N_SUBJECTS + + +def test_main_requires_a_source_without_a_run_dir( + split_dir: Path, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr( + sys, + "argv", + [ + "cohort_counts", + "--split", + f"x={split_dir}", + "--output-json", + str(tmp_path / "o.json"), + ], + ) + with pytest.raises(SystemExit, match="--source"): + cohort_counts.main() diff --git a/tests/scripts/test_compare_runs_paired.py b/tests/scripts/test_compare_runs_paired.py new file mode 100644 index 00000000..1bb8f67f --- /dev/null +++ b/tests/scripts/test_compare_runs_paired.py @@ -0,0 +1,250 @@ +"""Tests for the paired two-run comparison (scripts/compare_runs_paired.py).""" + +import json +import sys +from pathlib import Path + +import numpy as np +import polars as pl +import pytest + +from scripts import compare_runs_paired as crp +from scripts.compare_runs_paired import ( + RowSetMismatchError, + compare, + inference_block, + join_event, + markdown_table, + score_pair, + summarise, +) +from tests.scripts.test_alerts_cis import _dump + + +def _pair(shift: float = 0.0, seed: int = 0) -> tuple[pl.DataFrame, pl.DataFrame]: + """Two dumps in the alerts writer format on identical rows. + + Dump B's hazard is A's plus ``shift`` on the positives, so a positive + shift makes B strictly better separated; ``shift == 0`` gives + identical scores. B's GBM is a different refit (re-drawn noise). + """ + a = _dump(seed=seed) + b = _dump(seed=seed + 1) # same rows and labels, different gbm noise + hazard_b = np.clip(a["hazard@8h"].to_numpy() + shift * a["y@8h"].to_numpy(), 0, 1) + b = b.with_columns(pl.Series("hazard@8h", hazard_b)) + assert a.select("subject_id", "visit_id", "time_hours").equals( + b.select("subject_id", "visit_id", "time_hours") + ) + return a, b + + +def _joined(a: pl.DataFrame, b: pl.DataFrame) -> pl.DataFrame: + joined, counts = join_event(a, b, "death", label_a="a", label_b="b") + assert counts["n_unmatched_a"] == 0 and counts["n_unmatched_b"] == 0 + return joined + + +def test_shifted_scores_give_an_interval_that_excludes_zero() -> None: + a, b = _pair(shift=0.3) + cell = score_pair(_joined(a, b), "hazard", 8.0, n_boot=200, seed=0) + assert cell is not None and "unscoreable" not in cell + assert cell["n_at_risk"] == 240 and cell["n_positive"] == 80 + assert cell["n_subjects"] == 60 + assert cell["auroc_b"] > cell["auroc_a"] + delta = cell["delta_b_minus_a"] + assert delta["point"] == pytest.approx(cell["auroc_b"] - cell["auroc_a"]) + assert delta["ci_low"] > 0 and delta["separated"] is True + assert summarise({"death@8h": cell}) == { + "b_beats_a": ["death@8h"], + "a_beats_b": [], + "ties": [], + } + + +def test_identical_scores_give_a_zero_point_and_zero_width_interval() -> None: + a, b = _pair(shift=0.0) + cell = score_pair(_joined(a, b), "hazard", 8.0, n_boot=100, seed=0) + assert cell is not None + assert cell["auroc_a"] == cell["auroc_b"] + delta = cell["delta_b_minus_a"] + assert delta["point"] == 0.0 + assert delta["ci_low"] == 0.0 and delta["ci_high"] == 0.0 + assert delta["separated"] is False + assert summarise({"death@8h": cell})["ties"] == ["death@8h"] + + +def test_dropped_rows_are_refused_with_both_counts_in_the_message() -> None: + a, b = _pair() + b_short = b.head(228) # 5% of 240 rows dropped + with pytest.raises(RowSetMismatchError) as excinfo: + join_event(a, b_short, "death", label_a="bottleneck", label_b="baseline") + msg = str(excinfo.value) + assert "bottleneck has 240 rows" in msg + assert "baseline has 228 rows" in msg + assert "the join keeps 228" in msg + assert "12 of bottleneck's 240 rows are unmatched" in msg + + +def test_one_extra_row_in_a_large_dump_is_within_tolerance() -> None: + a = _dump(n_subjects=300, rows_per=4) # 1200 rows + b = a.head(1199) # one row missing: 0.08%, under the 0.1% limit + joined, counts = join_event(a, b, "death", label_a="a", label_b="b") + assert joined.height == 1199 + assert counts["n_unmatched_a"] == 1 and counts["n_unmatched_b"] == 0 + + +def test_duplicate_keys_are_refused() -> None: + a, b = _pair() + with pytest.raises(RowSetMismatchError, match="duplicate row keys"): + join_event(pl.concat([a, a.head(1)]), b, "death", label_a="a", label_b="b") + + +def test_label_disagreement_is_refused() -> None: + a, b = _pair() + b = b.with_columns((1.0 - pl.col("y@8h")).alias("y@8h")) + with pytest.raises(RowSetMismatchError, match="different labels"): + score_pair(_joined(a, b), "hazard", 8.0, n_boot=10, seed=0) + + +def test_two_subject_fixture_shows_subject_clustering() -> None: + """Skipped single-class resamples prove subjects are the resampling unit. + + With one all-positive and one all-negative subject, a resample that + draws the same subject twice is single-class and must be skipped; + that only happens when subjects, not rows, are resampled. + """ + rng = np.random.default_rng(0) + y = np.array([1.0] * 20 + [0.0] * 20) + frame = pl.DataFrame( + { + "event": ["death"] * 40, + "subject_id": np.repeat([1, 2], 20), + "visit_id": np.repeat([1, 2], 20), + "time_hours": np.arange(40, dtype=float), + "y@8h": y, + "hazard@8h": np.where(y == 1, 0.6, 0.3) + rng.uniform(0, 0.2, 40), + } + ) + b = frame.with_columns( + (pl.col("hazard@8h") + 0.1 * pl.col("y@8h")).alias("hazard@8h") + ) + cell = score_pair(_joined(frame, b), "hazard", 8.0, n_boot=200, seed=0) + assert cell is not None + assert cell["n_subjects"] == 2 + delta = cell["delta_b_minus_a"] + assert delta["n_boot_used"] + delta["n_boot_skipped"] == 200 + # P(both draws are the same subject) = 1/2, so about half are skipped + assert 60 <= delta["n_boot_skipped"] <= 140 + + +def test_compare_reports_gbm_refits_without_a_bootstrap_and_rejects_missing_events() -> ( + None +): + a, b = _pair(shift=0.2) + out = compare( + a, + b, + label_a="a", + label_b="b", + scorer="hazard", + events=None, + horizons=None, + n_boot=50, + seed=0, + ) + assert list(out["cells"]) == ["death@8h"] + assert out["events"]["death"]["n_rows_joined"] == 240 + g = out["gbm"]["death@8h"] + assert g["auroc_a"] != g["auroc_b"] # two refits + assert g["delta_b_minus_a"] == pytest.approx(g["auroc_b"] - g["auroc_a"]) + assert "ci_low" not in g + table = markdown_table(out["cells"], out["gbm"], label_a="a", label_b="b") + assert "| death | 8h | 240 | 80 | 60 |" in table + with pytest.raises(RowSetMismatchError, match="not in both dumps"): + compare( + a, + b, + label_a="a", + label_b="b", + scorer="hazard", + events=["sepsis"], + horizons=None, + n_boot=10, + seed=0, + ) + + +def test_inference_block_reports_side_by_side_deltas() -> None: + inf_a = { + "set_top1_accuracy": 0.80, + "top1_accuracy": 0.37, + "cross_entropy": 3.5, + "top5_accuracy": 0.7, + "n_predictions": 100, + "n_set_predictions": 99, + } + inf_b = {**inf_a, "top1_accuracy": 0.39} + block = inference_block(inf_a, inf_b) + assert block is not None + assert block["delta_b_minus_a"]["top1_accuracy"] == pytest.approx(0.02) + assert block["delta_b_minus_a"]["n_predictions"] is None + assert inference_block(None, None) is None + + +def test_main_writes_json_and_refuses_on_mismatch( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + a, b = _pair(shift=0.3) + pa, pb = tmp_path / "a.parquet", tmp_path / "b.parquet" + a.with_columns(pl.lit(4).alias("landmark_protocol_version")).write_parquet(pa) + b.with_columns(pl.lit(4).alias("landmark_protocol_version")).write_parquet(pb) + inf = tmp_path / "inference_results.json" + inf.write_text( + json.dumps( + { + "task_metrics": { + "top1_accuracy": 0.37, + "set_top1_accuracy": 0.8, + "cross_entropy": 3.5, + } + } + ) + ) + out = tmp_path / "paired.json" + argv = [ + "compare_runs_paired", + "--dump-a", + str(pa), + "--dump-b", + str(pb), + "--label-a", + "bottleneck", + "--label-b", + "baseline", + "--inference-a", + str(inf), + "--inference-b", + str(inf), + "--horizons", + "8", + "--n-boot", + "20", + "--seed", + "0", + "--output-json", + str(out), + ] + monkeypatch.setattr(sys, "argv", argv) + crp.main() + payload = json.loads(out.read_text()) + assert payload["row_key"] == ["event", "subject_id", "visit_id", "time_hours"] + assert payload["summary"]["b_beats_a"] == ["death@8h"] + assert payload["landmark_protocol_version"] == {"a": 4, "b": 4} + assert payload["inference"]["delta_b_minus_a"]["top1_accuracy"] == 0.0 + assert payload["gbm"]["death@8h"]["auroc_a"] is not None + + b.head(200).write_parquet(pb) + monkeypatch.setattr(sys, "argv", argv) + with pytest.raises(SystemExit) as excinfo: + crp.main() + assert excinfo.value.code == 2 diff --git a/tests/scripts/test_make_hparams_table.py b/tests/scripts/test_make_hparams_table.py new file mode 100644 index 00000000..e0b0320b --- /dev/null +++ b/tests/scripts/test_make_hparams_table.py @@ -0,0 +1,184 @@ +"""The hyperparameter table renders the banked configs and the GBM's code settings.""" + +from __future__ import annotations + +import json +import sys +from pathlib import Path +from typing import Any + +import pytest + +from scripts import make_hparams_table +from scripts.make_hparams_table import NOT_EXPORTED, TBD, gbm_rows, render, run_name + + +REPO = Path(__file__).resolve().parents[2] +BANKED = { + "MIMIC-IV": REPO / "research_journal/figure_data/vm1/full_run_v10/config.json", + "eICU-CRD": REPO / "research_journal/figure_data/vm2/eicu_full_v10/config.json", +} + +# The values the two flagship configs carry (research_journal is gitignored, +# so a clean checkout falls back to this copy of the exported keys). +_FLAGSHIP: dict[str, Any] = { + "output_dir": "/home/amritkrishnan/runs/full_run_v10", + "model_kind": "bottleneck", + "backbone": "hybrid", + "max_context": 4096, + "hidden_size": 256, + "num_hidden_layers": 8, + "mamba_state_size": 128, + "mamba_headdim": 64, + "mamba_chunk_size": 256, + "attn_num_heads": 8, + "embedding_dim": 32, + "vocab_min_count": 5, + "vocab_max_size": 20000, + "vocab_backoff": "icd3", + "quantile_n_bins": 5, + "quantile_min_count": 100, + "num_lanes": 64, + "chunk_size": 512, + "learning_rate": 0.0003, + "weight_decay": 0.01, + "grad_clip_norm": 1.0, + "num_epochs": 2, + "concept_weight": 1.0, + "orthogonality_weight": 0.1, + "observability_weight": 0.1, + "task_weight": 1.0, + "checkpoint_every": 2000, + "time_weight": 1.0, + "event_hazard_weight": 1.0, + "randint_prob": 0.0, + "early_stopping_patience": 15, + "seed": 0, +} + + +def _config_paths(tmp_path: Path) -> dict[str, Path]: + """Use the real banked configs when present, else fixtures with their values.""" + out: dict[str, Path] = {} + for label, banked in BANKED.items(): + if banked.exists(): + out[label] = banked + continue + path = tmp_path / f"{label}.json" + cfg = dict(_FLAGSHIP) + if label == "eICU-CRD": + cfg["output_dir"] = "/home/amritkrishnan/runs/eicu_full_v10" + path.write_text(json.dumps(cfg)) + out[label] = path + return out + + +def _rows(text: str) -> dict[str, list[str]]: + rows: dict[str, list[str]] = {} + for raw in text.splitlines(): + if not raw.endswith("\\\\") or "&" not in raw: + continue + label, *cells = raw[: -len(" \\\\")].split(" & ") + rows[label] = cells + return rows + + +def test_model_and_training_rows_carry_the_banked_values(tmp_path: Path) -> None: + paths = _config_paths(tmp_path) + runs = [(label, json.loads(p.read_text())) for label, p in paths.items()] + rows = _rows(render([*runs, ("GEMINI", None)])) + + assert rows["Banked config"] == ["yes", "yes", NOT_EXPORTED] + assert rows["Backbone"][:2] == ["hybrid Mamba-2 + chunk attention"] * 2 + assert rows["Hidden size"] == ["256", "256", "--"] + assert rows["Layers"][0] == "8" + assert rows["Attention heads"][0] == "8" + assert rows["Mamba state size"][0] == "128" + assert rows["Mamba head dim"][0] == "64" + assert rows["Mamba chunk size"][0] == "256" + assert rows["Concept embedding dim"][0] == "32" + assert rows["Lanes $\\times$ chunk (tokens)"][0] == "64 $\\times$ 512" + assert rows["Max context (tokens)"][0] == "4{,}096" + assert rows["Optimizer"] == ["AdamW (constant LR)"] * 2 + ["--"] + assert rows["Learning rate"][0] == "$3\\times10^{-4}$" + assert rows["Weight decay"][0] == "0.01" + assert rows["Gradient clip (norm)"][0] == "1" + assert rows["Epochs"][0] == "2" + assert rows["Early-stopping patience (evals)"][0] == "15" + assert rows["Checkpoint every (steps)"][0] == "2{,}000" + assert rows["Seed"][0] == "0" + assert rows["RandInt probability"][0] == "0" + assert rows["Concept"][0] == "1" + assert rows["Orthogonality"][0] == "0.1" + assert rows["Observability"][0] == "0.1" + assert rows["Task (next token)"][0] == "1" + assert rows["Time to event"][0] == "1" + assert rows["Event hazard"][0] == "1" + assert rows["Vocabulary min count"][0] == "5" + assert rows["Vocabulary max size"][0] == "20{,}000" + assert rows["Vocabulary backoff"][0] == "icd3" + assert rows["Quantile bins per lab"][0] == "5" + assert rows["Parameters"] == [TBD, TBD, TBD] + + +def test_parameter_counts_key_by_run_name_or_label(tmp_path: Path) -> None: + paths = _config_paths(tmp_path) + runs = [(label, json.loads(p.read_text())) for label, p in paths.items()] + assert run_name(runs[0][1], "MIMIC-IV") == "full_run_v10" + assert run_name(None, "GEMINI") == "GEMINI" + rows = _rows( + render( + [*runs, ("GEMINI", None)], + {"full_run_v10": 12_345_678, "GEMINI": 9_000_000}, + ) + ) + assert rows["Parameters"] == ["12{,}345{,}678", TBD, "9{,}000{,}000"] + + +def test_gbm_block_reads_the_code_not_a_config() -> None: + rows = dict(gbm_rows()) + assert "HistGradientBoostingClassifier" in rows["Estimator"] + assert ( + rows["Search grid (LR, max leaves, min leaf)"] + == "(0.05, 31, 20); (0.05, 63, 100); (0.1, 15, 20); (0.1, 63, 100)" + ) + assert rows["Boosting rounds"].startswith("up to 400") + assert "subject-grouped" in rows["Validation"] + assert "200{,}000 rows" in rows["Validation"] + assert rows["Feature panel"] == "609 features, 110 counts" + + +def test_main_writes_the_table_with_one_column_per_database( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + paths = _config_paths(tmp_path) + params = tmp_path / "params.json" + params.write_text(json.dumps({"eicu_full_v10": 1_000_000})) + out = tmp_path / "tables" / "hparams.tex" + argv = ["make_hparams_table"] + for label, path in paths.items(): + argv += ["--run", label, str(path)] + argv += ["--missing", "GEMINI", "--params", str(params), "--output", str(out)] + monkeypatch.setattr(sys, "argv", argv) + make_hparams_table.main() + + text = out.read_text() + assert text.startswith("% GENERATED by scripts/make_hparams_table.py") + assert "\\begin{tabular}{@{}lrrr@{}}" in text + assert "Setting & MIMIC-IV & eICU-CRD & GEMINI \\\\" in text + assert "\\toprule" in text and "\\bottomrule" in text + assert "\\multicolumn{3}{l}{scikit-learn" in text + rows = _rows(text) + assert rows["Parameters"] == [TBD, "1{,}000{,}000", TBD] + # the paper rule: no em dashes anywhere in a generated table + assert chr(0x2014) not in text # em dash + + +def test_main_refuses_an_empty_column_set( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr( + sys, "argv", ["make_hparams_table", "--output", str(tmp_path / "x.tex")] + ) + with pytest.raises(SystemExit, match="at least one"): + make_hparams_table.main() diff --git a/tests/scripts/test_panel_coverage.py b/tests/scripts/test_panel_coverage.py new file mode 100644 index 00000000..0d1335c7 --- /dev/null +++ b/tests/scripts/test_panel_coverage.py @@ -0,0 +1,82 @@ +"""Tests for the signal-panel coverage report (scripts/panel_coverage.py).""" + +import json +import sys +from pathlib import Path + +import polars as pl +import pytest + +from odyssey.data.signal_panel import N_PANEL_SIGNALS, SIGNAL_PANEL +from scripts import panel_coverage +from scripts.panel_coverage import OUTPUT_KEYS, SOURCES, format_table + + +@pytest.mark.parametrize("source", SOURCES) +def test_every_source_resolves_a_subset_of_the_panel(source: str) -> None: + report = panel_coverage.panel_coverage(source) + assert set(report) == set(OUTPUT_KEYS) + assert report["source"] == source + assert report["n_panel_signals"] == N_PANEL_SIGNALS == len(SIGNAL_PANEL) + assert 0 < report["n_resolved"] <= report["n_panel_signals"] + assert report["n_resolved"] == len(report["resolved"]) + assert len(report["resolved"]) + len(report["unresolved"]) == N_PANEL_SIGNALS + assert not set(report["resolved"]) & set(report["unresolved"]) + assert set(report["prefixes"]) == set(report["resolved"]) + # Without an inventory the observed split is explicitly absent. + assert report["codes_inventory"] is False + assert report["resolved_observed"] is None + + +def test_gemini_lacks_the_noninvasive_blood_pressure_panel() -> None: + """The review's guess, checked against the in-repo table.""" + report = panel_coverage.panel_coverage("gemini") + assert "sbp_noninvasive" in report["unresolved"] + assert "map_noninvasive" in report["unresolved"] + assert "creatinine" in report["resolved"] + assert "lactate" in report["resolved"] + + +def test_observed_codes_split_the_resolved_signals() -> None: + """A resolved prefix that no charted code matches is reported as unobserved.""" + heart_rate = panel_coverage.panel_coverage("gemini")["prefixes"]["heart_rate"][0] + report = panel_coverage.panel_coverage( + "gemini", observed_codes=[heart_rate + "bpm::3", "SOMETHING_ELSE"] + ) + assert report["codes_inventory"] is True + assert report["resolved_observed"] == ["heart_rate"] + assert "creatinine" in report["resolved_unobserved"] + assert set(report["resolved_observed"]) | set(report["resolved_unobserved"]) == set( + report["resolved"] + ) + assert "no codes" in format_table(report) + + +def test_main_writes_json_from_a_codes_parquet( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str] +) -> None: + codes = tmp_path / "codes.parquet" + pl.DataFrame( + {"code": ["LAB//3020564//umol/L", "VITALS//3027018//"], "count": [5, 7]} + ).write_parquet(codes) + out = tmp_path / "panel_coverage.json" + monkeypatch.setattr( + sys, + "argv", + [ + "panel_coverage", + "--source", + "gemini", + "--codes-parquet", + str(codes), + "--output-json", + str(out), + ], + ) + panel_coverage.main() + report = json.loads(out.read_text()) + assert set(report) == set(OUTPUT_KEYS) + assert sorted(report["resolved_observed"]) == ["creatinine", "heart_rate"] + # Only names leave: the inventory's counts never reach the output. + assert "count" not in json.dumps(report) + assert "resolved: " in capsys.readouterr().out diff --git a/tests/scripts/test_probe_channel.py b/tests/scripts/test_probe_channel.py new file mode 100644 index 00000000..81ccec2f --- /dev/null +++ b/tests/scripts/test_probe_channel.py @@ -0,0 +1,440 @@ +"""Tests for scripts/probe_channel.py on a tiny synthetic run (CPU).""" + +import json +from datetime import datetime, timedelta +from pathlib import Path +from typing import Any + +import polars as pl +import pytest +import torch +import torch.nn.functional as F # noqa: N812 + +from odyssey.data.concepts import concepts_for_source +from odyssey.data.streaming import PackedLaneSampler +from odyssey.data.value_binning import QuantileBinner +from odyssey.data.vocabulary import Vocabulary +from odyssey.models.backbones.tiny_gru import TinyGRUBackbone +from odyssey.models.sequence_model import ( + BaselineSequenceModel, + ConceptBottleneckSequenceModel, +) +from odyssey.training.data import iter_patient_sequences +from odyssey.training.train import TrainingConfig +from scripts import probe_channel + + +T0 = datetime(2024, 1, 1) +NUM_CONCEPTS = 3 +CODES = [f"LAB//{i}//" for i in range(10)] +SUBJECTS = range(1, 13) +EVENTS_PER_SUBJECT = 24 +SCHEMA = { + "subject_id": pl.Int64, + "code": pl.Utf8, + "time": pl.Datetime, + "numeric_value": pl.Float32, + "hadm_id": pl.Int64, +} + + +def _events(subjects: range) -> pl.DataFrame: + # The next code is a function of the current one, so a trained model + # forecasts it exactly and the full bottleneck carries that forecast. + rows = [ + ( + sid, + CODES[(sid * 3 + i) % len(CODES)], + T0 + timedelta(hours=i), + None, + 100 + sid, + ) + for sid in subjects + for i in range(EVENTS_PER_SUBJECT) + ] + return pl.DataFrame(rows, schema=SCHEMA, orient="row") + + +def _vocab() -> Vocabulary: + tokens = {"[PAD]": 0, "[UNK]": 1} + tokens.update({c: i + 2 for i, c in enumerate(CODES)}) + return Vocabulary(tokens) + + +def _trained_model(vocab: Vocabulary) -> ConceptBottleneckSequenceModel: + torch.manual_seed(0) + model = ConceptBottleneckSequenceModel( + backbone=TinyGRUBackbone( + vocab_size=len(vocab), hidden_size=16, num_layers=1, padding_idx=0 + ), + vocab_size=len(vocab), + num_concepts=NUM_CONCEPTS, + embedding_dim=4, + padding_idx=0, + ) + opt = torch.optim.Adam(model.parameters(), lr=1e-2) + model.train() + for _ in range(6): + state = None + sampler = PackedLaneSampler( + iter_patient_sequences(_events(SUBJECTS), vocab), + num_lanes=2, + chunk_size=16, + reset_prob=0.0, + ) + for chunk in sampler: + opt.zero_grad() + logits, _, state = model( + chunk.batch, state=state, reset_mask=chunk.reset_mask + ) + real = chunk.real_mask + loss = F.cross_entropy(logits[real], chunk.targets[real]) + loss.backward() + opt.step() + state = state.detach() if hasattr(state, "detach") else None + model.eval() + return model + + +def _write_shards(shard_dir: Path) -> None: + shard_dir.mkdir() + _events(range(1, 7)).write_parquet(shard_dir / "0.parquet") + _events(range(7, 13)).write_parquet(shard_dir / "1.parquet") + + +def _config(tmp_path: Path, **overrides: Any) -> TrainingConfig: + return TrainingConfig( + train_shard_dir=str(tmp_path / "train"), + tuning_shard_dir=str(tmp_path / "tuning"), + output_dir=str(tmp_path / "run"), + source="mimic_iv", + task_set="v1", + concept_supervision="visit", + **overrides, + ) + + +def _patch_run( + monkeypatch: pytest.MonkeyPatch, model: torch.nn.Module, config: TrainingConfig +) -> None: + vocab = _vocab() + binner = QuantileBinner(boundaries={}, n_bins=3) + monkeypatch.setattr( + probe_channel, "load_run", lambda *a, **k: (model, vocab, binner, config) + ) + monkeypatch.setattr( + probe_channel, + "resolve_concepts_for_run", + lambda *a, **k: concepts_for_source("mimic_iv", task_set="v1")[:NUM_CONCEPTS], + ) + + +@pytest.fixture(scope="module") +def payload(tmp_path_factory: pytest.TempPathFactory) -> dict[str, Any]: + tmp_path = tmp_path_factory.mktemp("probe") + _write_shards(tmp_path / "held_out") + with pytest.MonkeyPatch.context() as mp: + _patch_run(mp, _trained_model(_vocab()), _config(tmp_path)) + out = tmp_path / "run" / "channel_probes.json" + probe_channel.main( + [ + "--run-dir", + str(tmp_path / "run"), + "--held-out-shard-dir", + str(tmp_path / "held_out"), + "--output-json", + str(out), + "--max-positions", + "0", + "--num-lanes", + "2", + "--chunk-size", + "16", + "--epochs", + "40", + "--patience", + "5", + "--n-boot", + "50", + "--ctl-epochs", + "3", + "--train-frac", + "0.6", + "--tune-frac", + "0.15", + ] + ) + loaded: dict[str, Any] = json.loads(out.read_text()) + return loaded + + +def test_split_by_subject_has_no_subject_overlap() -> None: + gen = torch.Generator().manual_seed(1) + subject_ids = torch.randint(0, 40, (1000,), generator=gen) + split = probe_channel.split_by_subject(subject_ids, seed=3) + parts = { + name: set(subject_ids[getattr(split, name)].tolist()) + for name in ("train", "tune", "test") + } + assert parts["train"] & parts["test"] == set() + assert parts["train"] & parts["tune"] == set() + assert parts["tune"] & parts["test"] == set() + assert len(parts["train"] | parts["tune"] | parts["test"]) == 40 + total = split.train.numel() + split.tune.numel() + split.test.numel() + assert total == 1000 + + +def test_full_bottleneck_readout_is_at_least_as_accurate_as_k_only( + payload: dict[str, Any], +) -> None: + readouts = payload["readouts"] + assert readouts["h_bar"]["top1_accuracy"] >= readouts["k_only"]["top1_accuracy"] + # The trained model forecasts the deterministic cycle, and the full + # bottleneck is exactly what its head reads. + assert payload["model"]["top1_accuracy"] > payload["majority_class_accuracy"] + assert readouts["h_bar"]["retained"] > 0.5 + + +def test_payload_schema(payload: dict[str, Any]) -> None: + for key in ( + "run_dir", + "checkpoint", + "concept_names", + "protocol", + "n_positions", + "n_subjects", + "majority_class_accuracy", + "model", + "readouts", + "ctl", + "hazards", + "notes", + ): + assert key in payload, key + assert set(payload["readouts"]) == set(probe_channel.READOUT_NAMES) + for score in (payload["model"], *payload["readouts"].values()): + for key in ( + "n_positions", + "n_subjects", + "top1_accuracy", + "ci95", + "retained", + "retained_ci95", + "completeness_score", + ): + assert key in score, key + lo, hi = score["ci95"] + assert lo <= score["top1_accuracy"] <= hi + assert payload["model"]["retained"] == 1.0 + assert ( + payload["n_subjects"]["train"] + + payload["n_subjects"]["tune"] + + payload["n_subjects"]["test"] + == payload["n_subjects"]["bank"] + ) + assert payload["n_positions"]["bank"] == len(SUBJECTS) * (EVENTS_PER_SUBJECT - 1) + assert payload["n_shards"] == 2 + assert payload["hazards"] is None + assert payload["ctl"] is not None + assert "probs_only" in payload["ctl"] + assert len(payload["concept_names"]) == NUM_CONCEPTS + + +def test_refuses_a_baseline_run( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + model = BaselineSequenceModel( + TinyGRUBackbone(vocab_size=12, hidden_size=4), vocab_size=12 + ) + _patch_run(monkeypatch, model, _config(tmp_path, model_kind="baseline")) + with pytest.raises(ValueError, match="needs a concept bottleneck"): + probe_channel.run_channel_probes(tmp_path / "run", tmp_path / "held_out") + + +def test_refuses_a_decomposed_bottleneck( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + model = ConceptBottleneckSequenceModel( + TinyGRUBackbone(vocab_size=12, hidden_size=4), + vocab_size=12, + num_concepts=NUM_CONCEPTS, + embedding_dim=4, + bottleneck_kind="decomposed", + ) + _patch_run(monkeypatch, model, _config(tmp_path, bottleneck_kind="decomposed")) + with pytest.raises(ValueError, match="mixture bottleneck"): + probe_channel.run_channel_probes(tmp_path / "run", tmp_path / "held_out") + + +# --------------------------------------------------------------------------- +# Crash recovery: saved bank, staged JSON writes, append-only ctl fill +# --------------------------------------------------------------------------- + + +COMMON_ARGS = [ + "--max-positions", + "0", + "--num-lanes", + "2", + "--chunk-size", + "16", + "--epochs", + "10", + "--n-boot", + "20", + "--ctl-epochs", + "2", + "--train-frac", + "0.6", + "--tune-frac", + "0.15", +] + + +def _options(**overrides: Any) -> probe_channel.ProbeOptions: + base: dict[str, Any] = { + "max_positions": None, + "num_lanes": 2, + "chunk_size": 16, + "epochs": 10, + "n_boot": 20, + "ctl_epochs": 2, + "train_frac": 0.6, + "tune_frac": 0.15, + "skip_ctl": True, + } + base.update(overrides) + return probe_channel.ProbeOptions(**base) + + +def _readout_accuracies(payload: dict[str, Any]) -> dict[str, float]: + out = {name: s["top1_accuracy"] for name, s in payload["readouts"].items()} + out["model_head"] = payload["model"]["top1_accuracy"] + return out + + +def test_from_bank_round_trip_reproduces_the_readouts( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + _write_shards(tmp_path / "held_out") + _patch_run(monkeypatch, _trained_model(_vocab()), _config(tmp_path)) + bank_path = tmp_path / "saved.bank.pt" + streamed = probe_channel.run_channel_probes( + tmp_path / "run", + tmp_path / "held_out", + options=_options(bank_path=bank_path), + device="cpu", + ) + assert bank_path.exists() + + bank, meta = probe_channel.load_bank(bank_path) + assert len(bank) == streamed["n_positions"]["bank"] + assert meta.vocab_size == streamed["vocab_size"] + assert bank.concept_names == tuple(streamed["concept_names"]) + assert bank.concept_probs.dtype == torch.float16 + assert probe_channel.bank_nbytes(bank) > 0 + + # Reloading must not touch the checkpoint at all. + def _no_model(*a: Any, **k: Any) -> Any: + raise AssertionError("load_run must not be called with --from-bank") + + monkeypatch.setattr(probe_channel, "load_run", _no_model) + reloaded = probe_channel.run_channel_probes( + tmp_path / "run", + tmp_path / "held_out", + options=_options(from_bank=bank_path), + device="cpu", + ) + assert _readout_accuracies(reloaded) == _readout_accuracies(streamed) + assert reloaded["n_positions"] == streamed["n_positions"] + assert reloaded["n_subjects"] == streamed["n_subjects"] + assert reloaded["notes"] == streamed["notes"] + assert reloaded["ctl"] is None + + +def test_readouts_are_written_before_ctl_and_ctl_is_filled_in_place( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + _write_shards(tmp_path / "held_out") + _patch_run(monkeypatch, _trained_model(_vocab()), _config(tmp_path)) + out = tmp_path / "run" / "channel_probes.json" + args = [ + "--run-dir", + str(tmp_path / "run"), + "--held-out-shard-dir", + str(tmp_path / "held_out"), + "--output-json", + str(out), + *COMMON_ARGS, + ] + + # Stage 1: readouts only (as if the run died before/inside the CTL probes). + probe_channel.main([*args, "--skip-ctl"]) + first = json.loads(out.read_text()) + assert first["ctl"] is None + bank_path = probe_channel.default_bank_path(out) + assert bank_path.exists() + + # A rerun that would stream again refuses: the bank is already on disk. + with pytest.raises(FileExistsError, match="--from-bank"): + probe_channel.main(args) + # ... and --skip-ctl on an unfilled JSON has nothing to do. + with pytest.raises(ValueError, match="nothing to"): + probe_channel.main([*args, "--skip-ctl", "--from-bank", str(bank_path)]) + + # Stage 2: JSON exists without ctl, bank saved; fill the ctl block only. + monkeypatch.setattr( + probe_channel, + "load_run", + lambda *a, **k: (_ for _ in ()).throw(AssertionError("no model load")), + ) + probe_channel.main([*args, "--from-bank", str(bank_path)]) + second = json.loads(out.read_text()) + assert second["ctl"] is not None + assert "probs_only" in second["ctl"] + assert second["readouts"] == first["readouts"] + assert second["model"] == first["model"] + assert second["n_positions"] == first["n_positions"] + + # Stage 3: complete JSON is append-only. + with pytest.raises(FileExistsError, match="--overwrite"): + probe_channel.main([*args, "--from-bank", str(bank_path)]) + # --overwrite refits everything from the saved bank. + probe_channel.main([*args, "--from-bank", str(bank_path), "--overwrite"]) + third = json.loads(out.read_text()) + assert third["ctl"] is not None + assert _readout_accuracies(third) == _readout_accuracies(first) + + +def test_fill_refuses_a_payload_from_a_different_split( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + _write_shards(tmp_path / "held_out") + _patch_run(monkeypatch, _trained_model(_vocab()), _config(tmp_path)) + out = tmp_path / "channel_probes.json" + bank_path = tmp_path / "bank.pt" + streamed = probe_channel.run_channel_probes( + tmp_path / "run", + tmp_path / "held_out", + options=_options(bank_path=bank_path), + device="cpu", + output_json=out, + ) + stale = json.loads(out.read_text()) + stale["protocol"]["seed"] = streamed["protocol"]["seed"] + 1 + with pytest.raises(ValueError, match="does not match"): + probe_channel.run_channel_probes( + tmp_path / "run", + tmp_path / "held_out", + options=_options(from_bank=bank_path, skip_ctl=False), + device="cpu", + output_json=out, + existing=stale, + ) + + +def test_load_bank_rejects_a_foreign_file(tmp_path: Path) -> None: + path = tmp_path / "not_a_bank.pt" + torch.save({"format": 99}, path) + with pytest.raises(ValueError, match="not a probe_channel bank"): + probe_channel.load_bank(path)