A Multi-Phase Co-Training Framework for Joint-Embedding Predictive Pretraining on EHR Patient Trajectories
Important
- [Sep 2026] ✨ Full pipeline released. This repository now contains the complete end-to-end pipeline behind the paper, from raw MIMIC-IV to every figure and table: data processing, training of all six models, latent rollout, and the three evaluations (rollout stability, latent geometry, and the 34-task downstream benchmark). Each stage has a step-by-step guide in
docs/. - [Sep 2026] 🤗 Open-source model release. The pretrained Clin-JEPA backbone, the co-trained encoder adapters together with the latent trajectory predictor, is being prepared for public release with loading code. Details will be announced here.
Note: This repository accompanies a paper currently under double-blind review. Author, citation, and funding information have been omitted from this version and will be added once review concludes.
Clin-JEPA is a latent world model for intensive-care patient trajectories. An LLM encoder reads each hour of the record as text, observations and interventions separately, so no feature engineering or vocabulary mapping is needed. An action-conditioned Transformer predictor rolls the patient state forward in the encoder's latent space under the recorded interventions, and is kept at inference.
Instead of freezing the encoder before the predictor is trained, Clin-JEPA trains the two together under one latent-prediction objective. A five-phase curriculum (predictor warmup, co-training, EMA target alignment, hard sync, predictor finalization) keeps this co-training stable: the representation does not collapse and the predictor's target space does not drift. In the paper's experiments the encoder is Qwen3-8B with LoRA adapters and the predictor has 92M parameters; docs/architecture.md describes the components, the curriculum and the path of each of the six models through the stages.
The code that takes MIMIC-IV v3.1 to every figure and table of the paper's experiments:
- the data pipeline: cohort, hourly observations and interventions, patient-level split, text trajectories;
- training of the six models the paper compares: Clin-JEPA, the V-JEPA 2-AC style and SFT baselines, the two curriculum ablations, and the variant with a predictor trained separately on the frozen Clin-JEPA encoder;
- one latent rollout per model, shared by all evaluations;
- the three evaluations: stability of the co-trained encoder and predictor, latent geometry, and 34 downstream tasks on three benchmarks against Ridge, LightGBM, LSTM, TCN and CLMBR-T-base.
git clone <repository-url> && cd Clin-JEPA
pip install -r requirements.txt && pip install -e .
pip install flash-attn==2.8.3 --no-build-isolation # on a GPU node with the CUDA toolkit
cp .env.example .env # paths to MIMIC-IV, the outputs and the Hugging Face cacheThen follow the guides in order, starting with setup: data access, environment variables and hardware. pip install -e ".[dev]" && pytest runs the unit tests on synthetic data in a few seconds.
Compute: the data stage runs on CPU nodes in under an hour as ten-task arrays. Clin-JEPA pretraining took about 430 GPU-hours on 8 H200 GPUs, encoder SFT about 40 and the V-JEPA 2-AC style refinement about 50. An embedding cache takes about 22 GPU-hours and 96 GB per encoder. The evaluations run on CPU nodes and a single GPU.
Each stage has a guide in docs/ with its inputs, commands, outputs and resources.
| stage | guide | code | produces |
|---|---|---|---|
| 00 | setup | environment, data access, paths | |
| 01 | data | clin_jepa.data |
cohort, hourly bins, split, trajectory shards, the feature table of Appendix A |
| 02 | encoders | clin_jepa.training |
encoder SFT, Clin-JEPA curriculum, V-JEPA 2-AC style refinement, the two ablations |
| 03 | embeddings | clin_jepa.encode |
one embedding cache per encoder |
| 04 | predictors | clin_jepa.training |
predictors on frozen encoders |
| 05 | rollout | clin_jepa.rollout |
pooled rollout blocks and per-step errors |
| 06 | stability | clin_jepa.stability |
Fig. 2, Table 2, Appendices C–D |
| 07 | geometry | clin_jepa.geometry |
Fig. 3, Table 3, Appendix E |
| 08 | downstream | clin_jepa.downstream |
Fig. 4, Table 4, the per-task tables |
clin_jepa/
data/ MIMIC-IV to text trajectories
model/ action-conditioned predictor, masked-prediction head
training/ encoder SFT, curriculum, V-JEPA 2-AC style refinement, predictors on frozen encoders
encode.py embedding caches
rollout.py latent rollout
stability/ rollout drift, rank diagnostics, report
geometry/ cohorts, native-space statistics, UMAP view, report
downstream/ labels, features, read-outs, baselines, bootstrap, report, clmbr/
configs/ one YAML per stage and per training run
docs/ numbered guides
slurm/ job templates
tests/ unit tests on synthetic data
Experiments use MIMIC-IV v3.1 (Johnson et al., 2023), available from PhysioNet under credentialed access. No patient data, derived features or model outputs are distributed with this repository.
Citation information has been omitted while this work is under double-blind review, and will be added once review concludes.
This code is released under the MIT License. MIMIC-IV data is governed separately by the PhysioNet Credentialed Health Data License.
We thank the MIMIC-IV team at the MIT Laboratory for Computational Physiology and Beth Israel Deaconess Medical Center for making the dataset publicly available. Funding acknowledgments have been omitted for anonymity and will be added once review concludes.
