Skip to content

tausurv.nn.functional.dsm

dsm_nll(predictions: Tensor, event_time: Tensor, event_indicator: Tensor, *, distribution: Literal['weibull', 'lognormal'] = 'weibull', reduction: str = 'mean') -> Tensor

Deep Survival Machines mixture NLL (Nagpal et al., 2021).

The model predicts 3K3K raw values per subject — interpreted as KK shape (or location), KK scale, and KK mixture-weight logits. After constraining shape/scale to be positive and softmaxing the weights, the marginal survival and density are mixtures:

S^(t∣x)=∑k=1Kwk(x) Sk(t; θk(x)),f^(t∣x)=∑k=1Kwk(x) fk(t; θk(x)).\hat S(t \mid x) = \sum_{k=1}^K w_k(x)\, S_k\big(t;\, \theta_k(x)\big), \qquad \hat f(t \mid x) = \sum_{k=1}^K w_k(x)\, f_k\big(t;\, \theta_k(x)\big).

The log-likelihood is computed with logsumexp for stability:

  • Event (δi=1\delta_i = 1): log⁡f^(ti∣xi)=logsumexp⁡k(log⁡wk+log⁡fk(ti))\log \hat f(t_i \mid x_i) = \operatorname{logsumexp}_k\big(\log w_k + \log f_k(t_i)\big).
  • Censored (δi=0\delta_i = 0): log⁡S^(ti∣xi)=logsumexp⁡k(log⁡wk+log⁡Sk(ti))\log \hat S(t_i \mid x_i) = \operatorname{logsumexp}_k\big(\log w_k + \log S_k(t_i)\big).
  • predictions — (n, 3K) tensor — Raw output of DSM.forward. Columns [:K], [K:2K], [2K:3K] are the raw shape (or location), raw scale, and raw mixture logits respectively.
  • event_time — (n,) tensor
  • event_indicator — (n,) tensor — δ=1\delta = 1 for events, 00 for right-censored.
  • distribution — lognormal, default “weibull” — Per-component distribution family. Must match the model’s config.distribution_family.
  • reduction — none, default “mean”

Nagpal, C., Li, X., Dubrawski, A. (2021). Deep Survival Machines: Fully Parametric Survival Regression and Representation Learning for Censored Data with Competing Risks. IEEE JBHI 25(8).