Skip to content

Latest commit

 

History

1 Commit

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

code/: Interpretable inputs and ECG foundation models under signal degradation

Reproduces the study "Transparent and Robust: Interpretable Inputs Improve Foundation Models in Mission-Critical Decision Pipelines" (XAI-EADM workshop, IEEE CBI/EDOC 2026, Springer LNBIP). The question: a frozen ECG foundation model sees only the waveform, so it cannot see patient age, one of the strongest predictors of one-year mortality. Concatenate that single column to the embedding and every one of six open encoders improves, in all 48 encoder-condition-cohort cells, by up to 0.160 AUROC, with the largest gains under the corruptions that cost the encoders most. A gradient-boosted model over 165 named ECG features plus age is equivalent to the best unaided encoder on clean data (TOST at a 0.02 AUROC margin). Scored at equal input, the encoders lead that named-feature model.

Read the age-control section below before using any number from results/robustness_*.json on its own: those margins are measured at unequal input, and the age control is the equal-input comparison.

Layout

This deposit is fully self-contained - one command rebuilds every number from the raw, credentialed inputs, including the age control.

  • degrade/ - the six real-world signal corruptions, applied in memory at the raw-waveform stage (unit-tested).
  • eval/ - build_labels.py (cohort + 1-yr mortality labels + age + split), extract_named_deg.py / extract_fm_deg.py (per-condition named-feature banks and 6 FM embeddings on the corrupted signal), evaluate.py (the robustness analysis), transport.py (cross-cohort transport under degradation), certificate.py (the per-decision certificate + novelty score), age_control.py (the age control: every arm at equal input, plus per-patient prediction dumps), paired_stats.py (the paired statistics over those dumps: DeLong, TOST, AUPRC, calibration).
  • pipeline/ - the feature/embedding pipeline: core/ (config + MIMIC/CODE-15 IO + the ECGFounder architecture), embed/ (base-cache + raw-cache builders, the FM zoo, and the vendored ECG-JEPA / xECG / ST-MEM encoders under embed/vendor/), features/ (rich_features.py = the 165 named features; band_feat.py), cohort/ (mimic_cohort.py = the MIMIC cohort + mortality table).
  • hpc/ - the SLURM DAG (run_all.sh), the chunked array/eval wrappers, the age-control array (run_age_control.slurm), the node-local scratch helper (scratch.sh), and the env builders (env_build_*.sh) + env_stage.sh.
  • figures/ - make_fig1.py (the mechanism plot, from committed results).

One command (cluster), from raw

# from the UTwente login node: submits the whole DAG and exits - never computes on the login node
bash code/hpc/run_all.sh
squeue -u $USER                       # watch; logs in code/hpc/logs/
# results land in results/robustness_{mimic,code15}.json + results/transport.json + results/age_control/

To re-run the age control alone against caches the DAG has already built, submit it from code/ (the runners resolve their paths from the submit directory):

cd code
sbatch --array=0-3  hpc/run_age_control.slurm     # seed 0
sbatch --array=4-19 hpc/run_age_control.slurm     # the training-resample seeds
sbatch hpc/run_cpu.slurm eval/paired_stats.py --dir ../results/age_control --out ../results/age_control/stats.json

run_all.sh assumes no pre-built intermediate. Stage 0 rebuilds the base ECGFounder caches (they define the MIMIC/CODE-15 subsets and carry labels + age) and the float16 raw-trace caches raw_{mimic,code15}.npz straight from the credentialed inputs. It fails loudly at the top if a credentialed dataset or a packed env is missing. The DAG, per corruption condition (clean + 16 degraded, x2 cohorts):

Stage 0  base ECGFounder caches (GPU) -> raw_<ds>.npz (CPU array + merge)   [pipeline/embed]
  -> labels_<ds>.npz (CPU, build_labels.py: y + split + age)
  -> raw_<ds>_test.npz (CPU, make_raw_test.py: the test-split slice)
  -> per CHUNK of conditions: stage the raw cache to /local/<jobid> ONCE, then per condition
       -> corrupt in memory (degrade.py)
       -> named features   (CPU, rich())                    -> $DEG/<cond>/rich_<ds>.npz
       -> 6 FM embeddings  (GPU, extract_fm_deg.py)         -> $DEG/<cond>/<ds>_<fm>_emb.npz
  -> evaluate.py     (CPU, per cohort)    -> results/robustness_<ds>.json
  -> transport.py    (CPU, needs both)    -> results/transport.json
  -> age_control.py  (CPU array, both)    -> results/age_control/*.json + *_preds.npz
  -> paired_stats.py (CPU, after array)   -> results/age_control/stats.json

The last three read both cohorts' caches, so they fire after both waves' arrays, submitted by the deferred CODE-15 helper. The age-control array is throttled to 4 concurrent tasks to stay inside the 8-running cap while the transport job and the CODE-15 tail are still draining.

Node-local scratch. Every job works on /local/$SLURM_JOB_ID (the compute nodes' NVMe) and copies results back to home as its last step; hpc/scratch.sh provides $SCRATCH, stage_in/stage_out_dir, a TMPDIR redirect, and the cleanup trap. The trap is required: this cluster's SLURM epilog is a no-op, so nothing creates or removes the dir for you. The packed env is the one thing kept OUT of per-job scratch, staying a node-cached tarball under /local/$USER (env_stage.sh), since re-extracting it per job would add tens of GB of NFS reads.

Chunked arrays. An array task is a CHUNK of conditions, not one condition (hpc/_lib.sh): task 0 = clean (the only condition needing the full cache, since it also embeds the train split), tasks 1-4 = 4 degraded conditions each. Each task stages its raw cache once and loops. Degraded conditions are test-only, so they read the pre-sliced raw_<ds>_test.npz (~40% of the bytes). One task per condition meant 68 tasks per cohort each re-reading a 1.7 GB / 7.8 GB cache off NFS (~646 GB per run); this is ~83 GB, at ~44 upfront submissions. eval/_rawsel.py accepts either the full or the pre-sliced cache and raises on any id misalignment, and the slice is order-preserving, so the numbers are unchanged either way.

The age control at equal input (eval/age_control.py + eval/paired_stats.py)

The named-feature arm receives patient age as an exact, corruption-immune input, while the six frozen encoders have to recover it from a degraded waveform. Comparing them as-is compares unequal inputs, so age_control.py fits every arm at equal input:

Arm Inputs
named 165 named features + age
named_noage the 165 named features alone
age_only age alone, one input
<fm> the frozen embedding
<fm>_age the frozen embedding + age

Heads are trained on clean data and applied to corrupted data, matching the deployment setting. Every arm scores the same held-out patients, so age_control.py dumps the per-patient probability vector for each (arm, condition) pair into results/age_control/*_preds.npz alongside the AUROC summary in *.json. Every downstream statistic then runs on those dumps without a refit.

paired_stats.py reads the dumps and reports, per condition: the paired AUROC difference with a DeLong interval and two-sided p-value; a TOST equivalence verdict against the 0.02 margin fixed in advance, which is what a parity claim needs and what a marginal interval does not answer; AUPRC with a paired patient-level bootstrap, since prevalence is 4.3 and 12.9 percent; and calibration under corruption as ECE, Brier, and recalibration slope and intercept.

Seeds are training resamples, not test resamples. --seed 0 fits on the full training split; seeds 1 and up fit on a bootstrap resample of it, so the spread reflects model variability. The test split is fixed across seeds by necessity: the corrupted waveforms were encoded for test rows only, so re-randomising which patients are held out would need the whole GPU encode again. Report training-resample stability, and do not describe the seed spread as a test-set interval.

Array task index is seed * 4 + config, with configs 0-3 covering within-MIMIC, within-CODE-15, and both transport directions. Indices 0-3 are therefore the entire seed-0 wave.

Shipped outputs. results/age_control/*.json (the 20 per-run AUROC summaries) and stats.json (every paired statistic) are in this repo, so each number can be checked without cluster access or credentialed data. The per-patient dumps they derive from are about 250 MB and are not shipped; run_age_control.slurm regenerates them.

Corruptions (degrade/degrade.py)

Corruption Levels What it simulates
lead_dropout 1 / 2 / 3 leads dead electrode / missing lead
emg_noise SNR 20 / 10 / 5 dB muscle interference
baseline_wander 0.5 / 1 / 2 x RMS respiratory / electrode drift
powerline 0.10 / 0.25 / 0.50 x RMS 50 Hz mains
bandwidth_loss 250 / 125 / 62 Hz low-rate acquisition
combined one preset a realistic "bad recording"

Each is deterministic given (dataset, corruption, level), so the named-feature job and the FM job corrupt identically.

Metrics (eval/evaluate.py)

evaluate.py computes a wider set of diagnostics than the study reports on. The headline comparison is the age control above; the rest is exploratory output of the pipeline.

  • Accuracy-robustness: AUROC vs. severity for the named model and each FM.
  • Certificate-robustness: size, coverage, and reason-set stability of the 7-feature certificate under degradation.
  • Novelty score: a named-feature novelty statistic scored on its ability to detect a degraded recording and to name the anomalous feature.
  • Cross-cohort transport under degradation (eval/transport.py -> results/transport.json): train on the source cohort's clean data, test on the target cohort's test split, clean and under every degradation, both directions. Reuses the arrays' cached features (no re-extraction). The job runs after both cohorts' waves finish.

The transport and robustness margins between the named-feature arm and the unaided encoders are measured at unequal input, since only the named arm gets age there. The equal-input comparison in age_control.py supersedes them.

Verified locally (no cluster data needed)

pip install -r requirements.txt
python degrade/test_degrade.py     # 6 corruptions do exactly what they claim (SNR, lead-drop, spectra)
python eval/certificate.py         # sufficient reason + self-flag self-test
python eval/evaluate.py --synth    # full analysis plumbing on fabricated caches (numbers are a smoke check, NOT results)

eval/paired_stats.py also runs locally, on prediction dumps pulled back from the cluster: it needs only results/age_control/*_preds.npz, not the caches or a GPU.

Prerequisites (cluster, shipping-impossible - checked at the top of run_all.sh)

Everything buildable is built by the DAG (Stage 0 rebuilds raw_{mimic,code15}.npz + the base caches from raw). What cannot be shipped and must exist first:

  • Packed envs in ~/envs/: gbdp_cpu_env, gbdp_gpu_env, gbdp_xecg_env, gbdp_ecgfm_env, each built ONCE on the login node via the matching hpc/env_build_{cpu,gpu,xecg,ecgfm}.sh (internet + node-local disk; never on NFS or a compute node). env_stage.sh extracts them to /local per node.
  • Credentialed data: ~/data/archive/mimic-iv-ecg-diagnostic-electrocardiogram-matched-subset-1.0.zip (PhysioNet), ~/data/code15/ (CODE-15 HDF5 + exams.csv), ~/data/mimic-iv/hosp/ (the patients table, for the age join).
  • FM checkpoints (compute nodes are offline, so pre-stage once on the login node): ECG-JEPA, ECG-FM, and ST-MEM weights in ~/models/glassbox-decision-pipelines/; ECGFounder, HuBERT-ECG, and xECG load from the HuggingFace cache (pre-download them there). Set $ECG_EMB_DIR to override the default cache dir (~/data/glassbox-decision-pipelines/emb).

Confirm-on-cluster note: eval/build_labels.py reads the CODE-15 mortality label field from the base cache (code15_ecgfounder_500hz.npz); it auto-detects y/death/label/mortality - verify the actual key on first run.

About

Code for "Transparent and Robust" (XAI-EADM 2026): one age column added to the embedding of six frozen ECG foundation models raises one-year mortality AUROC in all 48 encoder, corruption and cohort settings.

Topics

Resources

Stars

1 star

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages