pyhealth.tasks.mortality_prediction_stagenet_mimic4#

class pyhealth.tasks.mortality_prediction_stagenet_mimic4.MortalityPredictionStageNetMIMIC4(padding=0)[source]#

Bases: BaseTask

Task for predicting mortality using MIMIC-IV with StageNet format.

This task creates PATIENT-LEVEL samples (not visit-level) by aggregating all admissions for each patient. ICD codes (diagnoses + procedures) and lab results across all visits are combined with time intervals calculated from the patient’s first admission timestamp.

Time Calculation:
  • ICD codes: Hours from previous admission (0 for first visit, then time intervals between consecutive visits)

  • Labs: Hours from admission start (within-visit measurements)

Lab Processing:
  • 10-dimensional vectors (one per lab category)

  • Multiple itemids per category → take first observed value

  • Missing categories → None/NaN in vector

Parameters:

padding (int) – Additional padding for StageNet processor to handle sequences longer than observed during training. Default: 0.

task_name#

The name of the task.

Type:

str

input_schema#

The schema for input data: - icd_codes: Combined diagnosis + procedure ICD codes

(stagenet format, nested by visit)

  • labs: Lab results (stagenet_tensor, 10D vectors per timestamp)

Type:

Dict[str, str]

output_schema#

The schema for output data: - mortality: Binary indicator (1 if any admission had mortality)

Type:

Dict[str, str]

Examples

>>> from pyhealth.datasets import MIMIC4EHRDataset
>>> from pyhealth.tasks import MortalityPredictionStageNetMIMIC4
>>> dataset = MIMIC4EHRDataset(
...     root="/path/to/mimic-iv/2.2",
...     tables=["diagnoses_icd", "procedures_icd", "labevents"],
... )
>>> task = MortalityPredictionStageNetMIMIC4()
>>> samples = dataset.set_task(task)
task_name: str = 'MortalityPredictionStageNetMIMIC4'#
input_schema: Dict[str, Tuple[str, Dict[str, Any]]]#
output_schema: Dict[str, str]#
LAB_CATEGORIES: ClassVar[Dict[str, List[str]]] = {'Anion Gap': ['50868', '52500'], 'Bicarbonate': ['50803', '50804'], 'Calcium': ['50808', '51624'], 'Chloride': ['50806', '52434', '50902', '52535'], 'Glucose': ['50809', '52027', '50931', '52569'], 'Magnesium': ['50960'], 'Osmolality': ['52031', '50964', '51701'], 'Phosphate': ['50970'], 'Potassium': ['50822', '52452', '50971', '52610'], 'Sodium': ['50824', '52455', '50983', '52623']}#
LAB_CATEGORY_NAMES: ClassVar[List[str]] = ['Sodium', 'Potassium', 'Chloride', 'Bicarbonate', 'Glucose', 'Calcium', 'Magnesium', 'Anion Gap', 'Osmolality', 'Phosphate']#
LABITEMS: ClassVar[List[str]] = ['50824', '52455', '50983', '52623', '50822', '52452', '50971', '52610', '50806', '52434', '50902', '52535', '50803', '50804', '50809', '52027', '50931', '52569', '50808', '51624', '50960', '50868', '52500', '52031', '50964', '51701', '50970']#
pre_filter(df)#
Return type:

LazyFrame