pyhealth.tasks.mortality_prediction_stagenet_mimic4#
- class pyhealth.tasks.mortality_prediction_stagenet_mimic4.MortalityPredictionStageNetMIMIC4(padding=0)[source]#
Bases:
BaseTaskTask 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.
- 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)
- output_schema#
The schema for output data: - mortality: Binary indicator (1 if any admission had mortality)
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)
- 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