csai / code /a1_pipeline /targets.py
Mohith202's picture
Add core ROI masks, visualization script, and model profiles
8c5a642
Raw
History Blame Contribute Delete
5.89 kB
"""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