whuang06's picture
Add ICAInterp visualizer Space
37e13f7 verified
Raw
History Blame Contribute Delete
6.26 kB
"""ROI border overlays for Huth story-listening flatmaps."""
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
import h5py
import numpy as np
from PIL import Image, ImageDraw
from scipy import sparse
from scipy.ndimage import binary_closing, binary_dilation, binary_erosion
from flatmap_plotting import text_size
@dataclass(frozen=True)
class ROIGroup:
label: str
description: str
color: tuple[int, int, int]
rois: tuple[str, ...]
ROI_GROUPS = (
ROIGroup("AC", "auditory cortex", (0, 220, 255), ("AC",)),
ROIGroup("Lang-frontal", "Broca", (255, 180, 0), ("Broca",)),
ROIGroup("Lang-temporal", "STG/STS/LTC/TC", (255, 80, 0), ("STG", "STG_old", "STS", "LTC", "TC")),
ROIGroup("Lang-parietal", "LPC/cIPL", (255, 225, 0), ("LPC", "cIPL")),
ROIGroup(
"High-level visual",
"PPA/RSC/OPA/FFA/FBA/OFA/EBA/LO/VTC/hMT",
(255, 35, 210),
("PPA", "RSC", "OPA", "FFA", "FBA", "OFA", "EBA", "LO", "VTC", "hMT"),
),
ROIGroup("Dorsal attention", "IPS/FEF/SEF", (170, 105, 255), ("IPS", "FEF", "SEF")),
ROIGroup("Frontal control", "FC/FO/MPC", (65, 95, 255), ("FC", "FO", "MPC")),
)
def load_roi_group_masks(data_root: Path, subject: str, threshold: float = 0.25) -> list[dict[str, object]]:
mapper_path = data_root / "mappers" / f"{subject}_mappers.hdf"
groups = []
with h5py.File(mapper_path, "r") as h5:
shape = tuple(int(value) for value in h5["voxel_to_flatmap_shape"][:])
mapper = sparse.csr_matrix(
(
h5["voxel_to_flatmap_data"][:],
h5["voxel_to_flatmap_indices"][:],
h5["voxel_to_flatmap_indptr"][:],
),
shape=shape,
)
flatmap_mask = h5["flatmap_mask"][:].astype(bool)
denominator = np.asarray(mapper @ np.ones(shape[1], dtype=np.float32), dtype=np.float32)
active = denominator > 1e-6
for group in ROI_GROUPS:
available = [name for name in group.rois if f"roi_mask_{name}" in h5]
if not available:
continue
voxel_mask = np.zeros(shape[1], dtype=np.float32)
for name in available:
voxel_mask = np.maximum(voxel_mask, (np.asarray(h5[f"roi_mask_{name}"][:]) > 0.5).astype(np.float32))
weighted = np.asarray(mapper @ voxel_mask, dtype=np.float32)
coverage = np.zeros_like(weighted)
coverage[active] = weighted[active] / denominator[active]
flatmap_roi = np.zeros(flatmap_mask.shape, dtype=bool)
flatmap_roi[flatmap_mask] = coverage > threshold
if np.any(flatmap_roi):
groups.append(
{
"label": group.label,
"description": group.description,
"color": group.color,
"rois": available,
"mask": flatmap_roi,
}
)
return groups
def border_from_mask(mask: np.ndarray, panel_mask: np.ndarray, width: int, smooth_iterations: int) -> np.ndarray:
if not np.any(mask):
return np.zeros_like(mask, dtype=bool)
structure = np.ones((3, 3), dtype=bool)
smooth_mask = mask
if smooth_iterations > 0:
smooth_mask = binary_closing(mask, structure=structure, iterations=smooth_iterations)
border = binary_dilation(smooth_mask, structure=structure, iterations=width) & ~binary_erosion(
smooth_mask,
structure=structure,
iterations=width,
border_value=0,
)
return border & panel_mask
def overlay_roi_borders(
image: Image.Image,
roi_groups: list[dict[str, object]],
crop: tuple[slice, slice],
panel_mask: np.ndarray,
width: int = 1,
smooth_iterations: int = 1,
dot_step: int = 7,
) -> Image.Image:
if not roi_groups:
return image
rgba = np.asarray(image.convert("RGBA")).copy()
for group in roi_groups:
roi_mask = np.asarray(group["mask"], dtype=bool)[crop] & panel_mask
border = border_from_mask(roi_mask, panel_mask, width, smooth_iterations)
if not np.any(border):
continue
halo = binary_dilation(border, structure=np.ones((3, 3), dtype=bool), iterations=1) & panel_mask
rgba[halo, :3] = (20, 20, 20)
rgba[halo, 3] = 255
if dot_step > 0:
rows, cols = np.indices(roi_mask.shape)
dots = roi_mask & panel_mask & ~border & (((rows * 3 + cols * 5) % dot_step) == 0)
rgba[dots, :3] = group["color"]
rgba[dots, 3] = 255
rgba[border, :3] = group["color"]
rgba[border, 3] = 255
return Image.fromarray(rgba, mode="RGBA")
def roi_group_metadata(roi_groups: list[dict[str, object]]) -> list[dict[str, object]]:
return [
{
"label": str(group["label"]),
"description": str(group["description"]),
"color_rgb": list(group["color"]),
"rois": list(group["rois"]),
}
for group in roi_groups
]
def draw_roi_legend(
draw: ImageDraw.ImageDraw,
roi_groups: list[dict[str, object]],
x: int,
y: int,
max_width: int,
font,
title_font,
) -> int:
if not roi_groups:
return y
title = "ROI borders:"
draw.text((x, y), title, fill=(20, 20, 20), font=title_font)
title_w, _ = text_size(draw, title, title_font)
cursor_x = x + title_w + 18
cursor_y = y + 1
line_h = 24
for group in roi_groups:
label = f"{group['label']}: {group['description']}"
label_w, _ = text_size(draw, label, font)
item_w = 34 + label_w + 22
if cursor_x + item_w > x + max_width and cursor_x > x + title_w + 18:
cursor_x = x
cursor_y += line_h
color = tuple(int(value) for value in group["color"])
line_y = cursor_y + 11
draw.line((cursor_x, line_y, cursor_x + 24, line_y), fill=(0, 0, 0), width=5)
draw.line((cursor_x, line_y, cursor_x + 24, line_y), fill=color, width=3)
draw.text((cursor_x + 32, cursor_y), label, fill=(20, 20, 20), font=font)
cursor_x += item_w
return cursor_y + line_h