Skip to content
AI.info

Research

Let the Experts Speak: Improving Survival Prediction & Calibration via Mixture-of-Experts Heads

Overview Research area: Machine learning for healthcare — discrete-time survival analysis (time-to-event prediction under right-censoring), specifically mixture-of-experts (MoE) neural architectures e

Let the Experts Speak: Improving Survival Prediction & Calibration via Mixture-of-Experts Heads
arXiv
2511.09567
Published
2025-11-11
Authors
Todd Morrill, Aahlad Puli, Murad Megjhani, Soojin Park, Richard Zemel

AI summary

Overview

Research area: Machine learning for healthcare — discrete-time survival analysis (time-to-event prediction under right-censoring), specifically mixture-of-experts (MoE) neural architectures evaluated on calibration, discrimination, and latent patient clustering.

Technical level: Intermediate. The paper is written for readers comfortable with neural network heads, censorship-aware loss functions, and survival metrics, though the core architectural idea (how much freedom each expert gets) is explained conceptually.

Scope in one sentence: This paper designs and compares three discrete-time deep MoE survival heads that differ only in how expressive their experts are, and shows that per-patient ("personalized") experts simultaneously achieve latent clustering, better calibration, and better predictive accuracy on two real-world clinical datasets.

What This Paper Is About

Mixture-of-experts survival models are attractive in medicine because they can group similar patients into latent clusters, but this grouping usually comes with a cost: the model assumes a patient's predicted event distribution must look like the distribution of the group they were assigned to, which can hurt calibration and accuracy. The authors ask whether a model can still recover patient group structure where it genuinely exists while improving calibration and predictive accuracy. They answer this by building three MoE heads that are identical except for how much each expert can tailor its predicted event distribution to the individual patient.

Key Contributions

  1. Three new discrete-time deep MoE survival head architectures — Fixed MoE (each expert learns one event distribution shared by all patients assigned to it), Adjustable MoE (each expert learns a prototype distribution that is then warped per patient via a monotone bijection), and Personalized MoE (each expert constructs a custom event distribution for each patient from a chunked expert representation).
  2. A demonstration that expert expressivity is the key differentiator among discrete-time deep MoE survival models, isolated by a controlled experiment in which the three heads differ essentially only in that respect under matched parameter counts.
  3. Identification of one architecture that satisfies all three desiderata — the Personalized MoE achieves clustering, calibration, and predictive accuracy together, while the Fixed and Adjustable MoEs generally do not outperform the MTLR baseline on all metrics on real-world data.
  4. A routing and clustering analysis, including recovery of latent digit groups in a synthetic Survival MNIST dataset and clinician-informed signatures for the eight patient clusters discovered by a Personalized MoE trained on SUPPORT2, plus a stability measurement of routing across random seeds (Adjusted Rand Index of 0.36).

Main Findings

  • Personalized MoE wins on real-world data across all metrics. On SUPPORT2 it reports ECE 0.048 (average deviation −0.009 from MTLR), concordance 80.84 (+0.93), and Brier scores of 0.154, 0.142, and 0.138 at the 25th, 50th, and 75th percentile time bins (deviations −0.002, −0.007, −0.009) — outperforming CoxPH, RSF, and MTLR on every reported metric.

  • The same pattern holds on Sepsis. Personalized MoE reaches ECE 0.005 (−0.012 vs MTLR), concordance 89.77 (+1.41), and Brier scores 0.017, 0.030, 0.036 (−0.002, −0.003, −0.003), again best across all methods, including the very large errors of CoxPH (ECE 0.635) and RSF (ECE 0.604).

  • Fixed MoE is optimal when latent groups are perfectly specified. On the synthetic Survival MNIST dataset, the Fixed MoE achieves the best concordance (93.46, +0.65 over MTLR) and matches MTLR on Brier scores at all three time points, whereas the Personalized MoE is the best calibrated there (ECE 0.005, −0.001) and matches Fixed MoE on Brier at the 25th and 75th percentiles. The authors describe Survival MNIST as a "Platonic ideal" where the Fixed MoE is exactly specified with 10 expert heads, one per digit.

  • More expressive experts mean less sensitivity to model specification. Sweeping the number of experts from 2 to 20, the Fixed MoE is highly sensitive until enough experts exist to cover the latent groups; the Adjustable MoE is less sensitive; and the Personalized MoE is the least sensitive, because it can form custom event distributions regardless of how many experts exist. The authors describe this as a continuum of sensitivity based on expert expressivity.

  • The models recover latent groups. In Survival MNIST routing analysis, 93% of points routed to Fixed MoE expert 0 are digit 3 and 94% of those routed to expert 1 are digit 5, indicating strong per-expert specialization; the Personalized MoE shows slightly lower but still strong specialization.

  • Personalized MoE discovers clinically interpretable clusters on SUPPORT2. Eight clusters were found with sizes of 70, 3, 173, 47, 78, 221, 188, and 131, ranging from a lowest-risk cluster of 50–60s men with metastatic colon cancer and ESRD/HD (n=47) to a highest-risk cluster of 80s patients with dementia, early DNR, and ARF/MOSF ± sepsis/coma (n=221). The n=3 cluster is noted as possibly mergeable.

  • Routing is moderately stable across random seeds. Using the top-1 routing rule, the Adjusted Rand Index for the Personalized MoE is 0.36.

  • An explicit caveat on scope. The authors state they do not claim state-of-the-art across all survival benchmarks; their claims are relative and internal, with Survival MNIST serving as a counterexample showing when fixed prototypes are optimal.

Methodology in Plain English

All three models share the same backbone: raw patient records (categorical indicators embedded, continuous features standardized) are fed through a feedforward network to a final hidden representation, and the last layer is the MoE head. Everything is trained end-to-end with a discrete-time Multitask Logistic Regression (MTLR) style loss that predicts a monotone label sequence — once the event occurs, it stays on — with separate uncensored and censored negative log-likelihood terms, plus a load-balancing regularizer that penalizes low-entropy average expert weights so all experts get used.

The three heads differ in one dimension: how much a single expert's prediction can change from patient to patient.

  • Fixed MoE keeps a learnable matrix M ∈ ℝ^{n×m} where each row is an event-time distribution, and a linear router with a learnable temperature κ produces softmax weights over the n experts. The final prediction is a weighted average of those fixed rows — so two patients with the same routing weights get identical predictions.
  • Adjustable MoE starts from a prototype score vector per expert but warps it per patient. The warping is a strictly monotone bijection built from a normalized mixture of r = 2 logistic CDFs whose weights, slopes, and ordered centers are all learned linear functions of the patient's hidden state, then endpoint-normalized to map [0,1] to [0,1]. An inverse map (computed with a bisection solver) tells the model where each canonical time gridpoint lands on the expert's internal time axis, and scores are linearly interpolated there. This is essentially a flexible, monotone, per-patient time warp of the expert's distribution.
  • Personalized MoE gives experts the most freedom. The hidden state is projected into separate router (W_r ∈ ℝ^{h×h}) and expert (W_e ∈ ℝ^{h×h}) representations; the expert representation is split into n equal chunks, each passed through its own linear layer L_k ∈ ℝ^{m×(h/n)} to produce a fresh event distribution for that patient. Chunking keeps the parameter count efficient and may push experts to use independent information.

Experiments span three datasets. Survival MNIST is a synthetic dataset with a distinct event distribution per digit, 15% censoring, used to test latent-group recovery. SUPPORT2 has 9,105 examples with roughly 32% censoring. Sepsis has 40,336 patient records but only 2,932 positive sepsis instances, with the first 100 hours of each ICU stay summarized and administrative censoring at 100 hours.

Evaluation uses Harrell's concordance index, the time-dependent IPCW Brier score at the 25th, 50th, and 75th percentile time bins, and equal-mass expected calibration error (ECE) adjusted with IPCW and averaged over time bins. All neural models are parameter-count matched, all numbers are averaged over 5 random seeds, and the parenthetical values are the average per-seed gap from the MTLR baseline. Baselines are CoxPH, random survival forests (RSF), and MTLR.

Why This Matters

Research impact. The paper isolates a variable — expert expressivity — that prior MoE survival work had not carefully varied under matched conditions, and shows it can be the deciding factor between a model that clusters well, a model that is calibrated and accurate, and a model that does all three. It also shows that per-patient expert distributions can be implemented with straightforward end-to-end supervised training, avoiding the variational inference, spline baseline-hazard estimation, or EM-based cluster estimation used in much of the prior clustering-for-survival literature.

Real-world applications.

  • Clinical decision support systems for ICU patients, where calibrated risk probabilities (e.g., mortality or sepsis onset at 25/50/75th percentile time points) feed directly into treatment escalation decisions.
  • Reasoning by analogy to similar patients — the SUPPORT2 clusters (age bands, comorbidities such as COPD/CHF/cirrhosis, code status, income level) let a clinician see which historical patient group a new patient most resembles.
  • Retrospective cohort triage and resource planning, where identifying high-risk subgroups such as the 80s dementia/early-DNR cluster or the lowest-risk metastatic colon cancer cluster could inform monitoring intensity.
  • Accelerating human labeling of partially labeled historical datasets, which the authors cite as a motivation for the retrospective 100-hour Sepsis prediction setup.

Industry relevance. Any organization building risk models on censored longitudinal data — hospital systems, health insurers, clinical AI vendors, and EHR analytics teams — faces the same trade-off between model interpretability/grouping and probability quality. This paper's message is that these goals need not conflict: expressivity in the head, rather than a simplification of it, is what buys calibration, discrimination, and cluster structure at once. The code is released at https://github.com/ToddMorrill/survival-moe.

Future Directions

  • Expert pruning and consolidation. The Personalized MoE is relatively insensitive to the number of experts, which limits its use for discovering the number of latent groups. The authors suggest two heuristics: use the Fixed MoE with an elbow method to estimate the group count and then plug that into the Personalized MoE, or over-specify experts and prune/consolidate based on routing behavior — the latter is left for future work.
  • Simpler warping families for the Adjustable MoE. The authors want to explore alternative monotone transformation families that naturally map [0,1] to [0,1] with fewer moving parts.
  • Time-varying inputs and outputs. All methods can in principle be attached to a recurrent network or Transformer, but how to interpret patient groups when inputs and predictions change over time is an explicitly open question.
  • Broader comparative baselines. The limitations section states that adding further model classes such as DeepHit and continuous-time parametric mixtures is orthogonal to the central question and left for future comparative work. Relatedly, the paper notes an unanswered question about how many representation dimensions are needed when the event-time distribution is more complex than the datasets explored here (the current setup uses 100 discrete event times and typically a representation of size 16, i.e. 128 hidden dimensions divided by 8 experts).

Target Audience

This paper is most useful to machine learning researchers working on survival analysis and mixture-of-experts architectures, applied ML scientists building clinical risk prediction or decision support systems, and clinical informatics teams who need models that are simultaneously calibrated, accurate, and interpretable through patient similarity. It is also relevant to statisticians and biostatisticians familiar with Cox models, random survival forests, and MTLR who want to understand what deep MoE heads add — and where they do not help, as the Survival MNIST counterexample illustrates.

Authors’ abstract

Deep mixture-of-experts models have attracted a lot of attention for survival analysis problems, particularly for their ability to cluster similar patients together. In practice, grouping often comes at the expense of key metrics such as calibration error and predictive accuracy. This is due to the restrictive inductive bias that mixture-of-experts imposes, that predictions for individual patients must look like predictions for the group they're assigned to. Might we be able to discover patient group structure, where it exists, while improving calibration and predictive accuracy? In this work, we introduce several discrete-time deep mixture-of-experts (MoE)-based architectures for survival analysis problems, one of which achieves all desiderata: clustering, calibration, and predictive accuracy. We show that a key differentiator between this array of MoEs is how expressive their experts are. We find that more expressive experts that tailor predictions per patient outperform experts that rely on fixed group prototypes.

Read the original paper