Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion docs/api/tasks/pyhealth.tasks.mortality_prediction.rst
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,11 @@
:undoc-members:
:show-inheritance:

.. autoclass:: pyhealth.tasks.mortality_prediction.MultimodalMortalityPredictionMIMIC4
:members:
:undoc-members:
:show-inheritance:

.. autoclass:: pyhealth.tasks.mortality_prediction.MortalityPredictionEICU
:members:
:undoc-members:
Expand All @@ -24,4 +29,4 @@
.. autoclass:: pyhealth.tasks.mortality_prediction.MortalityPredictionOMOP
:members:
:undoc-members:
:show-inheritance:
:show-inheritance:
3 changes: 1 addition & 2 deletions examples/mortality_prediction/multimodal_mimic4_minimal.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,12 +24,11 @@
num_workers=8
)

# Apply multimodal task
# Apply multimodal task with X-rays available by prediction time
task = MultimodalMortalityPredictionMIMIC4()
samples = dataset.set_task(task, num_workers=8)

# Get and print sample
sample = samples[0]
print(sample)


21 changes: 18 additions & 3 deletions pyhealth/tasks/mortality_prediction.py
Original file line number Diff line number Diff line change
Expand Up @@ -351,7 +351,7 @@ class MultimodalMortalityPredictionMIMIC4(BaseTask):

Image Processing:
- Uses image_path from MIMIC-CXR metadata directly
- Returns first available X-ray image path across all X-rays
- Uses X-rays available by the last included admission's discharge
"""

task_name: str = "MultimodalMortalityPredictionMIMIC4"
Expand Down Expand Up @@ -575,6 +575,21 @@ def __call__(self, patient: Any) -> List[Dict[str, Any]]:
if len(admissions_to_process) == 0:
return []

prediction_time = None
for admission in reversed(admissions_to_process):
try:
admission_dischtime = datetime.strptime( # noqa: DTZ007
admission.dischtime, "%Y-%m-%d %H:%M:%S"
)
except (ValueError, AttributeError):
continue
if admission_dischtime >= admission.timestamp:
prediction_time = admission_dischtime
break

if prediction_time is None:
return []

# Get first admission time as reference for lab time calculations
first_admission_time = admissions_to_process[0].timestamp

Expand All @@ -591,8 +606,8 @@ def __call__(self, patient: Any) -> List[Dict[str, Any]]:

# Get X-ray data (patient-level, not admission-specific)
# Note: event types match table names in mimic4_cxr.yaml (negbio, metadata)
negbio_events = patient.get_events(event_type="negbio")
metadata_events = patient.get_events(event_type="metadata")
negbio_events = patient.get_events(event_type="negbio", end=prediction_time)
metadata_events = patient.get_events(event_type="metadata", end=prediction_time)

# Process X-ray findings (aggregate across all X-rays)
# NegBio findings attributes (from mimic4_cxr.yaml negbio table)
Expand Down
133 changes: 133 additions & 0 deletions tests/core/test_multimodal_mortality_leakage.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,133 @@
import unittest
from datetime import datetime, timedelta
from types import SimpleNamespace

import polars as pl

from pyhealth.tasks import MultimodalMortalityPredictionMIMIC4


class DummyPatient:
def __init__(self, events):
self.patient_id = "patient"
self.events = events

def get_events(
self,
event_type=None,
start=None,
end=None,
filters=None,
return_df=False,
):
events = list(self.events.get(event_type, []))
if start is not None:
events = [event for event in events if event.timestamp >= start]
if end is not None:
events = [event for event in events if event.timestamp <= end]
for field, operator, value in filters or []:
if operator == "==":
events = [
event for event in events if getattr(event, field, None) == value
]

if return_df:
return pl.DataFrame(
{
"timestamp": [event.timestamp for event in events],
"labevents/itemid": [event.itemid for event in events],
"labevents/valuenum": [event.valuenum for event in events],
"labevents/storetime": [event.storetime for event in events],
}
)
return events


def make_event(timestamp, **attributes):
return SimpleNamespace(timestamp=timestamp, **attributes)


def make_patient(*, include_available_xray):
admission_time = datetime(2025, 1, 1, 8) # noqa: DTZ001
cutoff = admission_time + timedelta(days=1)
death_admission_time = cutoff + timedelta(days=3)

metadata = [
make_event(
death_admission_time,
image_path="future.jpg",
)
]
negbio = [
make_event(
death_admission_time,
edema=1,
)
]
if include_available_xray:
metadata.insert(0, make_event(cutoff, image_path="available.jpg"))
negbio.insert(0, make_event(cutoff, cardiomegaly=1))

return DummyPatient(
{
"patients": [make_event(admission_time, anchor_age=50)],
"admissions": [
make_event(
admission_time,
hadm_id="history",
dischtime=cutoff.strftime("%Y-%m-%d %H:%M:%S"),
hospital_expire_flag=0,
),
make_event(
death_admission_time,
hadm_id="outcome",
dischtime=(death_admission_time + timedelta(days=1)).strftime(
"%Y-%m-%d %H:%M:%S"
),
hospital_expire_flag=1,
),
],
"diagnoses_icd": [
make_event(admission_time, hadm_id="history", icd_code="I10")
],
"procedures_icd": [
make_event(admission_time, hadm_id="history", icd_code="0W3P0ZZ")
],
"prescriptions": [
make_event(admission_time, hadm_id="history", ndc="0001")
],
"discharge": [
make_event(admission_time, hadm_id="history", text="Discharged")
],
"radiology": [],
"labevents": [
make_event(
admission_time + timedelta(hours=1),
itemid="50824",
valuenum=140.0,
storetime=(admission_time + timedelta(hours=1)).strftime(
"%Y-%m-%d %H:%M:%S"
),
)
],
"metadata": metadata,
"negbio": negbio,
}
)


class TestMultimodalMortalityLeakage(unittest.TestCase):
def test_excludes_future_xrays(self):
sample = MultimodalMortalityPredictionMIMIC4()(
make_patient(include_available_xray=True)
)[0]

self.assertEqual(sample["image_path"], "available.jpg")
self.assertEqual(sample["negbio_findings"], ["cardiomegaly"])

def test_requires_xray_by_prediction_time(self):
samples = MultimodalMortalityPredictionMIMIC4()(
make_patient(include_available_xray=False)
)

self.assertEqual(samples, [])
Loading