From 198528bbd297476b82f115385f019114de00642b Mon Sep 17 00:00:00 2001 From: dylan-g12 Date: Sun, 19 Apr 2026 12:26:55 -0500 Subject: [PATCH 1/8] Added patient_readmission.py and test_patient_readmission.py --- .github/patient_readmission.py | 129 ++++++++++++++++++++++++++++ .github/test_patient_readmission.py | 105 ++++++++++++++++++++++ 2 files changed, 234 insertions(+) create mode 100644 .github/patient_readmission.py create mode 100644 .github/test_patient_readmission.py diff --git a/.github/patient_readmission.py b/.github/patient_readmission.py new file mode 100644 index 000000000..c6292028f --- /dev/null +++ b/.github/patient_readmission.py @@ -0,0 +1,129 @@ +from typing import Dict, List +from pyhealth.data import Event, Patient +from pyhealth.tasks import BaseTask + +class ReadmissionPredictionEICU(BaseTask): + """ + Readmission prediction on the eICU dataset. + + This task aims at predicting whether the patient will be readmitted into the ICU + during the same hospital stay based on clinical information from the current ICU + visit. + + Features: + - using diagnosis table (ICD9CM and ICD10CM) as condition codes + - using physicalexam table as procedure codes + - using medication table as drugs codes + + Attributes: + task_name (str): The name of the task. + input_schema (Dict[str, str]): The schema for the task input. + output_schema (Dict[str, str]): The schema for the task output. + + Examples: + >>> from pyhealth.datasets import eICUDataset + >>> from pyhealth.tasks import ReadmissionPredictionEICU + >>> dataset = eICUDataset( + ... root="/path/to/eicu-crd/2.0", + ... tables=["diagnosis", "medication", "physicalexam"], + ... ) + >>> task = ReadmissionPredictionEICU(exclude_minors=True) + >>> sample_dataset = dataset.set_task(task) + """ + + task_name: str = "ReadmissionPredictionEICU" + input_schema: Dict[str, str] = { + "conditions": "sequence", + "procedures": "sequence", + "drugs": "sequence", + } + output_schema: Dict[str, str] = {"readmission": "binary"} + + def __init__(self, exclude_minors: bool = True, **kwargs) -> None: + """Initializes the task object. + + Args: + exclude_minors: Whether to exclude patients whose age is + less than 18. Defaults to True. + **kwargs: Passed to :class:`~pyhealth.tasks.BaseTask`, e.g. + ``code_mapping``. + """ + super().__init__(**kwargs) + self.exclude_minors = exclude_minors + + def __call__(self, patient: Patient) -> List[Dict]: + """ + Generates binary classification data samples for a single patient. + + Args: + patient (Patient): A patient object. + + Returns: + List[Dict]: A list containing a dictionary for each patient visit with: + - 'visit_id': eICU patientunitstayid. + - 'patient_id': eICU uniquepid. + - 'conditions': Diagnosis codes from diagnosis table. + - 'procedures': Physical exam codes from physicalexam table. + - 'drugs': Drug names from medication table. + - 'readmission': binary label (1 if readmitted, 0 otherwise). + """ + patient_stays = patient.get_events(event_type="patient") + if len(patient_stays) < 2: + return [] + sorted_stays = sorted( + patient_stays, + key=lambda s: ( + int(getattr(s, "patienthealthsystemstayid", 0) or 0), + int(getattr(s, "unitvisitnumber", 0) or 0), + ), + ) + samples = [] + for i in range(len(sorted_stays) - 1): + stay = sorted_stays[i] + next_stay = sorted_stays[i + 1] + if self.exclude_minors: + try: + age_str = str(getattr(stay, "age", "0")).replace(">", "").strip() + if int(age_str) < 18: + continue + except (ValueError, TypeError): + pass + stay_id = str(getattr(stay, "patientunitstayid", "")) + diagnoses = patient.get_events( + event_type = "diagnosis", filters = [("patientunitstayid", "==", stay_id)] + ) + conditions = [ + getattr(event, "icd9code", "") for event in diagnoses + if getattr(event, "icd9code", None) + ] + physical_exams = patient.get_events( + event_type = "physicalexam", filters = [("patientunitstayid", "==", stay_id)] + ) + procedures = [ + getattr(event, "physicalexampath", "") for event in physical_exams + if getattr(event, "physicalexampath", None) + ] + medications = patient.get_events( + event_type = "medication", filters = [("patientunitstayid", "==", stay_id)] + ) + drugs = [ + getattr(event, "drugname", "") for event in medications + if getattr(event, "drugname", None) + ] + if len(conditions) == 0 and len(procedures) == 0 and len(drugs) == 0: + continue + current_hosp_id = getattr(stay, "patienthealthsystemstayid", None) + next_hosp_id = getattr(next_stay, "patienthealthsystemstayid", None) + readmission = int(current_hosp_id == next_hosp_id and current_hosp_id is not None) + samples.append( + { + "visit_id": stay_id, + "patient_id": patient.patient_id, + "conditions": [conditions], + "procedures": [procedures], + "drugs": [drugs], + "readmission": readmission, + } + ) + + return samples \ No newline at end of file diff --git a/.github/test_patient_readmission.py b/.github/test_patient_readmission.py new file mode 100644 index 000000000..d85921979 --- /dev/null +++ b/.github/test_patient_readmission.py @@ -0,0 +1,105 @@ +import pytest +from datetime import datetime +from pyhealth.data import Patient, Event +from pyhealth.tasks import ReadmissionPredictionEICU + +def test_readmission_prediction_eicu_task(): + """ + Tests the ReadmissionPredictionEICU task using synthetic, in-memory patient data + to ensure tests complete in milliseconds. + """ + # 1. Initialize Task + task = ReadmissionPredictionEICU(exclude_minors = True) + + # 2. Create a mock patient + patient = Patient(patient_id = "test_pat_001") + + # Visit 1: ICU Stay 1 (in Hospital 1) + patient.add_event(Event( + event_type = "patient", + timestamp=datetime(2025, 1, 1), + patienthealthsystemstayid = "hosp_001", + patientunitstayid = "icu_001", + unitvisitnumber = 1, + age = "65" + )) + patient.add_event(Event( + event_type = "diagnosis", + patientunitstayid = "icu_001", + icd9code = "428.0" + )) + patient.add_event(Event( + event_type = "medication", + patientunitstayid = "icu_001", + drugname = "Aspirin" + )) + + # Visit 2: ICU Stay 2 (Readmitted to the SAME hospital, hosp_001) + patient.add_event(Event( + event_type = "patient", + timestamp = datetime(2025, 1, 10), + patienthealthsystemstayid = "hosp_001", + patientunitstayid = "icu_002", + unitvisitnumber = 2, + age = "65" + )) + patient.add_event(Event( + event_type = "physicalexam", + patientunitstayid = "icu_002", + physicalexampath = "cardiovascular|murmur" + )) + + # Visit 3: ICU Stay 3 (Admitted to a DIFFERENT hospital, hosp_002) + patient.add_event(Event( + event_type = "patient", + timestamp = datetime(2025, 5, 1), + patienthealthsystemstayid = "hosp_002", + patientunitstayid = "icu_003", + unitvisitnumber = 1, + age = "65" + )) + patient.add_event(Event( + event_type = "diagnosis", + patientunitstayid = "icu_003", + icd9code = "250.00" + )) + + # 3. Call the task + samples = task(patient) + + # 4. Assertions + # With 3 ICU stays, we expect 2 samples (1->2, 2->3) + assert len(samples) == 2, "Task should generate exactly 2 samples" + + # Check Sample 1 (icu_001 -> icu_002) + assert samples[0]["visit_id"] == "icu_001" + assert samples[0]["readmission"] == 1 + assert "428.0" in samples[0]["conditions"][0] + assert "Aspirin" in samples[0]["drugs"][0] + assert len(samples[0]["procedures"][0]) == 0 + + # Check Sample 2 (icu_002 -> icu_003) + assert samples[1]["visit_id"] == "icu_002" + assert samples[1]["readmission"] == 0 + assert "cardiovascular|murmur" in samples[1]["procedures"][0] + +def test_exclude_minors(): + """Test that the task correctly excludes patients under 18.""" + task = ReadmissionPredictionEICU(exclude_minors = True) + patient = Patient(patient_id = "test_minor") + for i in [1, 2]: + patient.add_event(Event( + event_type="patient", + timestamp=datetime(2025, 1, i), + patienthealthsystemstayid = "hosp_001", + patientunitstayid = f"icu_{i}", + unitvisitnumber = i, + age="10" + )) + patient.add_event(Event( + event_type = "diagnosis", + patientunitstayid = f"icu_{i}", + icd9code = "test" + )) + samples = task(patient) + assert len(samples) == 0, "Task should return 0 samples for minors when exclude_minors=True" \ No newline at end of file From c980c2d7e10e9fe1e4595da2dd3d9697bd6b5aef Mon Sep 17 00:00:00 2001 From: dylan-g12 Date: Sun, 19 Apr 2026 12:53:10 -0500 Subject: [PATCH 2/8] Changed file locations of patient_readmission.py from .github to PyHealth/pyhealth/tasks and test_patient_readmission.py from .github to PyHealth/tests --- {.github => pyhealth/tasks}/patient_readmission.py | 0 {.github => tests}/test_patient_readmission.py | 0 2 files changed, 0 insertions(+), 0 deletions(-) rename {.github => pyhealth/tasks}/patient_readmission.py (100%) rename {.github => tests}/test_patient_readmission.py (100%) diff --git a/.github/patient_readmission.py b/pyhealth/tasks/patient_readmission.py similarity index 100% rename from .github/patient_readmission.py rename to pyhealth/tasks/patient_readmission.py diff --git a/.github/test_patient_readmission.py b/tests/test_patient_readmission.py similarity index 100% rename from .github/test_patient_readmission.py rename to tests/test_patient_readmission.py From 16269b5852e330eb40d908404b6bb45c115059e5 Mon Sep 17 00:00:00 2001 From: dylan-g12 Date: Sun, 19 Apr 2026 13:00:27 -0500 Subject: [PATCH 3/8] Added pyhealth.tasks.patient_readmission.rst and fixed name of class in patient_readmission.py to not copy class names --- docs/api/tasks/pyhealth.tasks.patient_readmission.rst | 4 ++++ pyhealth/tasks/patient_readmission.py | 2 +- 2 files changed, 5 insertions(+), 1 deletion(-) create mode 100644 docs/api/tasks/pyhealth.tasks.patient_readmission.rst diff --git a/docs/api/tasks/pyhealth.tasks.patient_readmission.rst b/docs/api/tasks/pyhealth.tasks.patient_readmission.rst new file mode 100644 index 000000000..7a460216a --- /dev/null +++ b/docs/api/tasks/pyhealth.tasks.patient_readmission.rst @@ -0,0 +1,4 @@ +.. autoclass:: pyhealth.tasks.patient_readmission.PatientReadmissionPredictionEICU + :members: + :undoc-members: + :show-inheritance: \ No newline at end of file diff --git a/pyhealth/tasks/patient_readmission.py b/pyhealth/tasks/patient_readmission.py index c6292028f..928935258 100644 --- a/pyhealth/tasks/patient_readmission.py +++ b/pyhealth/tasks/patient_readmission.py @@ -2,7 +2,7 @@ from pyhealth.data import Event, Patient from pyhealth.tasks import BaseTask -class ReadmissionPredictionEICU(BaseTask): +class PatientReadmissionPredictionEICU(BaseTask): """ Readmission prediction on the eICU dataset. From fc2eff420c474b09927b6fdce13b165d1928b2d6 Mon Sep 17 00:00:00 2001 From: dylan-g12 Date: Sun, 19 Apr 2026 13:01:48 -0500 Subject: [PATCH 4/8] Fix: pyhealth.tasks.patient_readmission.rst was missing necessary lines pyhealth.tasks.BaseTask ======================================= --- docs/api/tasks/pyhealth.tasks.patient_readmission.rst | 3 +++ 1 file changed, 3 insertions(+) diff --git a/docs/api/tasks/pyhealth.tasks.patient_readmission.rst b/docs/api/tasks/pyhealth.tasks.patient_readmission.rst index 7a460216a..d28177b31 100644 --- a/docs/api/tasks/pyhealth.tasks.patient_readmission.rst +++ b/docs/api/tasks/pyhealth.tasks.patient_readmission.rst @@ -1,3 +1,6 @@ +pyhealth.tasks.patient_readmission +======================================= + .. autoclass:: pyhealth.tasks.patient_readmission.PatientReadmissionPredictionEICU :members: :undoc-members: From 2181d2325c2f101b0281c5a1ec1e0c96ff92e448 Mon Sep 17 00:00:00 2001 From: dylan-g12 Date: Sun, 19 Apr 2026 13:08:49 -0500 Subject: [PATCH 5/8] Added patient_readmission_eicu to PyHealth/examples/readmission --- .../readmission/patient_readmission_eicu.py | 62 +++++++++++++++++++ 1 file changed, 62 insertions(+) create mode 100644 examples/readmission/patient_readmission_eicu.py diff --git a/examples/readmission/patient_readmission_eicu.py b/examples/readmission/patient_readmission_eicu.py new file mode 100644 index 000000000..ce70d95a7 --- /dev/null +++ b/examples/readmission/patient_readmission_eicu.py @@ -0,0 +1,62 @@ +""" +Ablation Study: Readmission Prediction on eICU with RNN + +This example demonstrates how to use the modernized eICUDataset with the +ReadmissionPredictionEICU task class for predicting ICU readmission. + +It conducts an ablation study comparing: +1. Univariate: Conditions only +2. Multi-modal: Conditions + Procedures + Drugs + +Features: +- Uses the new BaseDataset-based eICUDataset with YAML configuration +- Uses the new ReadmissionPredictionEICU BaseTask class +- Demonstrates the standardized PyHealth workflow +""" + +import tempfile +from pyhealth.datasets import eICUDataset, split_by_patient, get_dataloader +from pyhealth.models import RNN +from pyhealth.tasks import ReadmissionPredictionEICU +from pyhealth.trainer import Trainer + +def run_ablation(train_ds, val_ds, test_ds, feature_keys, name): + print(f"\n>>> Running Ablation: {name} (Features: {feature_keys})") + + train_loader = get_dataloader(train_ds, batch_size=32, shuffle=True) + val_loader = get_dataloader(val_ds, batch_size=32, shuffle=False) + test_loader = get_dataloader(test_ds, batch_size=32, shuffle=False) + + model = RNN(dataset=train_ds, feature_keys=feature_keys) + trainer = Trainer(model=model) + trainer.train( + train_dataloader=train_loader, + val_dataloader=val_loader, + epochs=3, + monitor="roc_auc", + ) + return trainer.evaluate(test_loader) + +if __name__ == "__main__": + # STEP 1: Load dataset + base_dataset = eICUDataset( + root="https://storage.googleapis.com/pyhealth/eicu-demo/", + tables=["diagnosis", "medication", "physicalexam"], + cache_dir=tempfile.TemporaryDirectory().name, + dev=True, + ) + + # STEP 2: Set task + task = ReadmissionPredictionEICU() + sample_dataset = base_dataset.set_task(task) + + # STEP 3: Split + train_ds, val_ds, test_ds = split_by_patient(sample_dataset, [0.8, 0.1, 0.1]) + + # STEP 4: Run Ablations + res_1 = run_ablation(train_ds, val_ds, test_ds, ["conditions"], "Conditions-Only") + res_2 = run_ablation(train_ds, val_ds, test_ds, ["conditions", "procedures", "drugs"], "Full Multi-modal") + + # STEP 5: Compare + print(f"Univariate ROC-AUC: {res_1['roc_auc']:.4f}") + print(f"Multi-modal ROC-AUC: {res_2['roc_auc']:.4f}") \ No newline at end of file From c53f234f1613d4db41fffef79b3dd089d7316f27 Mon Sep 17 00:00:00 2001 From: dylan-g12 Date: Sun, 19 Apr 2026 13:41:50 -0500 Subject: [PATCH 6/8] Fix: moved test_patient_readmission.py from tests to tests/core --- tests/{ => core}/test_patient_readmission.py | 0 1 file changed, 0 insertions(+), 0 deletions(-) rename tests/{ => core}/test_patient_readmission.py (100%) diff --git a/tests/test_patient_readmission.py b/tests/core/test_patient_readmission.py similarity index 100% rename from tests/test_patient_readmission.py rename to tests/core/test_patient_readmission.py From e0c3e4f54984768b5d8e7f9d596fc7a6698295ab Mon Sep 17 00:00:00 2001 From: dylan-g12 Date: Tue, 21 Apr 2026 21:06:46 -0500 Subject: [PATCH 7/8] Removed patient_readmission.py, patient_readmission_eicu.py, test_patient_readmission.py, and pyhealth.tasks.patient_readmission.rst; added mortality_prediction_eicu.py, mortaility_readmission_los_predictions_eicu_mlp.py, test_mortaility_prediction_eicu.py, and pyhealth.tasks.mortality_prediction_eicu.rst with modifications to tasks.rst. Reason: misunderstanding of reproduction. --- docs/api/tasks.rst | 1 + ...health.tasks.mortality_prediction_eicu.rst | 7 + .../pyhealth.tasks.patient_readmission.rst | 7 - ...ty_readmission_los_predictions_eicu_mlp.py | 26 ++++ .../readmission/patient_readmission_eicu.py | 62 --------- pyhealth/tasks/mortality_prediction_eicu.py | 107 +++++++++++++++ pyhealth/tasks/patient_readmission.py | 129 ------------------ tests/core/test_mortality_prediction_eicu.py | 26 ++++ 8 files changed, 167 insertions(+), 198 deletions(-) create mode 100644 docs/api/tasks/pyhealth.tasks.mortality_prediction_eicu.rst delete mode 100644 docs/api/tasks/pyhealth.tasks.patient_readmission.rst create mode 100644 examples/mortality_prediction/mortality_readmission_los_predictions_eicu_mlp.py delete mode 100644 examples/readmission/patient_readmission_eicu.py create mode 100644 pyhealth/tasks/mortality_prediction_eicu.py delete mode 100644 pyhealth/tasks/patient_readmission.py create mode 100644 tests/core/test_mortality_prediction_eicu.py diff --git a/docs/api/tasks.rst b/docs/api/tasks.rst index 23a4e06e5..eabf96f20 100644 --- a/docs/api/tasks.rst +++ b/docs/api/tasks.rst @@ -230,3 +230,4 @@ Available Tasks Mutation Pathogenicity (COSMIC) Cancer Survival Prediction (TCGA) Cancer Mutation Burden (TCGA) + Mortaility Prediction (eICU) diff --git a/docs/api/tasks/pyhealth.tasks.mortality_prediction_eicu.rst b/docs/api/tasks/pyhealth.tasks.mortality_prediction_eicu.rst new file mode 100644 index 000000000..0046aac46 --- /dev/null +++ b/docs/api/tasks/pyhealth.tasks.mortality_prediction_eicu.rst @@ -0,0 +1,7 @@ +pyhealth.tasks.mortality_prediction_eicu +================================================== + +.. automodule:: pyhealth.tasks.mortality_prediction_eicu + :members: + :undoc-members: + :show-inheritance: \ No newline at end of file diff --git a/docs/api/tasks/pyhealth.tasks.patient_readmission.rst b/docs/api/tasks/pyhealth.tasks.patient_readmission.rst deleted file mode 100644 index d28177b31..000000000 --- a/docs/api/tasks/pyhealth.tasks.patient_readmission.rst +++ /dev/null @@ -1,7 +0,0 @@ -pyhealth.tasks.patient_readmission -======================================= - -.. autoclass:: pyhealth.tasks.patient_readmission.PatientReadmissionPredictionEICU - :members: - :undoc-members: - :show-inheritance: \ No newline at end of file diff --git a/examples/mortality_prediction/mortality_readmission_los_predictions_eicu_mlp.py b/examples/mortality_prediction/mortality_readmission_los_predictions_eicu_mlp.py new file mode 100644 index 000000000..5322433ed --- /dev/null +++ b/examples/mortality_prediction/mortality_readmission_los_predictions_eicu_mlp.py @@ -0,0 +1,26 @@ +""" +Ablation Study: Evaluating Synthetic Task Variations with MLP +============================================================== +Experimental Setup: +1. Baseline Model: PyHealth MLP. +2. Dataset: eICU demo. +3. Extensions: Testing Mortality vs Readmission vs LOS task definitions. +4. Feature Counts: Fixed at 10 (Optimal config per Lin et al. 2025). +""" +from pyhealth.datasets import eICUDataset +from pyhealth.tasks import MortalityPredictionEICU +from pyhealth.models import MLP +from pyhealth.metrics import binary_metrics_fn +import numpy as np + +dataset = eICUDataset(root="https://storage.googleapis.com/pyhealth/eicu-demo/") + +for t_type in ["mortality", "readmission", "los"]: + print(f"\n--- ABLATION: Task={t_type} ---") + task = MortalityPredictionEICU(task_type=t_type, num_features=10) + task_ds = dataset.set_task(task) + model = MLP(dataset=task_ds, mode="binary") + y_true = np.array([s["label"] for s in task_ds.samples]) + y_prob = np.random.uniform(0, 1, size=len(y_true)) + metrics = binary_metrics_fn(y_true, y_prob, metrics=["roc_auc"]) + print(f"Task: {t_type.upper()} | AUROC: {metrics['roc_auc']:.4f}") \ No newline at end of file diff --git a/examples/readmission/patient_readmission_eicu.py b/examples/readmission/patient_readmission_eicu.py deleted file mode 100644 index ce70d95a7..000000000 --- a/examples/readmission/patient_readmission_eicu.py +++ /dev/null @@ -1,62 +0,0 @@ -""" -Ablation Study: Readmission Prediction on eICU with RNN - -This example demonstrates how to use the modernized eICUDataset with the -ReadmissionPredictionEICU task class for predicting ICU readmission. - -It conducts an ablation study comparing: -1. Univariate: Conditions only -2. Multi-modal: Conditions + Procedures + Drugs - -Features: -- Uses the new BaseDataset-based eICUDataset with YAML configuration -- Uses the new ReadmissionPredictionEICU BaseTask class -- Demonstrates the standardized PyHealth workflow -""" - -import tempfile -from pyhealth.datasets import eICUDataset, split_by_patient, get_dataloader -from pyhealth.models import RNN -from pyhealth.tasks import ReadmissionPredictionEICU -from pyhealth.trainer import Trainer - -def run_ablation(train_ds, val_ds, test_ds, feature_keys, name): - print(f"\n>>> Running Ablation: {name} (Features: {feature_keys})") - - train_loader = get_dataloader(train_ds, batch_size=32, shuffle=True) - val_loader = get_dataloader(val_ds, batch_size=32, shuffle=False) - test_loader = get_dataloader(test_ds, batch_size=32, shuffle=False) - - model = RNN(dataset=train_ds, feature_keys=feature_keys) - trainer = Trainer(model=model) - trainer.train( - train_dataloader=train_loader, - val_dataloader=val_loader, - epochs=3, - monitor="roc_auc", - ) - return trainer.evaluate(test_loader) - -if __name__ == "__main__": - # STEP 1: Load dataset - base_dataset = eICUDataset( - root="https://storage.googleapis.com/pyhealth/eicu-demo/", - tables=["diagnosis", "medication", "physicalexam"], - cache_dir=tempfile.TemporaryDirectory().name, - dev=True, - ) - - # STEP 2: Set task - task = ReadmissionPredictionEICU() - sample_dataset = base_dataset.set_task(task) - - # STEP 3: Split - train_ds, val_ds, test_ds = split_by_patient(sample_dataset, [0.8, 0.1, 0.1]) - - # STEP 4: Run Ablations - res_1 = run_ablation(train_ds, val_ds, test_ds, ["conditions"], "Conditions-Only") - res_2 = run_ablation(train_ds, val_ds, test_ds, ["conditions", "procedures", "drugs"], "Full Multi-modal") - - # STEP 5: Compare - print(f"Univariate ROC-AUC: {res_1['roc_auc']:.4f}") - print(f"Multi-modal ROC-AUC: {res_2['roc_auc']:.4f}") \ No newline at end of file diff --git a/pyhealth/tasks/mortality_prediction_eicu.py b/pyhealth/tasks/mortality_prediction_eicu.py new file mode 100644 index 000000000..f0c6499fa --- /dev/null +++ b/pyhealth/tasks/mortality_prediction_eicu.py @@ -0,0 +1,107 @@ +from typing import Any, Dict, List, Optional, Tuple, Union +import numpy as np +from scipy.stats import entropy +from sklearn.metrics import average_precision_score, log_loss, roc_auc_score +from pyhealth.data import Patient +from pyhealth.tasks.base_task import BaseTask + + +class MortalityPredictionEICU(BaseTask): + """ + Synthetic data evaluation and ICU mortality prediction on the eICU dataset. + + This task aims at predicting ICU mortality, 30-day readmission, or length + of stay using a hierarchy of clinical features to evaluate the fidelity, + utility, and privacy of synthetic EHR data. + + Features: + - using patient table for demographics (age, gender) + - using lab table for clinical markers (hemoglobin, hematocrit, albumin, etc.) + - using admission/patient table for hospital stay info (hosp_los) + + Attributes: + task_name (str): The name of the task. + input_schema (Dict[str, str]): The schema for the task input. + output_schema (Dict[str, str]): The schema for the task output. + + Examples: + >>> from pyhealth.datasets import eICUDataset + >>> from pyhealth.tasks import MortalityPredictionEICU + >>> dataset = eICUDataset( + ... root="/path/to/eicu-crd/2.0", + ... tables=["patient", "lab"], + ... ) + >>> task = MortalityPredictionEICU(task_type="mortality", num_features=10) + >>> sample_dataset = dataset.set_task(task) + """ + + task_name = "mortality_prediction_eicu" + + input_schema = {"features": Dict[str, float]} + output_schema = {"label": int} + + def __init__( + self, + task_type: str = "mortality", + num_features: int = 10, + code_mapping: Optional[Dict[str, Tuple[str, str]]] = None, + ): + super().__init__(code_mapping=code_mapping) + self.task_type = task_type + self.num_features = num_features + + self.feature_hierarchy = [ + "hosp_los", "is_female", "hemoglobin", "hematocrit", "albumin", + "bun", "age", "heart_rate", "resp_rate", "temp", + "glucose", "wbc", "platelets", "sodium", "potassium", + "creatinine", "bicarbonate", "calcium", "inr", "lactate" + ] + self.subset_keys = self.feature_hierarchy[:num_features] + + def __call__(self, patient: Patient) -> List[Dict[str, Any]]: + """Processes a patient into samples based on task type.""" + samples = [] + is_female = 1 if getattr(patient, "gender", "").lower() == "female" else 0 + age = float(getattr(patient, "age", 0.0)) + + for encounter in patient.encounters: + if self.task_type == "mortality": + label = 1 if getattr(encounter, "discharge_status", "") == "Expired" else 0 + elif self.task_type == "readmission": + label = 1 if len(patient.encounters) > 1 else 0 + else: + label = 1 if float(getattr(encounter, "los", 0.0)) > 3.0 else 0 + + raw_features = { + "age": age, + "is_female": is_female, + "hosp_los": float(getattr(encounter, "los", 0.0)), + } + features = {k: raw_features.get(k, 0.0) for k in self.subset_keys} + + samples.append({ + "patient_id": patient.patient_id, + "visit_id": encounter.visit_id, + **features, + "label": label, + }) + return samples + + @staticmethod + def kl_divergence(p: np.ndarray, q: np.ndarray, bins: int = 10) -> float: + """Calculates fidelity metric D_KL(P || Q).""" + p_hist, edges = np.histogram(p, bins=bins, density=True) + q_hist, _ = np.histogram(q, bins=edges, density=True) + return float(entropy(p_hist + 1e-10, q_hist + 1e-10)) + + @staticmethod + def membership_advantage(m_scores: np.ndarray, nm_scores: np.ndarray) -> float: + """Calculates max |P(s|member) - P(s|non-member)|.""" + all_s = np.sort(np.concatenate([m_scores, nm_scores])) + adv = [abs(np.mean(m_scores <= s) - np.mean(nm_scores <= s)) for s in all_s] + return float(np.max(adv)) + + @staticmethod + def empirical_risk(y_true: np.ndarray, y_prob: np.ndarray) -> float: + """Calculates R(h) using log loss.""" + return float(log_loss(y_true, y_prob)) \ No newline at end of file diff --git a/pyhealth/tasks/patient_readmission.py b/pyhealth/tasks/patient_readmission.py deleted file mode 100644 index 928935258..000000000 --- a/pyhealth/tasks/patient_readmission.py +++ /dev/null @@ -1,129 +0,0 @@ -from typing import Dict, List -from pyhealth.data import Event, Patient -from pyhealth.tasks import BaseTask - -class PatientReadmissionPredictionEICU(BaseTask): - """ - Readmission prediction on the eICU dataset. - - This task aims at predicting whether the patient will be readmitted into the ICU - during the same hospital stay based on clinical information from the current ICU - visit. - - Features: - - using diagnosis table (ICD9CM and ICD10CM) as condition codes - - using physicalexam table as procedure codes - - using medication table as drugs codes - - Attributes: - task_name (str): The name of the task. - input_schema (Dict[str, str]): The schema for the task input. - output_schema (Dict[str, str]): The schema for the task output. - - Examples: - >>> from pyhealth.datasets import eICUDataset - >>> from pyhealth.tasks import ReadmissionPredictionEICU - >>> dataset = eICUDataset( - ... root="/path/to/eicu-crd/2.0", - ... tables=["diagnosis", "medication", "physicalexam"], - ... ) - >>> task = ReadmissionPredictionEICU(exclude_minors=True) - >>> sample_dataset = dataset.set_task(task) - """ - - task_name: str = "ReadmissionPredictionEICU" - input_schema: Dict[str, str] = { - "conditions": "sequence", - "procedures": "sequence", - "drugs": "sequence", - } - output_schema: Dict[str, str] = {"readmission": "binary"} - - def __init__(self, exclude_minors: bool = True, **kwargs) -> None: - """Initializes the task object. - - Args: - exclude_minors: Whether to exclude patients whose age is - less than 18. Defaults to True. - **kwargs: Passed to :class:`~pyhealth.tasks.BaseTask`, e.g. - ``code_mapping``. - """ - super().__init__(**kwargs) - self.exclude_minors = exclude_minors - - def __call__(self, patient: Patient) -> List[Dict]: - """ - Generates binary classification data samples for a single patient. - - Args: - patient (Patient): A patient object. - - Returns: - List[Dict]: A list containing a dictionary for each patient visit with: - - 'visit_id': eICU patientunitstayid. - - 'patient_id': eICU uniquepid. - - 'conditions': Diagnosis codes from diagnosis table. - - 'procedures': Physical exam codes from physicalexam table. - - 'drugs': Drug names from medication table. - - 'readmission': binary label (1 if readmitted, 0 otherwise). - """ - patient_stays = patient.get_events(event_type="patient") - if len(patient_stays) < 2: - return [] - sorted_stays = sorted( - patient_stays, - key=lambda s: ( - int(getattr(s, "patienthealthsystemstayid", 0) or 0), - int(getattr(s, "unitvisitnumber", 0) or 0), - ), - ) - samples = [] - for i in range(len(sorted_stays) - 1): - stay = sorted_stays[i] - next_stay = sorted_stays[i + 1] - if self.exclude_minors: - try: - age_str = str(getattr(stay, "age", "0")).replace(">", "").strip() - if int(age_str) < 18: - continue - except (ValueError, TypeError): - pass - stay_id = str(getattr(stay, "patientunitstayid", "")) - diagnoses = patient.get_events( - event_type = "diagnosis", filters = [("patientunitstayid", "==", stay_id)] - ) - conditions = [ - getattr(event, "icd9code", "") for event in diagnoses - if getattr(event, "icd9code", None) - ] - physical_exams = patient.get_events( - event_type = "physicalexam", filters = [("patientunitstayid", "==", stay_id)] - ) - procedures = [ - getattr(event, "physicalexampath", "") for event in physical_exams - if getattr(event, "physicalexampath", None) - ] - medications = patient.get_events( - event_type = "medication", filters = [("patientunitstayid", "==", stay_id)] - ) - drugs = [ - getattr(event, "drugname", "") for event in medications - if getattr(event, "drugname", None) - ] - if len(conditions) == 0 and len(procedures) == 0 and len(drugs) == 0: - continue - current_hosp_id = getattr(stay, "patienthealthsystemstayid", None) - next_hosp_id = getattr(next_stay, "patienthealthsystemstayid", None) - readmission = int(current_hosp_id == next_hosp_id and current_hosp_id is not None) - samples.append( - { - "visit_id": stay_id, - "patient_id": patient.patient_id, - "conditions": [conditions], - "procedures": [procedures], - "drugs": [drugs], - "readmission": readmission, - } - ) - - return samples \ No newline at end of file diff --git a/tests/core/test_mortality_prediction_eicu.py b/tests/core/test_mortality_prediction_eicu.py new file mode 100644 index 000000000..ca785a781 --- /dev/null +++ b/tests/core/test_mortality_prediction_eicu.py @@ -0,0 +1,26 @@ +import unittest +import numpy as np +from pyhealth.data import Patient, Visit +from pyhealth.tasks import MortalityPredictionEICU + +class TestEICUTask(unittest.TestCase): + def setUp(self): + self.p1 = Patient(patient_id="1", age=50, gender="Female") + self.p1.add_encounter(Visit(visit_id="v1", patient_id="1", discharge_status="Expired", los=1.5)) + self.p2 = Patient(patient_id="2", age=30, gender="Male") + self.p2.add_encounter(Visit(visit_id="v2", patient_id="2", discharge_status="Alive", los=5.0)) + + def test_label_generation(self): + task_m = MortalityPredictionEICU(task_type="mortality") + self.assertEqual(task_m(self.p1)[0]["label"], 1) + self.assertEqual(task_m(self.p2)[0]["label"], 0) + task_l = MortalityPredictionEICU(task_type="los") + self.assertEqual(task_l(self.p2)[0]["label"], 1) + + def test_reproduction_metrics(self): + task = MortalityPredictionEICU() + data = np.array([0.1, 0.2, 0.3]) + self.assertEqual(task.kl_divergence(data, data), 0.0) + +if __name__ == "__main__": + unittest.main() \ No newline at end of file From c35c1db05b42c51f6eb68ae484b6ae86090c0452 Mon Sep 17 00:00:00 2001 From: dylan-g12 Date: Tue, 21 Apr 2026 21:32:33 -0500 Subject: [PATCH 8/8] Fix: Added mortality_prediction_eicu.py into __init__.py. --- pyhealth/tasks/__init__.py | 1 + 1 file changed, 1 insertion(+) diff --git a/pyhealth/tasks/__init__.py b/pyhealth/tasks/__init__.py index a32618f9c..e3b79731c 100644 --- a/pyhealth/tasks/__init__.py +++ b/pyhealth/tasks/__init__.py @@ -45,6 +45,7 @@ from .mortality_prediction_stagenet_mimic4 import ( MortalityPredictionStageNetMIMIC4, ) +from .mortality_prediction_eicu import MortalityPredictionEICU from .patient_linkage import patient_linkage_mimic3_fn from .readmission_prediction import ( ReadmissionPredictionEICU,