| """LitCoder-style fsaverage5 surface plotting helpers.""" |
|
|
| from __future__ import annotations |
|
|
| import os |
| import warnings |
| from io import BytesIO |
| from pathlib import Path |
| from typing import Any |
|
|
| import numpy as np |
|
|
| DEFAULT_SCRATCH_ROOT = Path(os.environ.get("BRAINENCODING_SCRATCH_ROOT", "/storage/scratch1/8/whuang409")) |
| DEFAULT_FSAVERAGE5 = Path("/storage/project/r-aivanova7-0/shared/.env/freesurfer/subjects/fsaverage5") |
| N_HEMI = 10242 |
| N_VERTICES = 20484 |
|
|
| os.environ.setdefault("MPLCONFIGDIR", str(DEFAULT_SCRATCH_ROOT / "matplotlib")) |
| import matplotlib |
|
|
| matplotlib.use("Agg") |
| import matplotlib.pyplot as plt |
| from matplotlib.colors import Normalize |
|
|
|
|
| def save_litcoder_style_map( |
| values: np.ndarray, |
| output_path: Path, |
| *, |
| title: str, |
| fsaverage5: Path = DEFAULT_FSAVERAGE5, |
| significant_mask: np.ndarray | None = None, |
| positive: bool = False, |
| mask_zero: bool = True, |
| vmax: float | None = None, |
| cmap: str | Any | None = None, |
| dpi: int = 150, |
| ) -> None: |
| fig = plot_litcoder_style_map( |
| values, |
| title=title, |
| fsaverage5=fsaverage5, |
| significant_mask=significant_mask, |
| positive=positive, |
| mask_zero=mask_zero, |
| vmax=vmax, |
| cmap=cmap, |
| ) |
| output_path.parent.mkdir(parents=True, exist_ok=True) |
| fig.savefig(output_path, format="png", bbox_inches="tight", dpi=dpi) |
| plt.close(fig) |
|
|
|
|
| def litcoder_style_png_bytes( |
| values: np.ndarray, |
| *, |
| title: str, |
| fsaverage5: Path = DEFAULT_FSAVERAGE5, |
| significant_mask: np.ndarray | None = None, |
| positive: bool = False, |
| mask_zero: bool = True, |
| vmax: float | None = None, |
| cmap: str | Any | None = None, |
| dpi: int = 130, |
| ) -> bytes: |
| fig = plot_litcoder_style_map( |
| values, |
| title=title, |
| fsaverage5=fsaverage5, |
| significant_mask=significant_mask, |
| positive=positive, |
| mask_zero=mask_zero, |
| vmax=vmax, |
| cmap=cmap, |
| ) |
| buffer = BytesIO() |
| fig.savefig(buffer, format="png", bbox_inches="tight", dpi=dpi) |
| plt.close(fig) |
| return buffer.getvalue() |
|
|
|
|
| def plot_litcoder_style_map( |
| values: np.ndarray, |
| *, |
| title: str, |
| fsaverage5: Path = DEFAULT_FSAVERAGE5, |
| significant_mask: np.ndarray | None = None, |
| positive: bool = False, |
| mask_zero: bool = True, |
| vmax: float | None = None, |
| cmap: str | Any | None = None, |
| ) -> plt.Figure: |
| """Plot a 20484-vector in the LitCoder four-panel fsaverage5 layout.""" |
| from nilearn import plotting |
| from nilearn.plotting.cm import cold_hot |
|
|
| display_values = prepare_values( |
| values, |
| significant_mask=significant_mask, |
| positive=positive, |
| mask_zero=mask_zero, |
| ) |
| vmax = value_limit(display_values, vmax=vmax, positive=positive) |
| vmin = 0.0 if positive else -vmax |
| cmap = cmap or ("viridis" if positive else cold_hot) |
| norm = Normalize(vmin=vmin, vmax=vmax) |
|
|
| fsaverage5 = Path(fsaverage5) |
| meshes = { |
| "lh": fsaverage5 / "surf" / "lh.inflated", |
| "rh": fsaverage5 / "surf" / "rh.inflated", |
| } |
| for path in meshes.values(): |
| if not path.exists(): |
| raise FileNotFoundError(f"Missing fsaverage5 surface mesh: {path}") |
|
|
| fig = plt.figure(figsize=(15, 10)) |
| views = [ |
| ("lh", "left", "lateral", 231, "Left Lateral"), |
| ("lh", "left", "medial", 232, "Left Medial"), |
| ("rh", "right", "lateral", 234, "Right Lateral"), |
| ("rh", "right", "medial", 235, "Right Medial"), |
| ] |
| for hemi_key, hemi_name, view, subplot, pane_title in views: |
| hemi_values = display_values[:N_HEMI] if hemi_key == "lh" else display_values[N_HEMI:] |
| ax = fig.add_subplot(subplot, projection="3d") |
| with warnings.catch_warnings(): |
| warnings.filterwarnings("ignore", message="All-NaN slice encountered", category=RuntimeWarning) |
| plotting.plot_surf_stat_map( |
| str(meshes[hemi_key]), |
| hemi_values, |
| hemi=hemi_name, |
| view=view, |
| colorbar=False, |
| axes=ax, |
| cmap=cmap, |
| vmin=vmin, |
| vmax=vmax, |
| title=pane_title, |
| ) |
|
|
| sm = plt.cm.ScalarMappable(norm=norm, cmap=cmap) |
| sm.set_array([]) |
| cax = fig.add_axes([0.92, 0.15, 0.02, 0.7]) |
| fig.colorbar(sm, cax=cax) |
| fig.suptitle(title, fontsize=16) |
| with warnings.catch_warnings(): |
| warnings.filterwarnings("ignore", message="This figure includes Axes that are not compatible with tight_layout.*") |
| fig.tight_layout(rect=[0.03, 0.03, 0.9, 0.97]) |
| return fig |
|
|
|
|
| def prepare_values( |
| values: np.ndarray, |
| *, |
| significant_mask: np.ndarray | None, |
| positive: bool, |
| mask_zero: bool, |
| ) -> np.ndarray: |
| out = np.asarray(values, dtype=np.float32).reshape(-1).copy() |
| if out.shape != (N_VERTICES,): |
| raise ValueError(f"Expected {(N_VERTICES,)} values, got {out.shape}") |
| shown = np.isfinite(out) |
| if significant_mask is not None: |
| mask = np.asarray(significant_mask, dtype=bool).reshape(-1) |
| if mask.shape != (N_VERTICES,): |
| raise ValueError(f"Expected {(N_VERTICES,)} mask values, got {mask.shape}") |
| shown &= mask |
| if mask_zero: |
| shown &= (out > 0.0) if positive else (np.abs(out) > 0.0) |
| out[~shown] = np.nan |
| return out |
|
|
|
|
| def value_limit(values: np.ndarray, *, vmax: float | None, positive: bool) -> float: |
| if vmax is not None and np.isfinite(vmax) and vmax > 0: |
| return float(vmax) |
| finite = values[np.isfinite(values)] |
| if finite.size == 0: |
| return 1.0 |
| finite = finite[finite > 0.0] if positive else np.abs(finite) |
| if finite.size == 0: |
| return 1.0 |
| limit = float(np.nanmax(finite)) |
| return limit if np.isfinite(limit) and limit > 0 else 1.0 |
|
|