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