diff --git a/docs/api/tasks/pyhealth.tasks.drug_recommendation.rst b/docs/api/tasks/pyhealth.tasks.drug_recommendation.rst index eaf23ded5..3cd20c4e8 100644 --- a/docs/api/tasks/pyhealth.tasks.drug_recommendation.rst +++ b/docs/api/tasks/pyhealth.tasks.drug_recommendation.rst @@ -34,9 +34,9 @@ Task Functions (Legacy) neither of which the current ``pyhealth.data.Patient``/``Visit`` classes provide (``Visit`` is now a deprecated no-op stub). As a result they cannot currently be run through ``BaseDataset.set_task()``. Prefer the - task classes above (``DrugRecommendationMIMIC3``/``MIMIC4``/``EICU``), - which use the current API and are actively maintained. + task classes above (``DrugRecommendationMIMIC3``/``MIMIC4``/``EICU``/ + ``OMOP``), which use the current API and are actively maintained. .. autofunction:: pyhealth.tasks.drug_recommendation.drug_recommendation_mimic3_fn .. autofunction:: pyhealth.tasks.drug_recommendation.drug_recommendation_mimic4_fn -.. autofunction:: pyhealth.tasks.drug_recommendation.drug_recommendation_omop_fn \ No newline at end of file +.. autofunction:: pyhealth.tasks.drug_recommendation.drug_recommendation_omop_fn diff --git a/examples/drug_recommendation/drug_recommendation_mimic4_gamenet.py b/examples/drug_recommendation/drug_recommendation_mimic4_gamenet.py index 04eed866e..7804db873 100644 --- a/examples/drug_recommendation/drug_recommendation_mimic4_gamenet.py +++ b/examples/drug_recommendation/drug_recommendation_mimic4_gamenet.py @@ -3,7 +3,7 @@ # import mimic4 dataset and drug recommendaton task from pyhealth.datasets import MIMIC4Dataset -from pyhealth.tasks import drug_recommendation_mimic4_fn +from pyhealth.tasks import DrugRecommendationMIMIC4 # import dataloader related functions from pyhealth.datasets.splitter import split_by_patient @@ -32,11 +32,10 @@ def prepare_drug_task_data(): print("info") mimicvi.info() - # NOTE: drug_recommendation_mimic4_fn is a legacy, pre-2.0 task function - # (expects an indexable Patient with Visit.get_code_list()) and is not - # compatible with the current BaseDataset.set_task(), which requires a - # BaseTask instance. Use DrugRecommendationMIMIC4() instead. - mimic4_sample = mimicvi.set_task(drug_recommendation_mimic4_fn) + # drug_recommendation_mimic4_fn is a pre-2.0 task function (it expects an + # indexable Patient with Visit.get_code_list) and cannot be passed to + # BaseDataset.set_task(), which requires a BaseTask instance. + mimic4_sample = mimicvi.set_task(DrugRecommendationMIMIC4()) print(mimic4_sample[0]) return mimic4_sample diff --git a/examples/drug_recommendation/drug_recommendation_omop.py b/examples/drug_recommendation/drug_recommendation_omop.py new file mode 100644 index 000000000..a2fe91af9 --- /dev/null +++ b/examples/drug_recommendation/drug_recommendation_omop.py @@ -0,0 +1,30 @@ +"""Drug recommendation on an OMOP CDM dataset. + +Run with a local OMOP CDM v5.3 export, e.g. the CMS SynPUF 1k sample. +""" + +from pyhealth.datasets import OMOPDataset, get_dataloader, split_by_patient +from pyhealth.tasks import DrugRecommendationOMOP + + +def main() -> None: + dataset = OMOPDataset( + root="/path/to/omop_cdm", + tables=[ + "condition_occurrence", + "procedure_occurrence", + "drug_exposure", + ], + ) + dataset.stats() + + samples = dataset.set_task(DrugRecommendationOMOP()) + print(samples[0]) + + train, _val, _test = split_by_patient(samples, [0.8, 0.1, 0.1]) + train_loader = get_dataloader(train, batch_size=32, shuffle=True) + print(next(iter(train_loader)).keys()) + + +if __name__ == "__main__": + main() diff --git a/pyhealth/tasks/__init__.py b/pyhealth/tasks/__init__.py index f28b8ca76..612bc2546 100644 --- a/pyhealth/tasks/__init__.py +++ b/pyhealth/tasks/__init__.py @@ -14,6 +14,8 @@ from .covid19_cxr_classification import COVID19CXRClassification from .deid_ner import DeIDNERTask from .dka import DKAPredictionMIMIC4, T1DDKAPredictionMIMIC4 +# New exports must use the redundant `X as X` form: this module has no +# __all__, and the PR lint gate flags F401 on newly added import lines. from .drug_recommendation import ( DrugRecommendationEICU, DrugRecommendationMIMIC3, diff --git a/pyhealth/tasks/drug_recommendation.py b/pyhealth/tasks/drug_recommendation.py index 3e4fea3bb..13c30d84d 100644 --- a/pyhealth/tasks/drug_recommendation.py +++ b/pyhealth/tasks/drug_recommendation.py @@ -1,5 +1,6 @@ from typing import Any, Dict, Iterable, List, Optional -from typing import ClassVar +from typing import ClassVar # keep off line 1: merging would re-flag pre-existing UP035/I001 +from collections import defaultdict import polars as pl @@ -648,23 +649,30 @@ def __call__(self, patient: Any) -> List[Dict[str, Any]]: class DrugRecommendationOMOP(BaseTask): """Task for drug recommendation using an OMOP CDM dataset. - Drug recommendation aims at recommending a set of drugs given the patient health - history (e.g., conditions and procedures). This task creates samples with - cumulative history, where each visit includes all previous visit information. + Drug recommendation aims at recommending a set of drugs given the patient + health history (e.g., conditions and procedures). This task creates one + sample per qualifying visit with cumulative history: ``conditions`` and + ``procedures`` include the current visit, while ``drugs_hist`` excludes it + so the prediction target never appears in its own history. Features key-value pairs: - using condition_occurrence table as condition codes - using procedure_occurrence table as procedure codes - using drug_exposure table as drug codes + Concept ids equal to ``0`` are dropped: in OMOP, ``0`` is the + "no matching concept" sentinel, not a real code. + Attributes: task_name (str): The name of the task. - input_schema (Dict[str, str]): The schema for input data: - - conditions: Nested list of condition concept ids (history + current) - - procedures: Nested list of procedure concept ids (history + current) - - drugs_hist: Nested list of drug concept ids from history (current - visit excluded) - output_schema (Dict[str, str]): The schema for output data: + input_schema (dict[str, str]): The schema for input data: + - conditions: Nested list of condition concept ids (history + + current visit) + - procedures: Nested list of procedure concept ids (history + + current visit) + - drugs_hist: Nested list of drug concept ids from history; the + current visit's slot is always empty + output_schema (dict[str, str]): The schema for output data: - drugs: List of drug concept ids to predict for current visit Examples: @@ -672,13 +680,19 @@ class DrugRecommendationOMOP(BaseTask): >>> from pyhealth.tasks import DrugRecommendationOMOP >>> dataset = OMOPDataset( ... root="/path/to/omop", - ... tables=["condition_occurrence", "procedure_occurrence", "drug_exposure"], + ... tables=[ + ... "condition_occurrence", + ... "procedure_occurrence", + ... "drug_exposure", + ... ], ... ) - >>> task = DrugRecommendationOMOP() - >>> sample_dataset = dataset.set_task(task) + >>> sample_dataset = dataset.set_task(DrugRecommendationOMOP()) """ task_name: str = "DrugRecommendationOMOP" + # ClassVar is required here: ruff's RUF012 flags mutable class attributes, + # and the PR lint gate checks added lines. The sibling tasks predate that + # gate, hence the local inconsistency. input_schema: ClassVar[dict[str, str]] = { "conditions": "nested_sequence", "procedures": "nested_sequence", @@ -686,71 +700,92 @@ class DrugRecommendationOMOP(BaseTask): } output_schema: ClassVar[dict[str, str]] = {"drugs": "multilabel"} + # (sample key, event type, concept id column) + _SOURCES: ClassVar[tuple[tuple[str, str, str], ...]] = ( + ("conditions", "condition_occurrence", "condition_concept_id"), + ("procedures", "procedure_occurrence", "procedure_concept_id"), + ("drugs", "drug_exposure", "drug_concept_id"), + ) + _NULLISH: ClassVar[frozenset[str]] = frozenset({"", "nan", "none", ""}) + + @classmethod + def _norm(cls, value: Any) -> str | None: + """Normalizes a raw column value to a stable string, or None. + + CSV sources are loaded as all-string with pyarrow + (``strings_can_be_null=False``), so a blank cell arrives as ``""`` + rather than ``None``; Parquet sources keep their native dtype. This + collapses both cases. + """ + if value is None: + return None + text = str(value).strip() + return None if text.lower() in cls._NULLISH else text + + @classmethod + def _concept_id(cls, value: Any) -> str | None: + """Normalizes a concept id, dropping OMOP's 0 = 'no matching concept'.""" + code = cls._norm(value) + return None if code == "0" else code + + def _codes_by_visit( + self, patient: Any, event_type: str, field: str + ) -> dict[str, list[str]]: + """Groups one table's concept ids by visit in a single pass. + + Avoids one ``get_events`` call per (visit, table), which is O(V*N). + """ + grouped: dict[str, list[str]] = defaultdict(list) + for event in patient.get_events(event_type=event_type): + visit_id = self._norm(getattr(event, "visit_occurrence_id", None)) + code = self._concept_id(getattr(event, field, None)) + if visit_id is None or code is None: + continue + grouped[visit_id].append(code) + return grouped + def __call__(self, patient: Any) -> list[dict[str, Any]]: - """Process a patient to create drug recommendation samples. + """Processes a patient into drug recommendation samples. - Creates one sample per visit (after first visit) with cumulative history. - Each sample includes all previous visits' conditions, procedures, and drugs. + Emits one sample per visit that has at least one condition, one + procedure and one drug. Patients with fewer than two such visits are + dropped. Visits are consumed in chronological order (``Patient`` + sorts its event source by timestamp). Args: - patient: Patient object with get_events method + patient: Patient object exposing ``get_events``. Returns: - List of samples, each with patient_id, visit_id, conditions history, - procedures history, drugs history, and target drugs + List of samples with patient_id, visit_id, cumulative conditions + and procedures, leak-free drugs history, and the target drugs. """ - samples = [] - - # Get all visit occurrences - visit_occurrences = patient.get_events(event_type="visit_occurrence") - if len(visit_occurrences) < 2: - # Need at least 2 visits for history-based prediction + visits = patient.get_events(event_type="visit_occurrence") + if len(visits) < 2: return [] - # Process each visit - for visit in visit_occurrences: - condition_events = patient.get_events( - event_type="condition_occurrence", - filters=[("visit_occurrence_id", "==", visit.visit_occurrence_id)], - ) - conditions = [ - str(event.condition_concept_id) - for event in condition_events - if getattr(event, "condition_concept_id", None) is not None - ] - - procedure_events = patient.get_events( - event_type="procedure_occurrence", - filters=[("visit_occurrence_id", "==", visit.visit_occurrence_id)], - ) - procedures = [ - str(event.procedure_concept_id) - for event in procedure_events - if getattr(event, "procedure_concept_id", None) is not None - ] - - drug_events = patient.get_events( - event_type="drug_exposure", - filters=[("visit_occurrence_id", "==", visit.visit_occurrence_id)], - ) - drugs = [ - str(event.drug_concept_id) - for event in drug_events - if getattr(event, "drug_concept_id", None) is not None - ] + grouped = { + key: self._codes_by_visit(patient, event_type, field) + for key, event_type, field in self._SOURCES + } + samples: list[dict[str, Any]] = [] + for visit in visits: + visit_id = self._norm(getattr(visit, "visit_occurrence_id", None)) + if visit_id is None: + continue + conditions = grouped["conditions"].get(visit_id, []) + procedures = grouped["procedures"].get(visit_id, []) + drugs = grouped["drugs"].get(visit_id, []) # Exclude visits without condition, procedure, or drug code - if len(conditions) * len(procedures) * len(drugs) == 0: + if not (conditions and procedures and drugs): continue - samples.append( { - "visit_id": visit.visit_occurrence_id, + "visit_id": visit_id, "patient_id": patient.patient_id, "conditions": conditions, "procedures": procedures, "drugs": drugs, - "drugs_hist": drugs, } ) @@ -758,26 +793,19 @@ def __call__(self, patient: Any) -> list[dict[str, Any]]: if len(samples) < 2: return [] - # Add cumulative history for first sample - samples[0]["conditions"] = [samples[0]["conditions"]] - samples[0]["procedures"] = [samples[0]["procedures"]] - samples[0]["drugs_hist"] = [samples[0]["drugs_hist"]] - - # Add cumulative history for subsequent samples - for i in range(1, len(samples)): - samples[i]["conditions"] = samples[i - 1]["conditions"] + [ - samples[i]["conditions"] - ] - samples[i]["procedures"] = samples[i - 1]["procedures"] + [ - samples[i]["procedures"] - ] - samples[i]["drugs_hist"] = samples[i - 1]["drugs_hist"] + [ - samples[i]["drugs_hist"] - ] - - # Remove target drug from history (set current visit drugs_hist to empty) - for i in range(len(samples)): - samples[i]["drugs_hist"][i] = [] + # Snapshot before rewriting, then rebuild each sample from fresh lists + # so that no two samples ever share a list object. + per_visit = [ + (list(s["conditions"]), list(s["procedures"]), list(s["drugs"])) + for s in samples + ] + for index, sample in enumerate(samples): + window = per_visit[: index + 1] + sample["conditions"] = [list(codes) for codes, _, _ in window] + sample["procedures"] = [list(codes) for _, codes, _ in window] + sample["drugs_hist"] = [list(codes) for _, _, codes in window] + # The target visit's own drugs must not appear in its own history. + sample["drugs_hist"][index] = [] return samples diff --git a/tests/core/test_drug_recommendation_omop.py b/tests/core/test_drug_recommendation_omop.py new file mode 100644 index 000000000..f75a484b7 --- /dev/null +++ b/tests/core/test_drug_recommendation_omop.py @@ -0,0 +1,305 @@ +import csv +import sys +import tempfile +import unittest +from pathlib import Path + +from pyhealth.datasets import OMOPDataset +from pyhealth.tasks import DrugRecommendationOMOP + +TABLES = ["condition_occurrence", "procedure_occurrence", "drug_exposure"] + +PERSON_COLS = [ + "person_id", "gender_concept_id", "year_of_birth", "month_of_birth", + "day_of_birth", "race_concept_id", "ethnicity_concept_id", +] +VISIT_COLS = [ + "visit_occurrence_id", "person_id", "visit_concept_id", "visit_start_date", + "visit_start_datetime", "visit_end_date", "visit_end_datetime", + "visit_type_concept_id", +] +DEATH_COLS = ["person_id", "death_date", "death_datetime", "death_type_concept_id"] +CONDITION_COLS = [ + "person_id", "visit_occurrence_id", "condition_concept_id", + "condition_start_date", "condition_start_datetime", "condition_end_date", + "condition_end_datetime", "condition_type_concept_id", +] +PROCEDURE_COLS = [ + "person_id", "visit_occurrence_id", "procedure_concept_id", + "procedure_date", "procedure_datetime", "procedure_type_concept_id", +] +DRUG_COLS = [ + "person_id", "visit_occurrence_id", "drug_concept_id", + "drug_exposure_start_date", "drug_exposure_start_datetime", + "drug_exposure_end_date", "drug_exposure_end_datetime", + "drug_type_concept_id", +] + +# (person_id, visit_id, "YYYY-MM-DD", conditions, procedures, drugs) +VISITS = [ + # P_ORDER: three fully-coded visits, all codes distinct. + ("P_ORDER", "V13", "2020-03-01", ["C13"], ["P13"], ["D13"]), + ("P_ORDER", "V11", "2020-01-01", ["C11"], ["P11"], ["D11"]), + ("P_ORDER", "V12", "2020-02-01", ["C12"], ["P12"], ["D12"]), + # P_SKIP: middle visit has no drug -> dropped from the sample list. + ("P_SKIP", "V21", "2021-01-01", ["C21"], ["P21"], ["D21"]), + ("P_SKIP", "V22", "2021-02-01", ["C22"], ["P22"], []), + ("P_SKIP", "V23", "2021-03-01", ["C23"], ["P23"], ["D23"]), + # P_SINGLE: only one qualifying visit -> patient dropped entirely. + ("P_SINGLE", "V31", "2022-01-01", ["C31"], ["P31"], ["D31"]), + ("P_SINGLE", "V32", "2022-02-01", [], [], []), + # P_DIRTY: blank codes and OMOP's 0 sentinel must be filtered out. + ("P_DIRTY", "V41", "2023-01-01", ["C41", ""], ["P41"], ["D41", "0"]), + ("P_DIRTY", "V42", "2023-02-01", ["C42"], ["P42"], ["D42"]), +] + + +def _write(path: Path, columns, rows) -> None: + with path.open("w", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=columns) + writer.writeheader() + for row in rows: + writer.writerow({column: row.get(column, "") for column in columns}) + + +def write_omop_fixture(root: Path) -> None: + """Writes a minimal, fully-coded OMOP CDM v5.3 fixture under `root`.""" + root.mkdir(parents=True, exist_ok=True) + + persons = sorted({person for person, *_ in VISITS}) + _write(root / "person.csv", PERSON_COLS, [ + { + "person_id": person, "gender_concept_id": "8507", + "year_of_birth": "1970", "month_of_birth": "01", + "day_of_birth": "01", "race_concept_id": "0", + "ethnicity_concept_id": "0", + } + for person in persons + ]) + _write(root / "death.csv", DEATH_COLS, []) + + visits, conditions, procedures, drugs = [], [], [], [] + for person, visit, day, cond_codes, proc_codes, drug_codes in VISITS: + stamp = f"{day} 12:00:00" + visits.append({ + "visit_occurrence_id": visit, "person_id": person, + "visit_concept_id": "9201", "visit_start_date": day, + "visit_start_datetime": stamp, "visit_end_date": day, + "visit_end_datetime": stamp, "visit_type_concept_id": "32817", + }) + for code in cond_codes: + conditions.append({ + "person_id": person, "visit_occurrence_id": visit, + "condition_concept_id": code, "condition_start_date": day, + "condition_start_datetime": stamp, "condition_end_date": day, + "condition_end_datetime": stamp, + "condition_type_concept_id": "32020", + }) + for code in proc_codes: + procedures.append({ + "person_id": person, "visit_occurrence_id": visit, + "procedure_concept_id": code, "procedure_date": day, + "procedure_datetime": stamp, + "procedure_type_concept_id": "32020", + }) + for code in drug_codes: + drugs.append({ + "person_id": person, "visit_occurrence_id": visit, + "drug_concept_id": code, + "drug_exposure_start_date": day, + "drug_exposure_start_datetime": stamp, + "drug_exposure_end_date": day, + "drug_exposure_end_datetime": stamp, + "drug_type_concept_id": "32020", + }) + + # An orphan drug row: no visit_occurrence_id, must be ignored silently. + drugs.append({ + "person_id": "P_DIRTY", "visit_occurrence_id": "", + "drug_concept_id": "D_ORPHAN", + "drug_exposure_start_date": "2023-01-01", + "drug_exposure_start_datetime": "2023-01-01 12:00:00", + "drug_exposure_end_date": "2023-01-01", + "drug_exposure_end_datetime": "2023-01-01 12:00:00", + "drug_type_concept_id": "32020", + }) + + _write(root / "visit_occurrence.csv", VISIT_COLS, visits) + _write(root / "condition_occurrence.csv", CONDITION_COLS, conditions) + _write(root / "procedure_occurrence.csv", PROCEDURE_COLS, procedures) + _write(root / "drug_exposure.csv", DRUG_COLS, drugs) + + +class _OMOPFixtureCase(unittest.TestCase): + """Builds the fixture and the dataset once, in throwaway directories.""" + + @classmethod + def setUpClass(cls): + # litdata keeps its chunk files memory-mapped; on Windows the handles + # are still open when the directory is removed. Elsewhere a cleanup + # failure is a real defect and must surface. + windows = sys.platform == "win32" + cls._root_dir = tempfile.TemporaryDirectory( + ignore_cleanup_errors=windows + ) + cls._cache_dir = tempfile.TemporaryDirectory( + ignore_cleanup_errors=windows + ) + write_omop_fixture(Path(cls._root_dir.name)) + cls.dataset = OMOPDataset( + root=cls._root_dir.name, + tables=TABLES, + cache_dir=cls._cache_dir.name, + ) + + @classmethod + def tearDownClass(cls): + cls._cache_dir.cleanup() + cls._root_dir.cleanup() + + def samples_for(self, person_id: str): + return DrugRecommendationOMOP()(self.dataset.get_patient(person_id)) + + +class TestDrugRecommendationOMOPUnit(_OMOPFixtureCase): + """Unit-level behaviour of DrugRecommendationOMOP.__call__.""" + + def test_one_sample_per_qualifying_visit(self): + samples = self.samples_for("P_ORDER") + self.assertEqual(len(samples), 3) + + def test_visits_are_chronological_not_file_order(self): + # P_ORDER's rows are written V13, V11, V12 in the CSV. + samples = self.samples_for("P_ORDER") + self.assertEqual([s["visit_id"] for s in samples], ["V11", "V12", "V13"]) + + def test_history_excludes_current_visit(self): + for sample_index, sample in enumerate(self.samples_for("P_ORDER")): + with self.subTest(visit=sample["visit_id"]): + self.assertEqual(sample["drugs_hist"][sample_index], []) + + def test_history_holds_the_right_visit_in_the_right_slot(self): + samples = self.samples_for("P_ORDER") + self.assertEqual(samples[0]["drugs_hist"], [[]]) + self.assertEqual(samples[1]["drugs_hist"], [["D11"], []]) + self.assertEqual(samples[2]["drugs_hist"], [["D11"], ["D12"], []]) + + def test_conditions_and_procedures_include_current_visit(self): + samples = self.samples_for("P_ORDER") + self.assertEqual(samples[2]["conditions"], [["C11"], ["C12"], ["C13"]]) + self.assertEqual(samples[2]["procedures"], [["P11"], ["P12"], ["P13"]]) + + def test_target_is_the_current_visit_drugs(self): + samples = self.samples_for("P_ORDER") + self.assertEqual([s["drugs"] for s in samples], [["D11"], ["D12"], ["D13"]]) + + def test_incomplete_visit_is_dropped_from_history(self): + samples = self.samples_for("P_SKIP") + self.assertEqual([s["visit_id"] for s in samples], ["V21", "V23"]) + # V22 had no drug: it must not occupy a history slot. + self.assertEqual(samples[1]["drugs_hist"], [["D21"], []]) + self.assertNotIn(["C22"], samples[1]["conditions"]) + + def test_patient_with_one_qualifying_visit_is_dropped(self): + self.assertEqual(self.samples_for("P_SINGLE"), []) + + def test_blank_and_zero_concept_ids_are_filtered(self): + samples = self.samples_for("P_DIRTY") + self.assertEqual(samples[0]["conditions"][0], ["C41"]) + self.assertEqual(samples[0]["drugs"], ["D41"]) + for sample in samples: + for slot in sample["drugs_hist"]: + self.assertNotIn("0", slot) + self.assertNotIn("", slot) + + def test_orphan_event_without_visit_id_is_ignored(self): + for sample in self.samples_for("P_DIRTY"): + self.assertNotIn("D_ORPHAN", sample["drugs"]) + for slot in sample["drugs_hist"]: + self.assertNotIn("D_ORPHAN", slot) + + def test_samples_do_not_share_list_objects(self): + # Regression guard: history slots used to alias the target lists, + # so mutating one sample silently rewrote another. + samples = self.samples_for("P_ORDER") + self.assertIsNot(samples[1]["drugs_hist"][0], samples[0]["drugs"]) + samples[0]["drugs"].append("MUTATED") + self.assertEqual(samples[1]["drugs_hist"][0], ["D11"]) + self.assertEqual(samples[2]["drugs_hist"][0], ["D11"]) + + def test_visit_ids_are_strings(self): + for sample in self.samples_for("P_ORDER"): + self.assertIsInstance(sample["visit_id"], str) + + +class TestDrugRecommendationOMOPSchema(unittest.TestCase): + """Declared schema — the contract set_task() relies on.""" + + def test_schema_is_declared_on_the_class(self): + self.assertIn("task_name", vars(DrugRecommendationOMOP)) + self.assertIn("input_schema", vars(DrugRecommendationOMOP)) + self.assertIn("output_schema", vars(DrugRecommendationOMOP)) + self.assertEqual( + "DrugRecommendationOMOP", DrugRecommendationOMOP.task_name + ) + + def test_schema_matches_the_sibling_tasks(self): + self.assertEqual( + DrugRecommendationOMOP.input_schema, + { + "conditions": "nested_sequence", + "procedures": "nested_sequence", + "drugs_hist": "nested_sequence", + }, + ) + self.assertEqual( + DrugRecommendationOMOP.output_schema, {"drugs": "multilabel"} + ) + + def test_code_mapping_does_not_mutate_the_class(self): + task = DrugRecommendationOMOP( + code_mapping={"conditions": ("ICD9CM", "CCSCM")} + ) + self.assertIsInstance(task.input_schema["conditions"], tuple) + self.assertEqual( + DrugRecommendationOMOP.input_schema["conditions"], "nested_sequence" + ) + + +class TestDrugRecommendationOMOPIntegration(_OMOPFixtureCase): + """The claim the PR actually needs to prove: it runs through set_task().""" + + @classmethod + def setUpClass(cls): + super().setUpClass() + cls.sample_dataset = cls.dataset.set_task(DrugRecommendationOMOP()) + + @classmethod + def tearDownClass(cls): + try: + cls.sample_dataset.close() + except OSError: + if sys.platform != "win32": + raise + super().tearDownClass() + + def test_set_task_produces_the_expected_number_of_samples(self): + expected = sum( + len(DrugRecommendationOMOP()(self.dataset.get_patient(person))) + for person in ("P_ORDER", "P_SKIP", "P_SINGLE", "P_DIRTY") + ) + self.assertEqual(len(self.sample_dataset), expected) + self.assertGreater(expected, 0) + + def test_processed_samples_expose_the_declared_keys(self): + for sample in self.sample_dataset: + self.assertIn("patient_id", sample) + self.assertIn("visit_id", sample) + self.assertIn("conditions", sample) + self.assertIn("procedures", sample) + self.assertIn("drugs_hist", sample) + self.assertIn("drugs", sample) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/core/test_drug_recommendation_omop_leakage.py b/tests/core/test_drug_recommendation_omop_leakage.py deleted file mode 100644 index 34ac6c0ce..000000000 --- a/tests/core/test_drug_recommendation_omop_leakage.py +++ /dev/null @@ -1,143 +0,0 @@ -import unittest -from pathlib import Path - -from pyhealth.datasets import OMOPDataset -from pyhealth.tasks import DrugRecommendationOMOP -from pyhealth.tasks.drug_recommendation import drug_recommendation_omop_fn - - -class _MockVisit: - """Minimal stand-in for the legacy pyhealth.data.Visit interface that - drug_recommendation_omop_fn expects (visit_id, get_code_list(table)). - """ - - def __init__(self, visit_id, codes): - self.visit_id = visit_id - self._codes = codes - - def get_code_list(self, table): - return self._codes[table] - - -class _MockPatient: - """Minimal stand-in for the legacy indexable/len-able Patient interface - that drug_recommendation_omop_fn expects (patient_id, len(), [i]). - """ - - def __init__(self, patient_id, visits): - self.patient_id = patient_id - self._visits = visits - - def __len__(self): - return len(self._visits) - - def __getitem__(self, i): - return self._visits[i] - - -class TestDrugRecommendationOMOPLeakage(unittest.TestCase): - """Regression test for the drugs_all self-leakage bug. - - drug_recommendation_omop_fn built a "drugs_all" history feature by - accumulating each visit's own drugs without ever excluding the current - visit -- unlike its class-based siblings (DrugRecommendationMIMIC3/4/ - EICU) and the drug_recommendation_mimic3_fn/mimic4_fn functions, all of - which zero out the current visit's slot in the history sequence. That - meant the last entry of "drugs_all" was identical to the "drugs" target - for every sample, so a model could trivially copy it instead of - predicting from history. - """ - - def setUp(self): - self.visits = [ - _MockVisit( - "v1", - { - "condition_occurrence": ["C1"], - "procedure_occurrence": ["P1"], - "drug_exposure": ["D1"], - }, - ), - _MockVisit( - "v2", - { - "condition_occurrence": ["C2"], - "procedure_occurrence": ["P2"], - "drug_exposure": ["D2"], - }, - ), - _MockVisit( - "v3", - { - "condition_occurrence": ["C3"], - "procedure_occurrence": ["P3"], - "drug_exposure": ["D3"], - }, - ), - ] - self.patient = _MockPatient("pt1", self.visits) - - def test_drugs_all_excludes_current_visit_drugs(self): - samples = drug_recommendation_omop_fn(self.patient) - self.assertEqual(len(samples), 3) - - for i, sample in enumerate(samples): - with self.subTest(visit=sample["visit_id"]): - self.assertEqual( - sample["drugs_all"][i], - [], - "current visit's own drugs must not leak into its own " - "history slot", - ) - - def test_drugs_all_preserves_prior_visit_history(self): - samples = drug_recommendation_omop_fn(self.patient) - - # visit 2's history should still contain visit 1's drugs - self.assertEqual(samples[1]["drugs_all"][0], ["D1"]) - # visit 3's history should still contain visits 1 and 2's drugs - self.assertEqual(samples[2]["drugs_all"][0], ["D1"]) - self.assertEqual(samples[2]["drugs_all"][1], ["D2"]) - - def test_drugs_target_unaffected(self): - samples = drug_recommendation_omop_fn(self.patient) - self.assertEqual(samples[0]["drugs"], ["D1"]) - self.assertEqual(samples[1]["drugs"], ["D2"]) - self.assertEqual(samples[2]["drugs"], ["D3"]) - - -class TestDrugRecommendationOMOP(unittest.TestCase): - """DrugRecommendationOMOP is the current-API, leak-free replacement for - the legacy drug_recommendation_omop_fn (which cannot even run under the - current dataset API -- see the docs note on this task family). Verified - against real demo OMOP data. - """ - - @classmethod - def setUpClass(cls): - root = str(Path(__file__).parents[2] / "test-resources" / "omop") - tables = ["condition_occurrence", "procedure_occurrence", "drug_exposure"] - cls.dataset = OMOPDataset(root=root, tables=tables) - - def test_drugs_hist_excludes_current_visit_and_preserves_history(self): - # person_id "1" has 4 chronological visits (ids "1".."4"), each with - # exactly one condition/procedure/drug code (all coded "1"). - patient = self.dataset.get_patient("1") - samples = DrugRecommendationOMOP()(patient) - self.assertEqual(len(samples), 4) - - for i, sample in enumerate(samples): - with self.subTest(visit=sample["visit_id"]): - self.assertEqual(sample["drugs"], ["1"]) - self.assertEqual( - sample["drugs_hist"][i], - [], - "current visit's own drugs must not leak into its own " - "history slot", - ) - for j in range(i): - self.assertEqual(sample["drugs_hist"][j], ["1"]) - - -if __name__ == "__main__": - unittest.main()