Skip to content

About

JEPA-style world model for digital liver disease progression — predicts future 8-D clinical states from observed trajectories.

Topics

Resources

Stars

0 stars

Watchers

0 watching

Forks

Repository files navigation

Digital Liver World Model

A JEPA-style world model for digital liver disease progression. Predicts future 8-D clinical states from observed trajectories.

State vector

idx Field Range Behaviour
0 F (fibrosis) [0,1] Ratchet, non-decreasing
1 D (ductopenia) [0,1] Ratchet, irreversible
2 S (strictures) [0,1] Ratchet, drops at ERCP
3 P (portal HTN) [0,1] Ratchet, non-decreasing
4 A (inflammation) [0,1] Mean-reverting
5 C (cholestasis) [0,1] Fast, with flares
6 M (malignancy) [0,2] Monotone accumulator
7 flare [0,1] Transient, decays

Architecture

encoder (GAT with causal disease-graph mask)
    ├─→ GRU latent predictor → predicted future latent
    │                                    ↓
    └─→ EMA target encoder → target latent   JEPA loss
                                         ↓
                              decoder (MLP → 8-D state)
                                         ↓
                           constraint enforcement (cummax/cummin)
                                         ↓
                           integrated gradients (explainability)
image

The final architecture combines:

  • Graph Attention Network (GAT) encoder — follows the disease graph (Inflammation → Fibrosis, Cholestasis → Fibrosis)
  • EMA target encoder (momentum 0.996) — stable prediction targets
  • GRU latent predictor — models latent dynamics
  • Clinical decoder (MLP → 8-D state)
  • Constraint enforcement module — post-hoc monotonicity via torch.cummax/torch.cummin
  • Integrated Gradients explainability

JEPA-style: predicts future latents, not raw values. The optimization objective combines:

L = λ_j * L_JEPA + λ_r * L_Recon + λ_v * L_Var + λ_c * L_Cov

VICReg variance regularization prevents latent dimensions from collapsing to constant values; covariance regularization reduces redundancy between latent channels. Constraints are enforced after decoding using cumulative maximum and cumulative minimum operators, guaranteeing biologically valid trajectories regardless of prediction error.

Usage

pip install -r requirements.txt
python train.py
python evaluate.py
python explain.py

MLOps quickstart

Training is deterministic for a fixed seed and writes metrics to outputs/logs/metrics.jsonl. Set MLFLOW_TRACKING_URI to mirror the same parameters, metrics, and run metadata to an MLflow server; MLflow remains optional for local runs. The CI workflow runs linting and tests on every push and pull request.

cd digital_liver_world_model
pip install -r requirements.txt
python train.py

The RequestMetrics helper in telemetry.py provides request counts, latency, errors, and constraint-violation counters for an API layer. Clinical data is not bundled: connect an approved, de-identified source through a separate ETL adapter and document its governance before training.

Results

Trained on 300 synthetic trajectories (120 months each), predicts 12-month horizon from 24-month context. 80/10/10 train/val/test split (20,400 / 2,550 / 2,550 samples). Early stopping at epoch 80 (patience 30).

Metric Value
MAE 0.207
Constraint violations (raw) 0.10%
Constraint violations (enforced) 0.00%
Latent cosine similarity 0.953
Latent std 4.285
Latent cov (off-diagonal) 6.776
High susceptibility MAE 0.201
Low susceptibility MAE 0.160
24-month rollout MAE 0.162
Rollout latent drift 1.290

Constraint projection eliminated all violations without increasing MAE. The model generalized to unseen susceptibility values, although autoregressive error accumulated during long rollouts; rollout latent drift confirms this is gradual drift in representation space rather than a discrete failure.

Explainability: For one representative high-susceptibility trajectory, Integrated Gradients identified prolonged cholestasis and repeated inflammatory flares as the dominant contributors to the predicted deterioration around month 30, ahead of portal hypertension, fibrosis, and malignancy — consistent with the generator, where sustained inflammatory activity raises cholestasis, accelerates fibrosis, and subsequently raises Malignancy.

Baselines and ablations

The reproducible experiment matrix is in experiments/. It covers persistence, GRU-only, full JEPA/VICReg, EMA removal, VICReg removal, and constraint removal. Run python experiments/run_matrix.py --dry-run to review it, then keep the seed and data manifest fixed for every run. Results must report short-horizon MAE, 24-month rollout MAE, raw/enforced violation rates, and latent standard deviation/covariance. Current results are synthetic simulator evidence and do not establish clinical utility.

Model Development

Key implementation issues that were corrected:

  • Incorrect temporal alignment between predictions and targets → fixed with autoregressive rollout
  • GRU outputs gradually lost variance → fixed with residual connections and layer normalization
  • Excessive constraint weighting encouraged overly flat trajectories → fixed with balanced loss weighting

Limitations and Future Work

The current implementation was evaluated only on trajectories generated by the synthetic simulator. Successful performance demonstrates recovery of simulator dynamics rather than clinical generalization.

Future work:

  1. Constraint-preserving latent parameterizations (monotonicity by construction, not post-processing) with explicit hazard-memory mechanism for the Fibrosis × Cholestasis interaction
  2. λ_v=0 ablation to isolate VICReg's contribution to collapse prevention
  3. Latent manifold critic to distinguish biologically valid trajectories from implausible latent transitions
  4. Counterfactual intervention scenarios under alternative treatment schedules (e.g., earlier UDCA, different ERCP timing)

About

JEPA-style world model for digital liver disease progression — predicts future 8-D clinical states from observed trajectories.

Topics

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages