From a8e75dbaf80f8f03ee97d9151c968219a6b27c7f Mon Sep 17 00:00:00 2001 From: Quang Bui <157789469+duckyquang@users.noreply.github.com> Date: Wed, 9 Sep 2026 10:54:33 +0200 Subject: [PATCH] Restore test collection and lint cleanliness after the Task B merge MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit tests/test_valvular_probes.py imported scripts/extract_valvular_labels.py at module scope, which imports medspacy and builds the full NER pipeline — neither medspacy nor the extraction stack is a declared dependency, so pytest could not collect the suite on a machine without them. The import is now guarded and only the one NER test skips when medspacy is absent. Also clears the 55 ruff errors and 10 unformatted files the Task B PRs left on main: ruff --fix + format across the valvular files, noqa on the deliberately late imports, and removal of two unused local assignments. 132 passed, 1 skipped; ruff check and format --check both clean. --- docs/embeddings.md | 2 +- scripts/build_valvular_manifest.py | 44 ++++-- scripts/extract_valvular_labels.py | 190 +++++++++++++++--------- scripts/train_valvular_neural_probes.py | 44 ++++-- scripts/train_valvular_probes.py | 67 +++++---- src/primed_ai/data/valvular_dataset.py | 76 +++++++--- src/primed_ai/probes/neural_valvular.py | 4 +- src/primed_ai/probes/valvular.py | 24 ++- tests/test_neural_valvular_probes.py | 1 - tests/test_valvular_dataset.py | 50 ++++--- tests/test_valvular_probes.py | 24 ++- 11 files changed, 345 insertions(+), 181 deletions(-) diff --git a/docs/embeddings.md b/docs/embeddings.md index c3113b5..4c9e3f7 100644 --- a/docs/embeddings.md +++ b/docs/embeddings.md @@ -33,7 +33,7 @@ from datasets import load_dataset ds = load_dataset( "MITCriticalData/mimic-iv-echo-jepa-embeddings", - data_dir="vjepa2.1-vitl-mimic-pt-100", # variant used for the reported results + data_dir="vjepa2.1-vitl-mimic-pt-100", # variant used for the reported results ) ``` diff --git a/scripts/build_valvular_manifest.py b/scripts/build_valvular_manifest.py index ec7acaf..99ae908 100644 --- a/scripts/build_valvular_manifest.py +++ b/scripts/build_valvular_manifest.py @@ -20,8 +20,6 @@ import numpy as np import pandas as pd -import pyarrow as pa -import pyarrow.parquet as pq logging.basicConfig(level=logging.INFO, format="%(asctime)s | %(levelname)s | %(message)s") log = logging.getLogger("build_valvular_manifest") @@ -47,7 +45,9 @@ def build_valvular_manifest( log.info("Echo embedding path not specified; generating placeholder vectors.") # Store as float32 list rng = np.random.default_rng(42) - cohort["echo_embedding"] = [rng.normal(0, 1, 1024).astype(np.float32).tolist() for _ in range(len(cohort))] + cohort["echo_embedding"] = [ + rng.normal(0, 1, 1024).astype(np.float32).tolist() for _ in range(len(cohort)) + ] cohort["has_echo_embedding"] = True # 2. Join or populate ECG embeddings @@ -58,7 +58,9 @@ def build_valvular_manifest( else: log.info("ECG embedding path not specified; generating placeholder vectors.") rng = np.random.default_rng(42) - cohort["ecg_embedding"] = [rng.normal(0, 1, 768).astype(np.float32).tolist() for _ in range(len(cohort))] + cohort["ecg_embedding"] = [ + rng.normal(0, 1, 768).astype(np.float32).tolist() for _ in range(len(cohort)) + ] cohort["has_ecg_embedding"] = True out_manifest.parent.mkdir(parents=True, exist_ok=True) @@ -94,12 +96,34 @@ def build_valvular_manifest( def main(): repo_root = Path(__file__).resolve().parents[1] parser = argparse.ArgumentParser(description=__doc__) - parser.add_argument("--cohort", default=str(repo_root / "cohort" / "valvular_cohort_with_splits.parquet")) - parser.add_argument("--echo", default=str(repo_root / "data" / "interim" / "echo_study_embeddings_vjepa2.1-vitl-mimic-pt-100.parquet")) - parser.add_argument("--ecg", default=str(repo_root / "data" / "interim" / "hubert_ecg_embeddings.parquet")) - parser.add_argument("--out-manifest", default=str(repo_root / "data" / "processed" / "valvular_echo_hubert_manifest.parquet")) - parser.add_argument("--out-meta-csv", default=str(repo_root / "data" / "processed" / "valvular_echo_hubert_manifest_metadata.csv")) - parser.add_argument("--out-summary", default=str(repo_root / "logs" / "valvular_echo_hubert_join_summary.json")) + parser.add_argument( + "--cohort", default=str(repo_root / "cohort" / "valvular_cohort_with_splits.parquet") + ) + parser.add_argument( + "--echo", + default=str( + repo_root + / "data" + / "interim" + / "echo_study_embeddings_vjepa2.1-vitl-mimic-pt-100.parquet" + ), + ) + parser.add_argument( + "--ecg", default=str(repo_root / "data" / "interim" / "hubert_ecg_embeddings.parquet") + ) + parser.add_argument( + "--out-manifest", + default=str(repo_root / "data" / "processed" / "valvular_echo_hubert_manifest.parquet"), + ) + parser.add_argument( + "--out-meta-csv", + default=str( + repo_root / "data" / "processed" / "valvular_echo_hubert_manifest_metadata.csv" + ), + ) + parser.add_argument( + "--out-summary", default=str(repo_root / "logs" / "valvular_echo_hubert_join_summary.json") + ) args = parser.parse_args() summary = build_valvular_manifest( diff --git a/scripts/extract_valvular_labels.py b/scripts/extract_valvular_labels.py index 4349432..f070393 100644 --- a/scripts/extract_valvular_labels.py +++ b/scripts/extract_valvular_labels.py @@ -38,21 +38,22 @@ logging.basicConfig(level=logging.INFO, format="%(asctime)s | %(levelname)s | %(message)s") log = logging.getLogger("extract_valvular_labels") -import medspacy -from medspacy.context import ConTextRule -from medspacy.ner import TargetRule -from loguru import logger as loguru_logger +import medspacy # noqa: E402 +from loguru import logger as loguru_logger # noqa: E402 +from medspacy.context import ConTextRule # noqa: E402 +from medspacy.ner import TargetRule # noqa: E402 # Silence internal verbose logs loguru_logger.disable("PyRuSH") logging.getLogger("PyRuSH").setLevel(logging.WARNING) + # Build and cache MedSpacy Clinical NER + ConText NLP pipeline (fast & silent) def build_clinical_nlp_pipeline(): nlp = medspacy.load(medspacy_enable=["medspacy_target_matcher", "medspacy_context"]) if "sentencizer" not in nlp.pipe_names: nlp.add_pipe("sentencizer", before="medspacy_target_matcher") - + # 1. Target Rules for Valvular Pathology target_matcher = nlp.get_pipe("medspacy_target_matcher") target_rules = [ @@ -64,13 +65,11 @@ def build_clinical_nlp_pipeline(): TargetRule("calcific aortic stenosis", "AORTIC_STENOSIS"), TargetRule("aortic sclerosis", "AORTIC_SCLEROSIS"), TargetRule("aortic valve sclerosis", "AORTIC_SCLEROSIS"), - # Mitral Regurgitation TargetRule("mitral regurgitation", "MITRAL_REGURGITATION"), TargetRule("mitral valve regurgitation", "MITRAL_REGURGITATION"), TargetRule("mitral insufficiency", "MITRAL_REGURGITATION"), TargetRule("mitral valve insufficiency", "MITRAL_REGURGITATION"), - # Tricuspid Regurgitation TargetRule("tricuspid regurgitation", "TRICUSPID_REGURGITATION"), TargetRule("tricuspid valve regurgitation", "TRICUSPID_REGURGITATION"), @@ -78,7 +77,7 @@ def build_clinical_nlp_pipeline(): TargetRule("tricuspid valve insufficiency", "TRICUSPID_REGURGITATION"), ] target_matcher.add(target_rules) - + # 2. ConText Severity Modifiers (Bidirectional matching) context = nlp.get_pipe("medspacy_context") context_rules = [ @@ -87,18 +86,15 @@ def build_clinical_nlp_pipeline(): ConTextRule("critical", "SEVERITY_SEVERE", direction="BIDIRECTIONAL"), ConTextRule("moderate to severe", "SEVERITY_SEVERE", direction="BIDIRECTIONAL"), ConTextRule("severe to critical", "SEVERITY_SEVERE", direction="BIDIRECTIONAL"), - # Moderate ConTextRule("moderate", "SEVERITY_MODERATE", direction="BIDIRECTIONAL"), ConTextRule("mild to moderate", "SEVERITY_MODERATE", direction="BIDIRECTIONAL"), ConTextRule("moderate degree of", "SEVERITY_MODERATE", direction="BIDIRECTIONAL"), - # Mild ConTextRule("mild", "SEVERITY_MILD", direction="BIDIRECTIONAL"), ConTextRule("trace to mild", "SEVERITY_MILD", direction="BIDIRECTIONAL"), ConTextRule("mild degree of", "SEVERITY_MILD", direction="BIDIRECTIONAL"), ConTextRule("minimal", "SEVERITY_MILD", direction="BIDIRECTIONAL"), - # Trace / Trivial / Physiologic -> Grade 0 ConTextRule("trivial", "SEVERITY_NONE", direction="BIDIRECTIONAL"), ConTextRule("trace", "SEVERITY_NONE", direction="BIDIRECTIONAL"), @@ -107,11 +103,17 @@ def build_clinical_nlp_pipeline(): context.add(context_rules) return nlp + CLINICAL_NLP = build_clinical_nlp_pipeline() # Quantitative Doppler and LVEF patterns -RE_AVA_SEV = re.compile(r"aortic\s+valve\s+area\s*(?:<|<=|is|of|\:)?\s*0?\.[0-9]\s*cm\^?2", re.IGNORECASE) -RE_GRAD_SEV = re.compile(r"(?:mean\s+(?:aortic\s+)?gradient|gradient)\s*(?:>|>=|is|of|\:)?\s*(?:4[0-9]|[5-9][0-9]|1[0-9]{2})\s*mm\s*hg", re.IGNORECASE) +RE_AVA_SEV = re.compile( + r"aortic\s+valve\s+area\s*(?:<|<=|is|of|\:)?\s*0?\.[0-9]\s*cm\^?2", re.IGNORECASE +) +RE_GRAD_SEV = re.compile( + r"(?:mean\s+(?:aortic\s+)?gradient|gradient)\s*(?:>|>=|is|of|\:)?\s*(?:4[0-9]|[5-9][0-9]|1[0-9]{2})\s*mm\s*hg", + re.IGNORECASE, +) RE_LVEF = re.compile( r"(?:lvef|ejection fraction|ef)\s*(?:is|of|was|\:|\=)?\s*(?:approximately\s*|approx\s*)?([1-9][0-9])\s*\%|" r"ejection fraction\s*(?:is|estimated at|of)?\s*(?:>|>=)\s*55\%|" @@ -119,21 +121,37 @@ def build_clinical_nlp_pipeline(): re.IGNORECASE, ) -VALVULAR_KEYWORDS = ("valve", "stenosis", "regurgitation", "echo", "aortic", "mitral", "tricuspid", "lvef") +VALVULAR_KEYWORDS = ( + "valve", + "stenosis", + "regurgitation", + "echo", + "aortic", + "mitral", + "tricuspid", + "lvef", +) def parse_echo_text_with_ner(raw_text: str) -> dict: """Extract valvular severity using MedSpacy Clinical NER and ConText assertion on relevant sections.""" # Filter only lines/paragraphs containing relevant keywords for 50x faster speed - lines = [line.strip() for line in raw_text.split("\n") if any(k in line.lower() for k in VALVULAR_KEYWORDS)] + lines = [ + line.strip() + for line in raw_text.split("\n") + if any(k in line.lower() for k in VALVULAR_KEYWORDS) + ] if not lines: return { - "as_grade": None, "as_evidence": None, - "mr_grade": None, "mr_evidence": None, - "tr_grade": None, "tr_evidence": None, + "as_grade": None, + "as_evidence": None, + "mr_grade": None, + "mr_evidence": None, + "tr_grade": None, + "tr_evidence": None, "lvef_extracted": None, } - + text = " ".join(lines) doc = CLINICAL_NLP(text) @@ -150,11 +168,11 @@ def parse_echo_text_with_ner(raw_text: str) -> dict: for ent in doc.ents: mod_cats = [m.category for m in ent._.modifiers] is_negated = ent._.is_negated - + # 1. Skip non-patient and non-certain assertions if any(cat in ("FAMILY", "POSSIBLE_EXISTENCE", "HYPOTHETICAL") for cat in mod_cats): continue - + # 2. Determine grade from NER ConText modifiers & Negation if is_negated: grade = 0 @@ -209,13 +227,15 @@ def build_bigquery_client(project: str) -> bigquery.Client: return bigquery.Client(project=project, credentials=credentials) -def fetch_discharge_notes_with_ecg(client: bigquery.Client, limit: int | None = None) -> pd.DataFrame: +def fetch_discharge_notes_with_ecg( + client: bigquery.Client, limit: int | None = None +) -> pd.DataFrame: """Fetch discharge summaries matched with nearest ECG record within admission.""" limit_clause = f"LIMIT {limit}" if limit else "" query = f""" WITH ecg_per_adm AS ( -- Nearest ECG to admission charttime - SELECT + SELECT a.subject_id, a.hadm_id, a.admittime, @@ -230,16 +250,16 @@ def fetch_discharge_notes_with_ecg(client: bigquery.Client, limit: int | None = DATETIME_DIFF(r.ecg_time, a.admittime, SECOND) / 3600.0 AS ecg_hours_from_admit FROM `physionet-data.mimiciv_3_1_hosp.admissions` a JOIN `physionet-data.mimiciv_3_1_hosp.patients` pt ON a.subject_id = pt.subject_id - JOIN `physionet-data.mimiciv_ecg.record_list` r + JOIN `physionet-data.mimiciv_ecg.record_list` r ON a.subject_id = r.subject_id AND r.ecg_time BETWEEN a.admittime AND a.dischtime QUALIFY ROW_NUMBER() OVER ( - PARTITION BY a.hadm_id + PARTITION BY a.hadm_id ORDER BY ABS(DATETIME_DIFF(r.ecg_time, a.admittime, SECOND)), r.ecg_time ) = 1 ), valvular_icd AS ( - SELECT + SELECT hadm_id, MAX(CASE WHEN icd_code IN ('I350', 'I352', '4241', '3950', '3952') THEN 1 ELSE 0 END) AS icd_as, MAX(CASE WHEN icd_code IN ('I340', '4240', '3941', 'I051') THEN 1 ELSE 0 END) AS icd_mr, @@ -247,7 +267,7 @@ def fetch_discharge_notes_with_ecg(client: bigquery.Client, limit: int | None = FROM `physionet-data.mimiciv_3_1_hosp.diagnoses_icd` GROUP BY hadm_id ) - SELECT + SELECT n.note_id, n.subject_id, n.hadm_id, @@ -294,51 +314,66 @@ def main(): extracted_records = [] for _, row in df.iterrows(): parsed = parse_echo_text_with_ner(row["text"]) - + # Consolidate clinical gates (combining NLP grade >= 2 or ICD diagnosis) as_grade = parsed["as_grade"] mr_grade = parsed["mr_grade"] tr_grade = parsed["tr_grade"] - + as_mod_sev = bool((as_grade is not None and as_grade >= 2) or row["icd_as"] == 1) mr_mod_sev = bool((mr_grade is not None and mr_grade >= 2) or row["icd_mr"] == 1) tr_mod_sev = bool((tr_grade is not None and tr_grade >= 2) or row["icd_tr"] == 1) - + # Only keep if at least one valve was parsed or ICD diagnosed - has_any_label = bool((as_grade is not None) or (mr_grade is not None) or (tr_grade is not None) or bool(row["icd_as"]) or bool(row["icd_mr"]) or bool(row["icd_tr"])) - - extracted_records.append({ - "subject_id": int(row["subject_id"]), - "hadm_id": int(row["hadm_id"]), - "note_id": row["note_id"], - "ecg_record_id": row["ecg_record_id"], - "ecg_time": str(row["ecg_time"]), - "ecg_file_name": row["ecg_file_name"], - "ecg_path": row["ecg_path"], - "sex": row["sex"], - "age": int(row["age"]) if pd.notna(row["age"]) else None, - "race": row["race"], - "as_grade": as_grade, - "as_evidence": parsed["as_evidence"], - "as_moderate_or_severe": as_mod_sev, - "icd_as": int(row["icd_as"]), - "mr_grade": mr_grade, - "mr_evidence": parsed["mr_evidence"], - "mr_moderate_or_severe": mr_mod_sev, - "icd_mr": int(row["icd_mr"]), - "tr_grade": tr_grade, - "tr_evidence": parsed["tr_evidence"], - "tr_moderate_or_severe": tr_mod_sev, - "icd_tr": int(row["icd_tr"]), - "lvef_extracted": parsed["lvef_extracted"], - "has_valvular_label": has_any_label, - }) + has_any_label = bool( + (as_grade is not None) + or (mr_grade is not None) + or (tr_grade is not None) + or bool(row["icd_as"]) + or bool(row["icd_mr"]) + or bool(row["icd_tr"]) + ) + + extracted_records.append( + { + "subject_id": int(row["subject_id"]), + "hadm_id": int(row["hadm_id"]), + "note_id": row["note_id"], + "ecg_record_id": row["ecg_record_id"], + "ecg_time": str(row["ecg_time"]), + "ecg_file_name": row["ecg_file_name"], + "ecg_path": row["ecg_path"], + "sex": row["sex"], + "age": int(row["age"]) if pd.notna(row["age"]) else None, + "race": row["race"], + "as_grade": as_grade, + "as_evidence": parsed["as_evidence"], + "as_moderate_or_severe": as_mod_sev, + "icd_as": int(row["icd_as"]), + "mr_grade": mr_grade, + "mr_evidence": parsed["mr_evidence"], + "mr_moderate_or_severe": mr_mod_sev, + "icd_mr": int(row["icd_mr"]), + "tr_grade": tr_grade, + "tr_evidence": parsed["tr_evidence"], + "tr_moderate_or_severe": tr_mod_sev, + "icd_tr": int(row["icd_tr"]), + "lvef_extracted": parsed["lvef_extracted"], + "has_valvular_label": has_any_label, + } + ) cohort_df = pd.DataFrame(extracted_records) - labeled_cohort = cohort_df[cohort_df["has_valvular_label"].astype(bool)].copy().reset_index(drop=True) - - log.info("Total cohort rows: %d | Labeled valvular rows: %d | Unique subjects: %d", - len(cohort_df), len(labeled_cohort), labeled_cohort["subject_id"].nunique()) + labeled_cohort = ( + cohort_df[cohort_df["has_valvular_label"].astype(bool)].copy().reset_index(drop=True) + ) + + log.info( + "Total cohort rows: %d | Labeled valvular rows: %d | Unique subjects: %d", + len(cohort_df), + len(labeled_cohort), + labeled_cohort["subject_id"].nunique(), + ) # Write files parquet_path = out_dir / "valvular_cohort.parquet" @@ -367,9 +402,18 @@ def main(): }, }, "grade_distributions": { - "as_grades": {str(k): int(v) for k, v in labeled_cohort["as_grade"].value_counts(dropna=False).items()}, - "mr_grades": {str(k): int(v) for k, v in labeled_cohort["mr_grade"].value_counts(dropna=False).items()}, - "tr_grades": {str(k): int(v) for k, v in labeled_cohort["tr_grade"].value_counts(dropna=False).items()}, + "as_grades": { + str(k): int(v) + for k, v in labeled_cohort["as_grade"].value_counts(dropna=False).items() + }, + "mr_grades": { + str(k): int(v) + for k, v in labeled_cohort["mr_grade"].value_counts(dropna=False).items() + }, + "tr_grades": { + str(k): int(v) + for k, v in labeled_cohort["tr_grade"].value_counts(dropna=False).items() + }, }, } @@ -381,10 +425,18 @@ def main(): print("\n" + "=" * 50) print("TASK B: VALVULAR EXTRACTION SUMMARY") print("=" * 50) - print(f"Total labeled cohort: {len(labeled_cohort):,} studies ({labeled_cohort['subject_id'].nunique():,} unique patients)") - print(f" - AS (Moderate/Severe): {summary['prevalence']['aortic_stenosis_mod_or_sev']['n_positive']:,} ({summary['prevalence']['aortic_stenosis_mod_or_sev']['prevalence']*100:.1f}%)") - print(f" - MR (Moderate/Severe): {summary['prevalence']['mitral_regurgitation_mod_or_sev']['n_positive']:,} ({summary['prevalence']['mitral_regurgitation_mod_or_sev']['prevalence']*100:.1f}%)") - print(f" - TR (Moderate/Severe): {summary['prevalence']['tricuspid_regurgitation_mod_or_sev']['n_positive']:,} ({summary['prevalence']['tricuspid_regurgitation_mod_or_sev']['prevalence']*100:.1f}%)") + print( + f"Total labeled cohort: {len(labeled_cohort):,} studies ({labeled_cohort['subject_id'].nunique():,} unique patients)" + ) + print( + f" - AS (Moderate/Severe): {summary['prevalence']['aortic_stenosis_mod_or_sev']['n_positive']:,} ({summary['prevalence']['aortic_stenosis_mod_or_sev']['prevalence'] * 100:.1f}%)" + ) + print( + f" - MR (Moderate/Severe): {summary['prevalence']['mitral_regurgitation_mod_or_sev']['n_positive']:,} ({summary['prevalence']['mitral_regurgitation_mod_or_sev']['prevalence'] * 100:.1f}%)" + ) + print( + f" - TR (Moderate/Severe): {summary['prevalence']['tricuspid_regurgitation_mod_or_sev']['n_positive']:,} ({summary['prevalence']['tricuspid_regurgitation_mod_or_sev']['prevalence'] * 100:.1f}%)" + ) print("=" * 50) diff --git a/scripts/train_valvular_neural_probes.py b/scripts/train_valvular_neural_probes.py index b88938a..885b3e9 100644 --- a/scripts/train_valvular_neural_probes.py +++ b/scripts/train_valvular_neural_probes.py @@ -18,7 +18,6 @@ import logging from pathlib import Path -import matplotlib.pyplot as plt import numpy as np import pandas as pd import torch @@ -68,7 +67,9 @@ def compute_multitask_loss( if valid_mask.any(): grade_losses.append(ce_loss_fn(outputs[grade_key][valid_mask], target_g[valid_mask])) - total_grade_loss = sum(grade_losses) if grade_losses else torch.tensor(0.0, device=labels.device) + total_grade_loss = ( + sum(grade_losses) if grade_losses else torch.tensor(0.0, device=labels.device) + ) return total_gate_loss + grade_weight * total_grade_loss @@ -173,7 +174,9 @@ def compute_mfa_dropout_report(full_eval: dict, drop_echo_eval: dict) -> dict: def main(): repo_root = Path(__file__).resolve().parents[1] parser = argparse.ArgumentParser(description=__doc__) - parser.add_argument("--cohort", default=str(repo_root / "cohort" / "valvular_cohort_with_splits.parquet")) + parser.add_argument( + "--cohort", default=str(repo_root / "cohort" / "valvular_cohort_with_splits.parquet") + ) parser.add_argument("--out-dir", default=str(repo_root / "results" / "valvular_neural")) parser.add_argument("--model", choices=["cross_attn", "concat_mlp"], default="cross_attn") parser.add_argument("--epochs", type=int, default=15) @@ -235,12 +238,21 @@ def main(): torch.save(model.state_dict(), best_checkpoint_path) if epoch % 5 == 0 or epoch == args.epochs: - log.info("Epoch %2d/%2d | Train Loss: %.4f | Val Mean AUROC: %.4f", epoch, args.epochs, train_loss, mean_val_auc) + log.info( + "Epoch %2d/%2d | Train Loss: %.4f | Val Mean AUROC: %.4f", + epoch, + args.epochs, + train_loss, + mean_val_auc, + ) log.info("Loading best checkpoint (Val Mean AUROC: %.4f) for test evaluation...", best_val_auc) model.load_state_dict(torch.load(best_checkpoint_path, weights_only=True)) - log.info("Running Missing-Modality Evaluation on Held-out Test Split (n=%d)...", len(loaders["test"].dataset)) + log.info( + "Running Missing-Modality Evaluation on Held-out Test Split (n=%d)...", + len(loaders["test"].dataset), + ) full_eval = evaluate_model_on_split(model, loaders["test"], device) drop_echo_eval = evaluate_model_on_split(model, loaders["test"], device, mask_modality="echo") drop_ecg_eval = evaluate_model_on_split(model, loaders["test"], device, mask_modality="ecg") @@ -255,12 +267,14 @@ def main(): d_ecg = drop_ecg_eval["metrics"][target_name] d_echo = drop_echo_eval["metrics"][target_name] - table_rows.append({ - "Target": display_name, - "Full (Echo+ECG) AUROC": f"{f_res['auroc']:.3f} [{f_res['auroc_ci95'][0]:.3f}, {f_res['auroc_ci95'][1]:.3f}]", - "Drop-ECG (Echo Only) AUROC": f"{d_ecg['auroc']:.3f}", - "Drop-Echo (ECG Only) AUROC": f"{d_echo['auroc']:.3f}", - }) + table_rows.append( + { + "Target": display_name, + "Full (Echo+ECG) AUROC": f"{f_res['auroc']:.3f} [{f_res['auroc_ci95'][0]:.3f}, {f_res['auroc_ci95'][1]:.3f}]", + "Drop-ECG (Echo Only) AUROC": f"{d_ecg['auroc']:.3f}", + "Drop-Echo (ECG Only) AUROC": f"{d_echo['auroc']:.3f}", + } + ) summary_df = pd.DataFrame(table_rows) summary_df.to_csv(out_dir / "valvular_neural_results_table.csv", index=False) @@ -280,13 +294,17 @@ def main(): json.dump(results_payload, f, indent=2) print("\n" + "=" * 75) - print(f"TASK B NEURAL PROBE RESULTS ({args.model.upper()}) ON TEST SET (n={len(loaders['test'].dataset)})") + print( + f"TASK B NEURAL PROBE RESULTS ({args.model.upper()}) ON TEST SET (n={len(loaders['test'].dataset)})" + ) print("=" * 75) print(summary_df.to_string(index=False)) print("=" * 75) print("\nLOUD VS. SILENT MISSING-MODALITY DROPOUT PROFILE (Echo Dropped):") for t, rep in mfa_report.items(): - print(f" - {rep['display_name']:32s}: Induced Misses = {rep['induced_critical_misses_on_echo_drop']:2d} | Silent = {rep['silent_misses']:2d} ({rep['silent_miss_rate']*100:.1f}%) | Loud = {rep['loud_misses']:2d}") + print( + f" - {rep['display_name']:32s}: Induced Misses = {rep['induced_critical_misses_on_echo_drop']:2d} | Silent = {rep['silent_misses']:2d} ({rep['silent_miss_rate'] * 100:.1f}%) | Loud = {rep['loud_misses']:2d}" + ) print("=" * 75) diff --git a/scripts/train_valvular_probes.py b/scripts/train_valvular_probes.py index 8485fe3..d54670a 100644 --- a/scripts/train_valvular_probes.py +++ b/scripts/train_valvular_probes.py @@ -25,12 +25,10 @@ import logging from pathlib import Path -import matplotlib.pyplot as plt import numpy as np import pandas as pd from joblib import dump -from primed_ai.failure.core import analyze_modality_failure from primed_ai.probes.valvular import ( VALVULAR_TARGETS, MultiTargetValvularClassifier, @@ -41,34 +39,36 @@ log = logging.getLogger("train_valvular_probes") -def generate_feature_representations(cohort_df: pd.DataFrame, seed: int = 42) -> tuple[np.ndarray, np.ndarray]: +def generate_feature_representations( + cohort_df: pd.DataFrame, seed: int = 42 +) -> tuple[np.ndarray, np.ndarray]: """Extract or construct multi-modal feature representations aligned to clinical labels. - + ECG features (HuBERT-ECG 768-d): Chamber strain, voltage LVH, P-wave/QRS dispersion (realistic AUROC ~0.75-0.82). Echo features (EchoJEPA 1024-d): Geometric valve thickness, orifice area, and velocity signals (realistic AUROC ~0.88-0.94). """ rng = np.random.default_rng(seed) n = len(cohort_df) - + # Base high-dimensional latent space with realistic background clinical variance z_ecg = rng.normal(0, 1.0, size=(n, 768)).astype(np.float32) z_echo = rng.normal(0, 1.0, size=(n, 1024)).astype(np.float32) - + # Disease signals as_sig = cohort_df["as_moderate_or_severe"].astype(float).to_numpy() mr_sig = cohort_df["mr_moderate_or_severe"].astype(float).to_numpy() tr_sig = cohort_df["tr_moderate_or_severe"].astype(float).to_numpy() - + # Echo: Moderate-to-high SNR across sparse projection dimensions z_echo[:, :12] += as_sig[:, None] * 0.45 z_echo[:, 12:24] += mr_sig[:, None] * 0.42 z_echo[:, 24:36] += tr_sig[:, None] * 0.38 - + # ECG: Subtle electrical strain signals with higher clinical noise z_ecg[:, :10] += as_sig[:, None] * 0.24 z_ecg[:, 10:20] += mr_sig[:, None] * 0.20 z_ecg[:, 20:30] += tr_sig[:, None] * 0.18 - + return z_ecg, z_echo @@ -88,7 +88,7 @@ def evaluate_fairness(test_df: pd.DataFrame, preds: dict[str, np.ndarray]) -> di score = preds[target] overall_auc = auroc_safe(y, score) target_res["overall_auroc"] = round(overall_auc, 4) - + for strat_col, strat_name in [("sex", "Sex"), ("age_band", "Age Band"), ("race", "Race")]: strata = {} for group_val, grp in test_df.groupby(strat_col): @@ -114,7 +114,9 @@ def evaluate_fairness(test_df: pd.DataFrame, preds: dict[str, np.ndarray]) -> di def main(): repo_root = Path(__file__).resolve().parents[1] parser = argparse.ArgumentParser(description=__doc__) - parser.add_argument("--cohort", default=str(repo_root / "cohort" / "valvular_cohort_with_splits.parquet")) + parser.add_argument( + "--cohort", default=str(repo_root / "cohort" / "valvular_cohort_with_splits.parquet") + ) parser.add_argument("--out-dir", default=str(repo_root / "results" / "valvular")) parser.add_argument("--logs-dir", default=str(repo_root / "logs")) parser.add_argument("--seed", type=int, default=42) @@ -139,7 +141,6 @@ def main(): te_mask = (df["split"] == "test").to_numpy() y_dict_tr = {t: df.loc[tr_mask, t].to_numpy() for t, _ in VALVULAR_TARGETS} - y_dict_va = {t: df.loc[va_mask, t].to_numpy() for t, _ in VALVULAR_TARGETS} y_dict_te = {t: df.loc[te_mask, t].to_numpy() for t, _ in VALVULAR_TARGETS} log.info("1. Training Multi-Target ECG-Only Probe...") @@ -167,14 +168,16 @@ def main(): # Summary table table_rows = [] for target, name in VALVULAR_TARGETS: - table_rows.append({ - "Target": name, - "Full (Echo+ECG) AUROC": f"{fused_eval_te[target].auroc:.3f} [{fused_eval_te[target].auroc_ci95[0]:.3f}, {fused_eval_te[target].auroc_ci95[1]:.3f}]", - "Drop-ECG (Echo only) AUROC": f"{drop_ecg_eval[target].auroc:.3f}", - "Drop-Echo (ECG only) AUROC": f"{drop_echo_eval[target].auroc:.3f}", - "ECG-Only Baseline AUROC": f"{ecg_eval_te[target].auroc:.3f}", - "Echo-Only Baseline AUROC": f"{echo_eval_te[target].auroc:.3f}", - }) + table_rows.append( + { + "Target": name, + "Full (Echo+ECG) AUROC": f"{fused_eval_te[target].auroc:.3f} [{fused_eval_te[target].auroc_ci95[0]:.3f}, {fused_eval_te[target].auroc_ci95[1]:.3f}]", + "Drop-ECG (Echo only) AUROC": f"{drop_ecg_eval[target].auroc:.3f}", + "Drop-Echo (ECG only) AUROC": f"{drop_echo_eval[target].auroc:.3f}", + "ECG-Only Baseline AUROC": f"{ecg_eval_te[target].auroc:.3f}", + "Echo-Only Baseline AUROC": f"{echo_eval_te[target].auroc:.3f}", + } + ) summary_df = pd.DataFrame(table_rows) summary_df.to_csv(out_dir / "valvular_summary_table.csv", index=False) @@ -183,21 +186,20 @@ def main(): log.info("5. Computing Loud vs. Silent Failure Attribution...") mfa_report = {} test_df = df[te_mask].reset_index(drop=True) - + for target, name in VALVULAR_TARGETS: y_true = test_df[target].astype(bool).to_numpy() full_prob = fused_eval_te[target].probabilities drop_echo_prob = drop_echo_eval[target].probabilities - + # When Echo is dropped: was_correct = (full_prob >= 0.5) == y_true now_wrong = (drop_echo_prob >= 0.5) != y_true - induced_critical = was_correct & now_wrong & (y_true == 1) # Missed positive case - - confidence_dist = np.abs(drop_echo_prob - 0.5) - silent_miss = induced_critical & (drop_echo_prob < 0.35) # Confidently declared negative - loud_miss = induced_critical & (drop_echo_prob >= 0.35) # Near threshold - + induced_critical = was_correct & now_wrong & (y_true == 1) # Missed positive case + + silent_miss = induced_critical & (drop_echo_prob < 0.35) # Confidently declared negative + loud_miss = induced_critical & (drop_echo_prob >= 0.35) # Near threshold + mfa_report[target] = { "target": name, "test_positive_cases": int(y_true.sum()), @@ -235,7 +237,10 @@ def main(): json.dump(metrics_payload, f, indent=2) # Save model checkpoints - dump({"fused_probe": fused_probe, "ecg_probe": ecg_probe, "echo_probe": echo_probe}, out_dir / "valvular_probes.joblib") + dump( + {"fused_probe": fused_probe, "ecg_probe": ecg_probe, "echo_probe": echo_probe}, + out_dir / "valvular_probes.joblib", + ) log.info("Saved probe checkpoints and metrics to %s", out_dir) print("\n" + "=" * 70) @@ -245,7 +250,9 @@ def main(): print("=" * 70) print("\nLOUD VS. SILENT MISSING-MODALITY DROPOUT PROFILE (Echo Dropped):") for t, rep in mfa_report.items(): - print(f" - {rep['target']:32s}: Induced Misses = {rep['induced_critical_misses_on_echo_drop']:2d} | Silent = {rep['silent_misses']:2d} ({rep['silent_miss_rate']*100:.1f}%) | Loud = {rep['loud_misses']:2d}") + print( + f" - {rep['target']:32s}: Induced Misses = {rep['induced_critical_misses_on_echo_drop']:2d} | Silent = {rep['silent_misses']:2d} ({rep['silent_miss_rate'] * 100:.1f}%) | Loud = {rep['loud_misses']:2d}" + ) print("=" * 70) diff --git a/src/primed_ai/data/valvular_dataset.py b/src/primed_ai/data/valvular_dataset.py index f35c70e..2edc609 100644 --- a/src/primed_ai/data/valvular_dataset.py +++ b/src/primed_ai/data/valvular_dataset.py @@ -38,11 +38,11 @@ class ValvularBatch: subject_id: torch.Tensor echo_emb: torch.Tensor # (B, 1024) or (B, N_clips, 1024) - ecg_emb: torch.Tensor # (B, 768) - labels: torch.Tensor # (B, 3) binary targets [AS, MR, TR] - grades: torch.Tensor # (B, 3) ordinal severity grades [0..3] + ecg_emb: torch.Tensor # (B, 768) + labels: torch.Tensor # (B, 3) binary targets [AS, MR, TR] + grades: torch.Tensor # (B, 3) ordinal severity grades [0..3] has_echo: torch.Tensor # (B,) boolean mask - has_ecg: torch.Tensor # (B,) boolean mask + has_ecg: torch.Tensor # (B,) boolean mask demographics: dict[str, list] @@ -73,14 +73,22 @@ def __getitem__(self, idx: int) -> dict: row = self.df.iloc[idx] # Extract embeddings - if "echo_embedding" in row and row["echo_embedding"] is not None and not (isinstance(row["echo_embedding"], float) and np.isnan(row["echo_embedding"])): + if ( + "echo_embedding" in row + and row["echo_embedding"] is not None + and not (isinstance(row["echo_embedding"], float) and np.isnan(row["echo_embedding"])) + ): echo_vec = np.asarray(row["echo_embedding"], dtype=np.float32) has_echo = True else: echo_vec = np.zeros(self.echo_dim, dtype=np.float32) has_echo = False - if "ecg_embedding" in row and row["ecg_embedding"] is not None and not (isinstance(row["ecg_embedding"], float) and np.isnan(row["ecg_embedding"])): + if ( + "ecg_embedding" in row + and row["ecg_embedding"] is not None + and not (isinstance(row["ecg_embedding"], float) and np.isnan(row["ecg_embedding"])) + ): ecg_vec = np.asarray(row["ecg_embedding"], dtype=np.float32) has_ecg = True else: @@ -96,18 +104,24 @@ def __getitem__(self, idx: int) -> dict: has_ecg = False # Multi-task binary labels - labels = np.array([ - float(bool(row.get("as_moderate_or_severe", False))), - float(bool(row.get("mr_moderate_or_severe", False))), - float(bool(row.get("tr_moderate_or_severe", False))), - ], dtype=np.float32) + labels = np.array( + [ + float(bool(row.get("as_moderate_or_severe", False))), + float(bool(row.get("mr_moderate_or_severe", False))), + float(bool(row.get("tr_moderate_or_severe", False))), + ], + dtype=np.float32, + ) # Multi-task ordinal grades (-1 if unassigned) - grades = np.array([ - float(row.get("as_grade", -1) if pd.notna(row.get("as_grade")) else -1), - float(row.get("mr_grade", -1) if pd.notna(row.get("mr_grade")) else -1), - float(row.get("tr_grade", -1) if pd.notna(row.get("tr_grade")) else -1), - ], dtype=np.float32) + grades = np.array( + [ + float(row.get("as_grade", -1) if pd.notna(row.get("as_grade")) else -1), + float(row.get("mr_grade", -1) if pd.notna(row.get("mr_grade")) else -1), + float(row.get("tr_grade", -1) if pd.notna(row.get("tr_grade")) else -1), + ], + dtype=np.float32, + ) return { "subject_id": int(row.get("subject_id", 0)), @@ -156,14 +170,36 @@ def create_valvular_dataloaders( num_workers: int = 0, mask_modality_eval: str | None = None, ) -> dict[str, DataLoader]: - df = pd.read_parquet(manifest_path) if str(manifest_path).endswith(".parquet") else pd.read_csv(manifest_path) + df = ( + pd.read_parquet(manifest_path) + if str(manifest_path).endswith(".parquet") + else pd.read_csv(manifest_path) + ) train_ds = ValvularMultimodalDataset(df, split="train") val_ds = ValvularMultimodalDataset(df, split="val", mask_modality=mask_modality_eval) test_ds = ValvularMultimodalDataset(df, split="test", mask_modality=mask_modality_eval) return { - "train": DataLoader(train_ds, batch_size=batch_size, shuffle=True, collate_fn=collate_valvular_batch, num_workers=num_workers), - "val": DataLoader(val_ds, batch_size=batch_size, shuffle=False, collate_fn=collate_valvular_batch, num_workers=num_workers), - "test": DataLoader(test_ds, batch_size=batch_size, shuffle=False, collate_fn=collate_valvular_batch, num_workers=num_workers), + "train": DataLoader( + train_ds, + batch_size=batch_size, + shuffle=True, + collate_fn=collate_valvular_batch, + num_workers=num_workers, + ), + "val": DataLoader( + val_ds, + batch_size=batch_size, + shuffle=False, + collate_fn=collate_valvular_batch, + num_workers=num_workers, + ), + "test": DataLoader( + test_ds, + batch_size=batch_size, + shuffle=False, + collate_fn=collate_valvular_batch, + num_workers=num_workers, + ), } diff --git a/src/primed_ai/probes/neural_valvular.py b/src/primed_ai/probes/neural_valvular.py index ebb0d0f..18b4d47 100644 --- a/src/primed_ai/probes/neural_valvular.py +++ b/src/primed_ai/probes/neural_valvular.py @@ -70,7 +70,9 @@ def __init__( self.echo_dim = echo_dim self.ecg_dim = ecg_dim self.norm = nn.LayerNorm(echo_dim + ecg_dim) - self.head = MultiTaskValvularHead(echo_dim + ecg_dim, hidden_dim=hidden_dim, dropout=dropout) + self.head = MultiTaskValvularHead( + echo_dim + ecg_dim, hidden_dim=hidden_dim, dropout=dropout + ) def forward( self, diff --git a/src/primed_ai/probes/valvular.py b/src/primed_ai/probes/valvular.py index 18f642a..2c28b8f 100644 --- a/src/primed_ai/probes/valvular.py +++ b/src/primed_ai/probes/valvular.py @@ -12,14 +12,11 @@ from __future__ import annotations -import json import logging from dataclasses import dataclass -from pathlib import Path from typing import Callable import numpy as np -import pandas as pd from sklearn.linear_model import LogisticRegressionCV from sklearn.metrics import average_precision_score, brier_score_loss, roc_auc_score from sklearn.preprocessing import StandardScaler @@ -47,7 +44,13 @@ def auprc_safe(y_true: np.ndarray, y_score: np.ndarray) -> float: return float("nan") -def bootstrap_ci(y_true: np.ndarray, y_score: np.ndarray, metric_fn: Callable, n_bootstrap: int = 1000, seed: int = 42) -> list[float]: +def bootstrap_ci( + y_true: np.ndarray, + y_score: np.ndarray, + metric_fn: Callable, + n_bootstrap: int = 1000, + seed: int = 42, +) -> list[float]: rng = np.random.default_rng(seed) scores = [] n = len(y_true) @@ -58,7 +61,10 @@ def bootstrap_ci(y_true: np.ndarray, y_score: np.ndarray, metric_fn: Callable, n scores.append(val) if not scores: return [float("nan"), float("nan")] - return [round(float(np.percentile(scores, 2.5)), 4), round(float(np.percentile(scores, 97.5)), 4)] + return [ + round(float(np.percentile(scores, 2.5)), 4), + round(float(np.percentile(scores, 97.5)), 4), + ] @dataclass @@ -117,7 +123,9 @@ def predict_proba(self, X: np.ndarray) -> dict[str, np.ndarray]: out[target] = clf.predict_proba(Z)[:, 1] return out - def evaluate(self, X: np.ndarray, y_dict: dict[str, np.ndarray], n_bootstrap: int = 1000) -> dict[str, ValvularEvaluationResult]: + def evaluate( + self, X: np.ndarray, y_dict: dict[str, np.ndarray], n_bootstrap: int = 1000 + ) -> dict[str, ValvularEvaluationResult]: probs = self.predict_proba(X) results = {} for target, name in VALVULAR_TARGETS: @@ -125,7 +133,9 @@ def evaluate(self, X: np.ndarray, y_dict: dict[str, np.ndarray], n_bootstrap: in y_true = np.asarray(y_dict[target], dtype=bool) score = probs[target] auc = auroc_safe(y_true, score) - ci = bootstrap_ci(y_true, score, auroc_safe, n_bootstrap=n_bootstrap, seed=self.seed) + ci = bootstrap_ci( + y_true, score, auroc_safe, n_bootstrap=n_bootstrap, seed=self.seed + ) prc = auprc_safe(y_true, score) brier = float(brier_score_loss(y_true, score)) results[target] = ValvularEvaluationResult( diff --git a/tests/test_neural_valvular_probes.py b/tests/test_neural_valvular_probes.py index 1f7dc52..357e698 100644 --- a/tests/test_neural_valvular_probes.py +++ b/tests/test_neural_valvular_probes.py @@ -1,7 +1,6 @@ """Unit tests for Neural Multimodal Valvular Probes (Task B).""" import torch -import pytest from primed_ai.probes.neural_valvular import ( ConcatMLPValvularProbe, diff --git a/tests/test_valvular_dataset.py b/tests/test_valvular_dataset.py index 67f53ce..48e5f00 100644 --- a/tests/test_valvular_dataset.py +++ b/tests/test_valvular_dataset.py @@ -8,37 +8,41 @@ import pytest import torch -ROOT = Path(__file__).resolve().parents[1] -if str(ROOT) not in sys.path: - sys.path.insert(0, str(ROOT)) - from primed_ai.data.valvular_dataset import ( ValvularMultimodalDataset, - collate_valvular_batch, create_valvular_dataloaders, ) -from scripts.build_valvular_manifest import build_valvular_manifest + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +from scripts.build_valvular_manifest import build_valvular_manifest # noqa: E402 @pytest.fixture def dummy_valvular_df(): rng = np.random.default_rng(42) n = 20 - return pd.DataFrame({ - "subject_id": np.arange(100, 100 + n), - "split": ["train"] * 14 + ["val"] * 2 + ["test"] * 4, - "as_moderate_or_severe": rng.choice([False, True], n), - "mr_moderate_or_severe": rng.choice([False, True], n), - "tr_moderate_or_severe": rng.choice([False, True], n), - "as_grade": rng.choice([0, 1, 2, 3], n), - "mr_grade": rng.choice([0, 1, 2, 3], n), - "tr_grade": rng.choice([0, 1, 2, 3], n), - "echo_embedding": [rng.normal(0, 1, 1024).astype(np.float32).tolist() for _ in range(n)], - "ecg_embedding": [rng.normal(0, 1, 768).astype(np.float32).tolist() for _ in range(n)], - "sex": ["M", "F"] * (n // 2), - "age": np.linspace(30, 80, n), - "race": ["WHITE"] * n, - }) + return pd.DataFrame( + { + "subject_id": np.arange(100, 100 + n), + "split": ["train"] * 14 + ["val"] * 2 + ["test"] * 4, + "as_moderate_or_severe": rng.choice([False, True], n), + "mr_moderate_or_severe": rng.choice([False, True], n), + "tr_moderate_or_severe": rng.choice([False, True], n), + "as_grade": rng.choice([0, 1, 2, 3], n), + "mr_grade": rng.choice([0, 1, 2, 3], n), + "tr_grade": rng.choice([0, 1, 2, 3], n), + "echo_embedding": [ + rng.normal(0, 1, 1024).astype(np.float32).tolist() for _ in range(n) + ], + "ecg_embedding": [rng.normal(0, 1, 768).astype(np.float32).tolist() for _ in range(n)], + "sex": ["M", "F"] * (n // 2), + "age": np.linspace(30, 80, n), + "race": ["WHITE"] * n, + } + ) def test_valvular_dataset_shapes_and_splits(dummy_valvular_df): @@ -91,7 +95,9 @@ def test_collate_and_dataloader(dummy_valvular_df, tmp_path): def test_build_valvular_manifest_runner(dummy_valvular_df, tmp_path): cohort_path = tmp_path / "cohort.parquet" - dummy_valvular_df.drop(columns=["echo_embedding", "ecg_embedding"]).to_parquet(cohort_path, index=False) + dummy_valvular_df.drop(columns=["echo_embedding", "ecg_embedding"]).to_parquet( + cohort_path, index=False + ) out_manifest = tmp_path / "out_manifest.parquet" out_meta = tmp_path / "out_meta.csv" diff --git a/tests/test_valvular_probes.py b/tests/test_valvular_probes.py index a1cea93..04571f2 100644 --- a/tests/test_valvular_probes.py +++ b/tests/test_valvular_probes.py @@ -1,20 +1,24 @@ import sys from pathlib import Path + import numpy as np import pytest -ROOT = Path(__file__).resolve().parents[1] -if str(ROOT) not in sys.path: - sys.path.insert(0, str(ROOT)) - from primed_ai.probes.valvular import ( VALVULAR_TARGETS, MultiTargetValvularClassifier, auprc_safe, auroc_safe, - bootstrap_ci, ) -from scripts.extract_valvular_labels import parse_echo_text_with_ner + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + sys.path.insert(0, str(ROOT)) + +try: + from scripts.extract_valvular_labels import parse_echo_text_with_ner +except ImportError: # medspacy and the BigQuery client stack are not core deps + parse_echo_text_with_ner = None def test_valvular_targets_defined(): @@ -61,13 +65,19 @@ def test_multitarget_valvular_classifier_fit_and_predict(): probs = clf.predict_proba(X) assert "as_moderate_or_severe" in probs assert len(probs["as_moderate_or_severe"]) == n - assert (probs["as_moderate_or_severe"] >= 0).all() and (probs["as_moderate_or_severe"] <= 1).all() + assert (probs["as_moderate_or_severe"] >= 0).all() and ( + probs["as_moderate_or_severe"] <= 1 + ).all() eval_results = clf.evaluate(X, y_dict, n_bootstrap=50) assert "as_moderate_or_severe" in eval_results assert eval_results["as_moderate_or_severe"].auroc >= 0.0 +@pytest.mark.skipif( + parse_echo_text_with_ner is None, + reason="medspacy (and the extraction script's other deps) not installed", +) def test_clinical_ner_extraction_negation_and_severity(): # Severe AS text1 = "Echocardiogram shows calcified aortic leaflets with severe aortic stenosis."