Skip to content

tausurv.nn.models.dsm

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

A neural net predicts a per-subject mixture of KK parametric distributions (Weibull or LogNormal). The marginal survival is a weighted sum of the per-component survivals:

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

DSM is the parametric counterpart to DeepHit: DeepHit puts a PMF over discrete bins, DSM uses a mixture of continuous parametric families. Everything (density, survival, prediction) is closed-form — no autograd-through-tt or Newton inversion needed.

The model outputs a flat (n, 3K) raw tensor: the first KK columns are raw shape (Weibull) / location (LogNormal), next KK are raw scale / standard deviation, last KK are mixture-weight logits. Constraints (softplus for positivity, softmax for weights) are applied by the loss and the predict methods, so the model’s forward stays a clean nn.Linear output.

Pair with tausurv.nn.functional.dsm_nll or tausurv.nn.losses.DSMLoss for training. Standard tausurv.nn.Trainer works — DSM doesn’t need a Trainer subclass.

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).

class DSMConfig(in_features: int, n_components: int = 4, distribution_family: Literal['weibull', 'lognormal'] = 'weibull', hidden_features: tuple[int, ...] = (64, 64), activation: Literal['gelu', 'relu'] = 'gelu', norm: Literal['layer', 'batch', 'none'] = 'layer', dropout: float = 0.1, residual: bool = True)

Architectural configuration for DSM.

class DSM(config: DSMConfig | None = None, **kwargs)

Deep Survival Machines model.

Construct with DSMConfig or kwargs (HF-style). Train with fit, or with the standard tausurv.nn.Trainer and tausurv.nn.functional.dsm_nll; either way the model inherits the unified prediction API from SurvivalPredictor.

The default predict grid is data-derived state: fit sets it to the training event times, set_time_grid sets it explicitly, and times= on the predict methods overrides it per call.