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."