Official code for Interventionally-guided representation learning for robust and interpretable AI models in cancer medicine.
PiCo leverages constrained representation learning to improve robustness and interpretability in high-dimensional machine learning models. This work shows that robust and interpretable predictions for a downstream task can be made by generating representations of some high-dimensional input data which are linked to auxiliary targets.
Many techniques for prediction from high-dimensional data either implicitly (through deep learning) or explicitly (through PCA, factor analysis etc.) use dimensionality reduction prior to fitting prediction models. This results in an unbiased compression of information in the input, preserving the largest sources of variance in the input. In many cases in biology, we have some prior on what information should be preserved in our low-dimensional representation for good prediction on a downstream task. It is popular to use priors derived from public databases for this. We take a different approach, using a data-driven method to preserve information related to auxiliary targets in representations, which is captured in an interpretable way for use in downstream tasks.
We specifically study the use of PiCo in a cancer biology setting, using gene expression data as the input and CRISPR knockout effect data as the auxiliary data. Then, we use representations generated using this data for downstream tasks such as drug response prediction in cancer cell lines and treatment response in patients.
All dependencies are pinned in requirements.txt:
matplotlib==3.10.8
numpy==2.4.2
optuna==3.2.0
pandas==1.4.1
pyreadr==0.4.7
scikit_learn==1.8.0
scipy==1.17.0
seaborn==0.13.2
statsmodels==0.13.2
torch==2.0.1+cu118
tqdm==4.63.0
umap_learn==0.5.6
| Component | Version |
|---|---|
| Python | 3.9.12 |
| PyTorch | 2.0.1 |
| CUDA | 11.8 |
| cuDNN | 8.9 |
| OS (full training) | Linux RHEL 8.10 (kernel 4.18), Cambridge CSD3 Wilkes3 cluster |
| OS (demo notebooks) | Linux (Google Colab, Ubuntu 22.04); also tested on macOS 14 with PyTorch CPU-only |
PyTorch 2.0 or later is required because training uses torch.compile (a PyTorch feature that just-in-time compiles models for faster training).
- For the Colab / local demo notebooks (
demo/): CPU-only is sufficient; any machine with ≥ 8 GB RAM. No GPU required. - For reproducing full-paper training runs (hyperparameter optimisation, 10 seeds per setting): an NVIDIA GPU with ≥ 16 GB memory and CUDA 11.8 support is required in practice. All paper results were generated on NVIDIA A100 (80 GB) GPUs.
Typical install time on a standard desktop (including downloading PyTorch CUDA wheels): ~5–10 minutes.
If you only intend to run the demo notebooks on Colab, skip this section — the notebook's first cell installs everything needed.
Clone the repository:
git clone https://github.com/domkirkham/pico.git
cd pico
Create a conda environment:
conda create --name pico python=3.9.12 -y
conda activate pico
Install dependencies:
# Install torch (pick one)
# GPU (CUDA 11.8) — required for full training:
pip install torch==2.0.1 torchvision==0.15.2 torchaudio==2.0.2 --index-url https://download.pytorch.org/whl/cu118
# CPU-only — sufficient for the demo notebooks:
pip install torch==2.0.1 torchvision==0.15.2 torchaudio==2.0.2 --index-url https://download.pytorch.org/whl/cpu
# Install the rest
pip install -r requirements.txt
Two self-contained Jupyter notebooks reproduce the headline figures of the paper end-to-end on CPU. The pretrained iCoVAE / VAE encoders and a small preprocessed slice of DepMap / GDSC / TransNEO are included directly in this repository under demo/assets/ (~420 MB total), so no external downloads are required. The notebooks call the paper's exact plotting code (extracted into src/utils/plot_helpers.py.
| Notebook | Open | Reproduces | Expected runtime (CPU) |
|---|---|---|---|
| demo/demo_ccl.ipynb | Fig. 3 (drug-response performance, both the aggregated boxplot and the per-drug pointplot) and Fig. 4 (permutation feature importance) for four headline drugs: AZD6738, trametinib, oxaliplatin, 5-fluorouracil | ~5 min | |
| demo/demo_transneo.ipynb | Fig. 5c (RCB regression two-panel CV + ARTemis+PBCP pointplot across all six paper feature sets), Fig. 5d (pCR classification, same shape), and the Clinical+z+RNA permutation feature importance panels | ~3 min |
Both notebooks focus on modelling: they forward-pass each iCoVAE / VAE encoder live on the bundled gene-expression matrices, refit the Stage-2 ElasticNet / LogisticRegression probe with the saved per-seed hyperparameters, score on the held-out cancer types (CCL) or the ARTemis+PBCP cohort (TransNEO), and compute permutation feature importance from the fitted regressor coefficients. Per-seed test metrics are identical to the cached values from the original training pipeline — no precomputed predictions are loaded from disk.
Click either Colab badge at the top of this README. The notebook's first cell clones the repo and pins a torch version that's compatible with the current Colab runtime; subsequent cells load the included pretrained encoders + preprocessed data, run the live encoding + Stage-2 fit, then plot via src/utils/plot_helpers.py.
After completing the local install above:
jupyter lab demo/demo_ccl.ipynb
# or
jupyter lab demo/demo_transneo.ipynb
Each notebook prints aggregated metric tables and saves figures (PNG + SVG, dpi 600) into demo/outputs/:
-
demo_ccl.ipynb— Spearman ρ of predicted vs observed log(IC50) on held-out cancer types, by drug × feature extractor. Expected numbers across the 10 included seeds (mean ± s.d.; the figures themselves reproduce slices of Fig. 3 of the paper):Drug PiCo VAE AZD6738 (ceralasertib) 0.45 ± 0.03 0.31 ± 0.03 Trametinib 0.37 ± 0.04 0.33 ± 0.04 Oxaliplatin 0.55 ± 0.001 0.46 ± 0.03 5-Fluorouracil 0.32 ± 0.01 0.24 ± 0.05 Plus a permutation-feature-importance bar plot per drug (Fig. 4b–e). The first 16 latent dims are labelled by their constraint gene (e.g. z_RAD17, z_HUS1 for AZD6738; z_MDM4, z_TTF2 for oxaliplatin); the remaining unconstrained dims appear as z_16, z_17, …
-
demo_transneo.ipynb— RCB Spearman correlation and pCR AUROC on the ARTemis+PBCP external validation cohort, across all six paper feature sets. Expected numbers (matching Fig. 5c,d):Metric PiCo Clinical+z+RNA VAE Clinical+z+RNA Clinical+RNA (no z) Clinical only RCB Spearman ρ ~0.76 ~0.74 ~0.78 ~0.59 pCR AUROC ~0.91 ~0.89 ~0.89 ~0.74 Plus permutation-feature-importance panels for the Clinical+z+RNA model. The pCR panel is dominated by z_ERBB2, PGR expression, and age; the RCB panel by ESR1 / PGR expression and the constrained z dimensions tied to therapy-relevant genes (taxane score, z_PSMC1, z_FANCF).
The demo/assets/ folder contains the exact files the paper notebooks read from data/outputs/, restricted to the demo's 4 drugs × 10 seeds × {iCoVAE, VAE} × 6 feature sets (TransNEO). Filenames match the training-script output convention so that re-running the training pipeline from src/scripts/ writes new files into the same directories — the included ones are simply a subset of what a full re-run produces. The whole demo/assets/ folder can be regenerated from scratch via scripts/build_demo_bundle.py.
The PiCo framework requires three data objects:
x: Input data e.g. gene expressions: Auxiliary data e.g. CRISPR gene effecty: Target data e.g. drug response
To use the framework with new data, add a function to src/utils/data_utils.py which loads your data and returns x, s, y, and a list test_samples. These should be of type pd.DataFrame, with index set as sample identifiers shared across x, s, and y. Examples: process_depmap_gdsc and process_depmap_gdsc_transneo. If test_samples is empty, random samples are held out.
Then add a line to process_data in the same file corresponding to your new dataset.
Fit the iCoVAE (stage 1):
python src/scripts/icovae_hopt.py \
-dataset <your_dataset> \
-target <drug_or_outcome> \
-constraints GENE1 GENE2 ... \
--experiment <experiment_name> \
--cuda
(For automatic constraint selection rather than a fixed list, see get_constraints and the example cluster-submission script at docs/examples/slurm/icovae_hopt.sh.)
Fit a PiCo prediction head (stage 2):
python src/scripts/pico_sk_hopt.py \
-dataset <your_dataset> \
-target <drug_or_outcome> \
--experiment <experiment_name> \
--cuda
The second stage accepts additional features via an optional c object and --confounders flag, as used for the TransNEO clinical/RNA features.
The full training runs are GPU-only and reproducing each setting requires hyperparameter optimisation over 150 trials driven by the Optuna framework, followed by 10-seed refits; see the walltime estimate below.
| Paper figure | Script / notebook |
|---|---|
| Fig. 2 (representation richness) | results_analysis/ccl_drug_resp.ipynb |
| Fig. 3 (out-of-distribution drug response, all drugs) | src/scripts/icovae_hopt.py + src/scripts/pico_sk_hopt.py, submission via src/scripts/schedule_jobs_depmap.py and src/scripts/schedule_jobs_depmap_sk.py; aggregation in results_analysis/ccl_drug_resp.ipynb |
| Fig. 4 (permutation feature importance, cell lines) | results_analysis/ccl_drug_resp.ipynb |
| Fig. 5 (TransNEO RCB / pCR) | src/scripts/schedule_jobs.py; aggregation and figures in results_analysis/transneo_treatment_resp.ipynb |
| Supp. Figs B1–B4 | Same pipelines as above; rendering in the notebooks listed. |
On a single NVIDIA A100 GPU, the stage-1 iCoVAE hyperparameter optimisation for the DepMap/GDSC data (1500 features, ~1000 samples, 150 Optuna trials, 300 epochs per trial) takes approximately 8 hours. Stage-2 PiCo head fitting and refit across 10 seeds takes ~15 min per drug. Full reproduction of the 65-drug cell-line experiment is therefore ≳ 200 GPU-hours.
- Hyperparameter-optimisation process — docs/examples/slurm/icovae_hopt.sh and the
schedule_jobs*.pyscripts undersrc/scripts/(these submit training jobs to a SLURM HPC cluster; see the SLURM templates README for site-specific edits). - Out-of-distribution (OOD) prediction on cancer cell lines — results_analysis/ccl_drug_resp.ipynb.
- Transfer from cell-line experimental data to patient treatment response — results_analysis/transneo_treatment_resp.ipynb.
| Dataset | Version | Link |
|---|---|---|
| DepMap | 23Q2 | DepMap |
| GDSC2 | Oct 2023 | GDSC2 |
| TransNEO & ARTemis+PBCP | — | (i) RNA-Seq (ii) Response |
The data and pretrained encoders included under demo/assets/ are derived from the DepMap, GDSC2, and TransNEO/ARTemis+PBCP sources above. The build script that produced them (scripts/build_demo_bundle.py) shows the exact selection / preprocessing applied to each file.
| Path | Description |
|---|---|
| demo/ | Two Colab/local demo notebooks (demo_ccl.ipynb, demo_transneo.ipynb) plus the pretrained encoders and small preprocessed data they read from assets/ |
| results_analysis/ | Notebooks used to produce the figures in the paper |
| src/ | Core PiCo package (see src/README.md for a directory map) |
| src/models/ | iCoVAE and PiCo model classes |
| src/scripts/ | Training / hyperparameter-optimisation entry points |
| scripts/ | Demo-build utilities (build_demo_bundle.py, build_demo_notebooks.py) |
| src/utils/ | Data-loading and comparison utilities (PerfComp, calculate_feat_imps, plot_helpers) |
| docs/examples/slurm/ | Example SLURM submission templates for full training runs on a Cambridge CSD3 Wilkes3-style cluster |
| MODEL_CARD.md | Model card (intended use, training data, evaluation, limitations, environmental impact) |
Please raise issues in the repository or email dom.kirkham@mrc-bsu.cam.ac.uk.
If you find this work interesting or use the code here, please cite our paper:
@misc{kirkham_interventionally-guided_2025,
title = {Interventionally-guided representation learning for robust and interpretable AI models in cancer medicine},
copyright = {© 2025, Posted by Cold Spring Harbor Laboratory. This pre-print is available under a Creative Commons License (Attribution 4.0 International), CC BY 4.0, as described at http://creativecommons.org/licenses/by/4.0/},
url = {https://www.biorxiv.org/content/10.1101/2025.07.21.662350},
doi = {10.1101/2025.07.21.662350},
language = {en},
urldate = {2025-07-22},
publisher = {bioRxiv},
author = {Kirkham, Dom and Masina, Riccardo and Sammut, Stephen-John and Mukherjee, Sach and Rueda, Oscar M.},
month = jul,
year = {2025},
}