File size: 6,661 Bytes
8c5a642 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 | """Spatial harmonization utilities for A1 baseline bootstrap."""
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
@dataclass(frozen=True)
class RunMaskQC:
"""Per-run QC record used when building the analysis mask."""
subject: str
run: int
derivatives_bold_path: str
run_mask_method: str
needs_resample_to_reference: bool
run_mask_voxels: int
def _make_reference_3d_from_manifest(manifest_df: pd.DataFrame) -> nib.Nifti1Image:
if manifest_df.empty:
raise ValueError("Manifest is empty; cannot create spatial reference")
first_path = Path(str(manifest_df.iloc[0]["derivatives_bold_path"]))
first_img = nib.load(str(first_path))
if len(first_img.shape) != 4:
raise ValueError(f"Expected 4D BOLD image but got shape={first_img.shape} for {first_path}")
shape_3d = first_img.shape[:3]
return nib.Nifti1Image(np.zeros(shape_3d, dtype=np.uint8), first_img.affine)
def _same_grid(img_a: nib.Nifti1Image, img_b: nib.Nifti1Image, atol: float = 1e-5) -> bool:
return img_a.shape == img_b.shape and np.allclose(img_a.affine, img_b.affine, atol=atol)
def _extract_run_mask_from_bold(
bold_img: nib.Nifti1Image,
method: str,
epsilon: float,
) -> nib.Nifti1Image:
shape = bold_img.shape
if len(shape) != 4:
raise ValueError(f"Expected 4D BOLD image, got shape={shape}")
if method == "first_volume_nonzero":
volume = np.asanyarray(bold_img.dataobj[..., 0])
run_mask = np.abs(volume) > epsilon
elif method == "temporal_any_nonzero":
run_mask = np.zeros(shape[:3], dtype=bool)
for time_index in range(shape[3]):
volume = np.asanyarray(bold_img.dataobj[..., time_index])
run_mask |= np.abs(volume) > epsilon
else:
raise ValueError(f"Unsupported run mask method: {method}")
return nib.Nifti1Image(run_mask.astype(np.uint8), bold_img.affine)
def _resample_binary_mask_to_reference(
binary_mask_img: nib.Nifti1Image,
reference_3d_img: nib.Nifti1Image,
) -> nib.Nifti1Image:
resampled = resample_from_to(
binary_mask_img,
(reference_3d_img.shape, reference_3d_img.affine),
order=0,
)
binary_data = (resampled.get_fdata() > 0.5).astype(np.uint8)
return nib.Nifti1Image(binary_data, reference_3d_img.affine)
def build_symmetric_analysis_mask(
manifest_df: pd.DataFrame,
run_mask_method: str = "first_volume_nonzero",
epsilon: float = 1e-6,
) -> tuple[nib.Nifti1Image, dict[tuple[str, int], np.ndarray], pd.DataFrame, dict[str, Any]]:
"""Build analysis mask in a canonical grid with left-right symmetry enforced.
Returns:
analysis_mask_img: symmetric boolean mask image in reference grid
run_mask_map: (subject, run) -> boolean 3D run mask in reference grid
run_mask_qc_df: per-run mask stats including whether resampling was needed
mask_qc: summary dictionary
"""
if manifest_df.empty:
raise ValueError("Manifest is empty; cannot build analysis mask")
required_columns = {"subject", "run", "derivatives_bold_path"}
missing = required_columns.difference(manifest_df.columns)
if missing:
raise ValueError(f"Manifest missing required columns: {sorted(missing)}")
reference_3d_img = _make_reference_3d_from_manifest(manifest_df)
reference_affine = reference_3d_img.affine
run_mask_map: dict[tuple[str, int], np.ndarray] = {}
run_qc_rows: list[RunMaskQC] = []
intersection_mask: np.ndarray | None = None
n_resampled = 0
for row in manifest_df.itertuples(index=False):
subject = str(getattr(row, "subject"))
run = int(getattr(row, "run"))
bold_path = Path(str(getattr(row, "derivatives_bold_path")))
bold_img = nib.load(str(bold_path))
run_mask_img = _extract_run_mask_from_bold(
bold_img=bold_img,
method=run_mask_method,
epsilon=epsilon,
)
needs_resample = not _same_grid(run_mask_img, reference_3d_img)
if needs_resample:
n_resampled += 1
run_mask_img = _resample_binary_mask_to_reference(run_mask_img, reference_3d_img)
run_mask_bool = run_mask_img.get_fdata() > 0.5
run_key = (subject, run)
run_mask_map[run_key] = run_mask_bool
if intersection_mask is None:
intersection_mask = run_mask_bool.copy()
else:
intersection_mask &= run_mask_bool
run_qc_rows.append(
RunMaskQC(
subject=subject,
run=run,
derivatives_bold_path=str(bold_path),
run_mask_method=run_mask_method,
needs_resample_to_reference=needs_resample,
run_mask_voxels=int(run_mask_bool.sum()),
)
)
if intersection_mask is None:
raise RuntimeError("Could not compute intersection mask from manifest entries")
swapped_mask = np.flip(intersection_mask, axis=0)
symmetric_mask = intersection_mask & swapped_mask
analysis_mask_img = nib.Nifti1Image(symmetric_mask.astype(np.uint8), reference_affine)
shape_x = symmetric_mask.shape[0]
mid_x = shape_x // 2
left_voxels = int(symmetric_mask[:mid_x, :, :].sum())
right_voxels = int(symmetric_mask[-mid_x:, :, :].sum()) if mid_x > 0 else 0
run_mask_qc_df = pd.DataFrame([asdict(row) for row in run_qc_rows])
if not run_mask_qc_df.empty:
run_mask_qc_df = run_mask_qc_df.sort_values(["subject", "run"]).reset_index(drop=True)
orientation_codes = "".join(nib.aff2axcodes(reference_affine))
mask_qc: dict[str, Any] = {
"run_mask_method": run_mask_method,
"n_manifest_rows": int(len(manifest_df)),
"n_run_masks": int(len(run_mask_map)),
"n_resampled_run_masks": int(n_resampled),
"reference_shape": [int(dim) for dim in analysis_mask_img.shape],
"reference_orientation": orientation_codes,
"intersection_voxels": int(intersection_mask.sum()),
"symmetric_mask_voxels": int(symmetric_mask.sum()),
"left_hemisphere_voxels": left_voxels,
"right_hemisphere_voxels": right_voxels,
}
if not run_mask_qc_df.empty:
mask_qc["run_mask_voxels_min"] = int(run_mask_qc_df["run_mask_voxels"].min())
mask_qc["run_mask_voxels_max"] = int(run_mask_qc_df["run_mask_voxels"].max())
return analysis_mask_img, run_mask_map, run_mask_qc_df, mask_qc
|