A JEPA-style world model for digital liver disease progression. Predicts future 8-D clinical states from observed trajectories.
| 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 |
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)
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.
pip install -r requirements.txt
python train.py
python evaluate.py
python explain.pyTraining 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.pyThe 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.
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.
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.
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
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:
- Constraint-preserving latent parameterizations (monotonicity by construction, not post-processing) with explicit hazard-memory mechanism for the Fibrosis × Cholestasis interaction
- λ_v=0 ablation to isolate VICReg's contribution to collapse prevention
- Latent manifold critic to distinguish biologically valid trajectories from implausible latent transitions
- Counterfactual intervention scenarios under alternative treatment schedules (e.g., earlier UDCA, different ERCP timing)