"""Target-region preparation and preservation QC for core language ROIs.""" from __future__ import annotations from dataclasses import asdict, dataclass from pathlib import Path from typing import Any import nibabel as nib from nibabel.processing import resample_from_to import numpy as np import pandas as pd CORE_ROI_NAMES: tuple[str, ...] = ( "TP", "aSTS", "pSTS", "BA44", "BA45", "BA47", "AG_TPJ", ) @dataclass(frozen=True) class CoreROICoverageRow: """Coverage diagnostics for one subject-run and one core ROI.""" subject: str run: int roi_name: str pre_mask_voxels: int post_mask_voxels: int retention_ratio: float centroid_shift_mm: float | None min_voxels_required: int passed: bool def _resample_binary_roi_to_reference( roi_img: nib.Nifti1Image, reference_img: nib.Nifti1Image, ) -> np.ndarray: resampled = resample_from_to( roi_img, (reference_img.shape, reference_img.affine), order=0, ) return resampled.get_fdata() > 0.5 def _compute_centroid_mm(mask_bool: np.ndarray, affine: np.ndarray) -> np.ndarray | None: indices = np.argwhere(mask_bool) if indices.size == 0: return None points_mm = nib.affines.apply_affine(affine, indices) return np.mean(points_mm, axis=0) def load_core_roi_masks( roi_mask_dir: Path, reference_img: nib.Nifti1Image, ) -> tuple[dict[str, np.ndarray], dict[str, int]]: """Load and resample all core ROI masks into reference analysis grid.""" roi_masks: dict[str, np.ndarray] = {} roi_voxel_counts: dict[str, int] = {} for roi_name in CORE_ROI_NAMES: roi_path = roi_mask_dir / f"{roi_name}.nii.gz" if not roi_path.exists(): raise FileNotFoundError(f"Missing core ROI mask: {roi_path}") roi_img = nib.load(str(roi_path)) roi_bool = _resample_binary_roi_to_reference(roi_img=roi_img, reference_img=reference_img) roi_masks[roi_name] = roi_bool roi_voxel_counts[roi_name] = int(roi_bool.sum()) return roi_masks, roi_voxel_counts def evaluate_core_roi_preservation( manifest_df: pd.DataFrame, run_mask_map: dict[tuple[str, int], np.ndarray], core_roi_masks: dict[str, np.ndarray], reference_affine: np.ndarray, min_voxels_required: int, ) -> tuple[pd.DataFrame, pd.DataFrame, dict[str, Any]]: """Evaluate whether each core ROI is preserved for every subject-run.""" rows: list[CoreROICoverageRow] = [] roi_centroids_pre: dict[str, np.ndarray | None] = { roi_name: _compute_centroid_mm(mask_bool=mask_bool, affine=reference_affine) for roi_name, mask_bool in core_roi_masks.items() } for row in manifest_df.itertuples(index=False): subject = str(getattr(row, "subject")) run = int(getattr(row, "run")) run_key = (subject, run) run_mask = run_mask_map.get(run_key) if run_mask is None: for roi_name, roi_mask in core_roi_masks.items(): rows.append( CoreROICoverageRow( subject=subject, run=run, roi_name=roi_name, pre_mask_voxels=int(roi_mask.sum()), post_mask_voxels=0, retention_ratio=0.0, centroid_shift_mm=None, min_voxels_required=min_voxels_required, passed=False, ) ) continue for roi_name, roi_mask in core_roi_masks.items(): pre_count = int(roi_mask.sum()) post_mask = roi_mask & run_mask post_count = int(post_mask.sum()) retention_ratio = float(post_count / pre_count) if pre_count > 0 else 0.0 centroid_pre = roi_centroids_pre.get(roi_name) centroid_post = _compute_centroid_mm(mask_bool=post_mask, affine=reference_affine) if centroid_pre is not None and centroid_post is not None: centroid_shift = float(np.linalg.norm(centroid_post - centroid_pre)) else: centroid_shift = None passed = bool(pre_count > 0 and post_count >= min_voxels_required) rows.append( CoreROICoverageRow( subject=subject, run=run, roi_name=roi_name, pre_mask_voxels=pre_count, post_mask_voxels=post_count, retention_ratio=retention_ratio, centroid_shift_mm=centroid_shift, min_voxels_required=min_voxels_required, passed=passed, ) ) coverage_df = pd.DataFrame([asdict(row) for row in rows]) if not coverage_df.empty: coverage_df = coverage_df.sort_values(["subject", "run", "roi_name"]).reset_index(drop=True) failures_df = coverage_df[coverage_df["passed"] == False].copy() if not coverage_df.empty else pd.DataFrame() if not failures_df.empty: failures_df = failures_df.sort_values(["subject", "run", "roi_name"]).reset_index(drop=True) per_run_failures = 0 if not failures_df.empty: per_run_failures = int(failures_df[["subject", "run"]].drop_duplicates().shape[0]) qc_summary: dict[str, Any] = { "n_coverage_rows": int(len(coverage_df)), "n_fail_rows": int(len(failures_df)), "n_subject_run_failures": per_run_failures, "min_voxels_required": int(min_voxels_required), } if not coverage_df.empty: qc_summary["retention_ratio_min"] = float(coverage_df["retention_ratio"].min()) qc_summary["retention_ratio_max"] = float(coverage_df["retention_ratio"].max()) return coverage_df, failures_df, qc_summary