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:
objectMULTImodal 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,
MaxFusionis 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
MultiSurvDatasetExample
>>> 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.