Survival

MultiSurv

class imml.survival.MultiSurv(input_dim: list = None, fusion: object = None, t_binds: int = 30, hidden_dim: int = 512, embed_size: int = 512, extractors: list = None, n_layers: int = 4, time_points: list = None, learning_rate: float = 0.0001, weight_decay: float = 0.0001)[source]

Bases: object

MULTImodal SURVival prediction (MultiSurv). [1] [2]

MultiSurv is a multimodal discrete-time survival model. Each modality is mapped to a shared latent feature size, available modality representations are fused, and a fully connected block predicts conditional survival probabilities for a sequence of time intervals.

This class provides training, validation, testing, and prediction logic compatible with the Lightning Trainer.

Parameters:
  • input_dim (list of int, default=None) -- A list specifying the input dimensions for each default tabular modality extractor.

  • fusion (torch.nn.Module, default=None) -- Fusion module used when more than one modality is available. If None, MaxFusion is used.

  • t_binds (int, default=30) -- Number of output survival intervals.

  • hidden_dim (int, default=512) -- Size of each modality representation, matching the original MultiSurv default.

  • embed_size (int, default=512) -- Size of the representation passed to the final risk layer.

  • extractors (list of torch.nn.Module, default=None) -- Custom feature extractors for each modality. If None, fully connected tabular extractors equivalent to the original omics submodels are used.

  • n_layers (int, default=4) -- Number of layers in the post-fusion fully connected block.

  • time_points (array-like, default=None) -- Interval break points in days. If None, yearly breaks from 0 to 30 years are used.

  • learning_rate (float, default=1e-4) -- Learning rate for the optimizer.

  • weight_decay (float, default=1e-4) -- Weight decay used by the optimizer.

References

See also

MultiSurvDataset

Example

>>> from lightning import Trainer
>>> import numpy as np
>>> import pandas as pd
>>> from torch.utils.data import DataLoader
>>> from imml.survival import MultiSurv
>>> from imml.load import MultiSurvDataset
>>> Xs = [pd.DataFrame(np.random.default_rng(42).random((8, 10))) for _ in range(2)]
>>> y = pd.DataFrame({"time": np.random.default_rng(42).uniform(0.5, 5, 8),
...                   "event": np.random.default_rng(42).integers(0, 2, 8)})
>>> train_data = MultiSurvDataset(Xs=Xs, y=y)
>>> train_dataloader = DataLoader(dataset=train_data, batch_size=4, shuffle=True)
>>> trainer = Trainer(max_epochs=1, logger=False, enable_checkpointing=False)
>>> estimator = MultiSurv(input_dim=[10, 10])
>>> trainer.fit(estimator, train_dataloader)
>>> trainer.predict(estimator, train_dataloader)
training_step(batch, batch_idx=None)[source]

Method required for training using Lightning Trainer.

validation_step(batch, batch_idx=None)[source]

Method required for validating using Lightning Trainer.

test_step(batch, batch_idx=None)[source]

Method required for testing using Lightning Trainer.

predict_step(batch, batch_idx=None)[source]

Method required for predicting using Lightning Trainer.

configure_optimizers()[source]

Method required for training using Lightning Trainer.