tausurv.nn.models.dsm
Deep Survival Machines (Nagpal et al., 2021).
A neural net predicts a per-subject mixture of parametric distributions (Weibull or LogNormal). The marginal survival is a weighted sum of the per-component survivals:
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- or Newton inversion needed.
The model outputs a flat (n, 3K) raw tensor: the first columns
are raw shape (Weibull) / location (LogNormal), next are raw scale
/ standard deviation, last 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.
References
Section titled “References”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).
Reference
Section titled “Reference”DSMConfig class
Section titled “DSMConfig ”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.
DSM class
Section titled “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.