PromptM_UNet_Dashboard / utils /nifti_handler.py
Kiuyha's picture
Fix slider JS event payload decoupling
637f72d
Raw
History Blame Contribute Delete
16.5 kB
"""NIfTI I/O, clinical windowing, slice rendering and quantitative metrics.
All functions in this module are strictly CPU-side (numpy / nibabel) and are
safe to call from Gradio callbacks without touching the GPU:
* Volume parsing of ``.nii`` / ``.nii.gz`` with canonical RAS+ reorientation.
* Spatial validation of Ground Truth masks against the source CT
(shape, voxel spacing, orientation/affine matrix).
* Hounsfield Unit (HU) windowing with clinical presets.
* Tri-planar slice extraction + mask/GT overlay compositing.
* Dice Similarity Coefficient, IoU and volumetric difference metrics.
* Prediction export as compressed ``.nii.gz`` retaining the original affine.
Array convention: after canonicalisation volumes are ``float32[X, Y, Z]`` in
RAS+ order (axis 0 -> Right, axis 1 -> Anterior, axis 2 -> Superior).
"""
from __future__ import annotations
import logging
from dataclasses import dataclass, field
from pathlib import Path
from typing import Dict, Optional, Tuple
import numpy as np
import nibabel as nib
logger = logging.getLogger(__name__)
PLANE_AXES = {
"axial": 2, # slice index along Z (superior)
"sagittal": 0, # slice index along X (right)
"coronal": 1, # slice index along Y (anterior)
}
# Clinical HU window presets: name -> (center, width)
WINDOW_PRESETS: Dict[str, Optional[Tuple[float, float]]] = {
"Soft Tissue (C40/W400)": (40.0, 400.0),
"Bone (C400/W1800)": (400.0, 1800.0),
"Lung (C-600/W1500)": (-600.0, 1500.0),
"Custom": None,
}
# Overlay palette (RGB) — clinical high-contrast scheme
COLOR_PREDICTION = (239, 68, 68) # red #EF4444 — model prediction
COLOR_GROUND_TRUTH = (16, 185, 129) # emerald #10b981 — ground truth
COLOR_OVERLAP = (250, 204, 21) # yellow #facc15 — prediction ∩ ground truth
def clamp_index(index: int | float | None, size: int) -> int:
"""Strictly clamp a slice index into ``[0, size - 1]`` (safe against None/float/out-of-bounds)."""
if size <= 0:
return 0
if index is None:
return size // 2
try:
val = int(round(float(index)))
return int(np.clip(val, 0, max(int(size) - 1, 0)))
except (ValueError, TypeError):
return size // 2
class VolumeMismatchError(ValueError):
"""Raised when a Ground Truth volume is spatially inconsistent with the CT."""
@dataclass
class MedicalVolume:
"""A canonical-RAS medical volume plus its spatial metadata."""
data: np.ndarray # float32 [X, Y, Z]
affine: np.ndarray # 4x4 world matrix of the canonical image
spacing: Tuple[float, float, float] # mm per voxel along (X, Y, Z)
source_path: str
name: str = "volume"
hu_min: float = field(default=0.0)
hu_max: float = field(default=0.0)
@property
def shape(self) -> Tuple[int, int, int]:
return tuple(int(v) for v in self.data.shape)
# ---------------------------------------------------------------------------
# Loading & validation
# ---------------------------------------------------------------------------
def _canonicalize(path: Path) -> tuple[nib.Nifti1Image, np.ndarray]:
img = nib.load(str(path))
canonical = nib.as_closest_canonical(img) # RAS+ orientation
return canonical, img.header.get_zooms()[:3]
def _volume_data(img: nib.Nifti1Image) -> np.ndarray:
"""Extract a strictly-3D float32 array (squeezes singleton trailing dims)."""
data = img.get_fdata(dtype=np.float32)
data = np.squeeze(data)
if data.ndim != 3:
raise ValueError(f"Expected a 3D volume, got shape {data.shape} after squeeze.")
return np.ascontiguousarray(data)
def load_medical_volume(path: str | Path, name: str = "ct") -> MedicalVolume:
"""Parse a NIfTI file, reorient to canonical RAS+, and extract metadata."""
path = Path(path)
if path.suffix not in {".nii", ".gz"} or not str(path).endswith((".nii", ".nii.gz")):
raise ValueError(f"Unsupported file type: '{path.name}'. Please upload a .nii or .nii.gz file.")
if not path.exists():
raise FileNotFoundError(path)
img, raw_zooms = _canonicalize(path)
data = _volume_data(img)
zooms = tuple(float(z) for z in img.header.get_zooms()[:3])
if any(not np.isfinite(z) or z <= 0 for z in zooms):
zooms = tuple(float(z) for z in raw_zooms)
logger.warning("Invalid canonical zooms; falling back to header zooms %s", zooms)
finite = data[np.isfinite(data)]
hu_min = float(finite.min()) if finite.size else 0.0
hu_max = float(finite.max()) if finite.size else 0.0
if not np.isfinite(data).all():
nan_count = int((~np.isfinite(data)).sum())
logger.warning("Volume contains %d non-finite voxels; they will render as the volume minimum.", nan_count)
data = np.nan_to_num(data, nan=hu_min, posinf=hu_max, neginf=hu_min)
return MedicalVolume(
data=data,
affine=np.asarray(img.affine, dtype=np.float64),
spacing=zooms,
source_path=str(path),
name=name,
hu_min=hu_min,
hu_max=hu_max,
)
def _check_close(actual, expected, tol, kind):
if actual is None or expected is None:
return None
a, e = np.asarray(actual, dtype=np.float64), np.asarray(expected, dtype=np.float64)
if a.shape != e.shape or not np.allclose(a, e, atol=tol, rtol=1e-4):
return kind
return None
def load_ground_truth_mask(gt_path: str | Path, ct: MedicalVolume) -> np.ndarray:
"""Load a GT segmentation and validate it against *ct*.
Checks dimensions, voxel spacing and the spatial orientation matrix.
Raises :class:`VolumeMismatchError` with a human-readable report on failure.
"""
gt_img, _ = _canonicalize(Path(gt_path))
gt_data = _volume_data(gt_img) > 0.5
problems = []
if gt_data.shape != ct.shape:
problems.append(
f"dimensions differ: GT {gt_data.shape} vs CT {ct.shape}"
)
else:
gt_zooms = tuple(float(z) for z in gt_img.header.get_zooms()[:3])
if (_r := _check_close(gt_zooms, ct.spacing, 1e-3, "voxel spacing")) is not None:
problems.append(f"{_r} differs: GT {tuple(round(z, 4) for z in gt_zooms)} vs CT {tuple(round(s, 4) for s in ct.spacing)}")
if (_r := _check_close(gt_img.affine, ct.affine, 1e-3, "orientation matrix")) is not None:
problems.append(f"{_r} differs (affine mismatch beyond tolerance)")
if problems:
raise VolumeMismatchError(
"Ground Truth does not match the loaded CT scan:\n- " + "\n- ".join(problems)
)
return np.ascontiguousarray(gt_data)
# ---------------------------------------------------------------------------
# HU windowing & tri-planar rendering
# ---------------------------------------------------------------------------
def resolve_window(preset: str, custom_center: float, custom_width: float) -> Tuple[float, float]:
"""Resolve a UI preset name into concrete (window_lo, window_hi) HU bounds."""
if preset in WINDOW_PRESETS and WINDOW_PRESETS[preset] is not None:
center, width = WINDOW_PRESETS[preset]
return (center - width / 2.0, center + width / 2.0)
width = max(1.0, float(custom_width))
center = float(custom_center)
return (center - width / 2.0, center + width / 2.0)
def apply_hu_window(slice_hu: np.ndarray, window_lo: float, window_hi: float) -> np.ndarray:
"""Map an HU slice to uint8 grayscale via a clinical window."""
denom = max(window_hi - window_lo, 1e-6)
norm = (np.clip(slice_hu.astype(np.float32), window_lo, window_hi) - window_lo) / denom
return (norm * 255.0).astype(np.uint8)
def orient_for_display(slice_2d: np.ndarray, plane: str) -> np.ndarray:
"""Rotate/flip a raw (i, j) plane so anatomy appears in radiological convention.
Axial : anterior at top, patient right on viewer left.
Coronal : superior at top, patient right on viewer left.
Sagittal: superior at top, anterior on viewer left.
All three conventions reduce to a transpose followed by a 180-degree
rotation of the extracted plane (rows reversed, columns reversed).
"""
if plane not in PLANE_AXES:
raise ValueError(f"Unknown plane '{plane}'")
return np.rot90(slice_2d.T, k=2)
def get_plane_slice(volume: np.ndarray, plane: str, index: int) -> np.ndarray:
"""Extract and display-orient a 2D plane from a [X, Y, Z] volume.
Axis mapping (canonical RAS+ NIfTI -> viewer):
array[X, Y, Z] -> tensor[B, C, X, Y, Z] (X->D, Y->H, Z->W)
axial slices along Z: volume[:, :, z]
sagittal slices along X: volume[x, :, :]
coronal slices along Y: volume[:, y, :]
``index`` is strictly clamped to ``[0, dim - 1]`` before access.
"""
axis = PLANE_AXES[plane]
idx = clamp_index(index, volume.shape[axis])
if plane == "axial":
sl = volume[:, :, idx]
elif plane == "sagittal":
sl = volume[idx, :, :]
else: # coronal
sl = volume[:, idx, :]
return orient_for_display(sl, plane)
def compose_overlay(
gray_slice: np.ndarray,
pred_slice: Optional[np.ndarray],
gt_slice: Optional[np.ndarray],
overlay_mode: str,
pred_opacity: float,
gt_opacity: float,
) -> np.ndarray:
"""Blend prediction / GT boolean slices over a grayscale slice.
``overlay_mode``: one of ``None (CT only)``, ``Prediction``, ``Ground Truth``,
``Combined (Pred vs GT)``. Colors: GT green, prediction red, overlap yellow.
All inputs must already be display-oriented via :func:`get_plane_slice`.
"""
rgb = np.stack([gray_slice] * 3, axis=-1).astype(np.float32)
show_pred = overlay_mode in ("Prediction", "Combined (Pred vs GT)") and pred_slice is not None
show_gt = overlay_mode in ("Ground Truth", "Combined (Pred vs GT)") and gt_slice is not None
def _blend(mask: np.ndarray, color: Tuple[int, int, int], alpha: float) -> None:
if not mask.any() or alpha <= 0:
return
m = mask[..., None].astype(np.float32)
color_arr = np.asarray(color, dtype=np.float32)[None, None, :]
rgb[:] = rgb * (1.0 - m * alpha) + color_arr * (m * alpha)
if show_pred and show_gt:
overlap = pred_slice & gt_slice
_blend(pred_slice & ~overlap, COLOR_PREDICTION, float(pred_opacity))
_blend(gt_slice & ~overlap, COLOR_GROUND_TRUTH, float(gt_opacity))
_blend(overlap, COLOR_OVERLAP, max(float(pred_opacity), float(gt_opacity)))
else:
if show_pred:
_blend(pred_slice, COLOR_PREDICTION, float(pred_opacity))
if show_gt:
_blend(gt_slice, COLOR_GROUND_TRUTH, float(gt_opacity))
return np.clip(rgb, 0, 255).astype(np.uint8)
# ---------------------------------------------------------------------------
# Quantitative metrics — Dice + NSD (raw patient grid)
# ---------------------------------------------------------------------------
def volume_ml(n_voxels: int, spacing: Tuple[float, float, float]) -> float:
"""Convert a voxel count to millilitres using voxel spacing (mm)."""
return float(n_voxels * float(np.prod(spacing)) / 1000.0)
def _get_surface(mask: np.ndarray) -> np.ndarray:
"""Extract 3-D surface voxels via binary erosion (6-connectivity via 3x3x3)."""
try:
from scipy.ndimage import binary_erosion # type: ignore
eroded = binary_erosion(mask, structure=np.ones((3, 3, 3)))
return mask & (~eroded)
except ImportError:
# Fallback: torch max_pool erosion on CPU (mirrors training code)
import torch
import torch.nn.functional as F
t = torch.from_numpy(mask.astype(np.float32)).unsqueeze(0).unsqueeze(0)
eroded = -F.max_pool3d(-t, kernel_size=3, stride=1, padding=1)
surf = (t.bool() & (eroded < 0.5)).squeeze(0).squeeze(0).numpy()
return surf.astype(bool)
def _nsd_score(
pred: np.ndarray,
gt: np.ndarray,
spacing: Tuple[float, float, float],
tolerance_mm: float = 2.0,
) -> float:
"""Normalized Surface Distance on the *raw* patient grid.
Mirrors ``promptm_unet.training.metrics.SegmentationMetrics``:
surface extraction via erosion, then symmetric distance check
``(psw + tsw) / (ps + ts)``. Distances are true Euclidean via
``distance_transform_edt`` with anisotropic ``sampling=spacing``.
"""
pred_surf = _get_surface(pred.astype(bool))
gt_surf = _get_surface(gt.astype(bool))
ps, ts = int(pred_surf.sum()), int(gt_surf.sum())
if ps == 0 or ts == 0:
return 0.0
try:
from scipy.ndimage import distance_transform_edt # type: ignore
# distance to nearest GT surface for every voxel
dt_to_gt = distance_transform_edt(np.logical_not(gt_surf), sampling=spacing)
dt_to_pred = distance_transform_edt(np.logical_not(pred_surf), sampling=spacing)
psw = int((dt_to_gt[pred_surf] <= tolerance_mm).sum())
tsw = int((dt_to_pred[gt_surf] <= tolerance_mm).sum())
return float((psw + tsw) / (ps + ts))
except ImportError:
# Fallback: torch BFS distance transform (isotropic approximation)
import torch
import torch.nn.functional as F
device = torch.device("cpu")
def _dt_torch(surf: np.ndarray) -> np.ndarray:
s = torch.from_numpy(surf).unsqueeze(0).unsqueeze(0).bool().to(device)
dist = torch.full(s.shape, 1e6, dtype=torch.float32, device=device)
dist[s] = 0.0
kernel = torch.ones(1, 1, 3, 3, 3, device=device)
max_steps = max(surf.shape) + 1
try:
import math as _math
max_steps = min(max_steps, _math.ceil(tolerance_mm / float(np.mean(spacing))) + 2)
except Exception:
pass
for step in range(1, max_steps):
reached = F.conv3d((dist < step).float(), kernel, padding=1) > 0
update = reached & (dist >= step)
if int(update.sum()) == 0:
break
dist[update] = float(step)
return (dist.squeeze(0).squeeze(0).numpy() * float(np.mean(spacing))).astype(np.float32)
dt_gt = _dt_torch(gt_surf)
dt_pred = _dt_torch(pred_surf)
psw = int((dt_gt[pred_surf] <= tolerance_mm).sum())
tsw = int((dt_pred[gt_surf] <= tolerance_mm).sum())
return float((psw + tsw) / (ps + ts))
def compute_metrics(
pred: np.ndarray,
gt: np.ndarray,
spacing: Tuple[float, float, float],
nsd_tolerance_mm: float = 2.0,
) -> Dict[str, object]:
"""Compute DSC, NSD and volumetric differences on the **raw** patient grid.
Both ``pred`` and ``gt`` must share the CT's original shape/spacing
(i.e. after paste-back, not the 96³ preprocessing grid).
"""
p = pred.astype(bool)
g = gt.astype(bool)
tp = int(np.logical_and(p, g).sum())
fp = int(np.logical_and(p, ~g).sum())
fn = int(np.logical_and(~p, g).sum())
dice = (2.0 * tp) / (2.0 * tp + fp + fn) if (2 * tp + fp + fn) > 0 else 0.0
nsd = _nsd_score(p, g, spacing, tolerance_mm=nsd_tolerance_mm)
pred_ml = volume_ml(int(p.sum()), spacing)
gt_ml = volume_ml(int(g.sum()), spacing)
rel_diff = ((pred_ml - gt_ml) / gt_ml) if gt_ml > 0 else 0.0
return {
"dice": float(dice),
"nsd": float(nsd),
# keep IoU for backward-compat display if needed
"iou": float(tp / (tp + fp + fn) if (tp + fp + fn) > 0 else 0.0),
"true_positives_voxels": tp,
"false_positives_voxels": fp,
"false_negatives_voxels": fn,
"predicted_volume_ml": round(pred_ml, 2),
"ground_truth_volume_ml": round(gt_ml, 2),
"relative_volume_error_pct": round(rel_diff * 100.0, 2),
}
# ---------------------------------------------------------------------------
# Export
# ---------------------------------------------------------------------------
def save_mask_nifti(mask: np.ndarray, reference: MedicalVolume, out_path: str | Path) -> Path:
"""Save a boolean mask as compressed .nii.gz retaining the reference spatial metadata."""
out_path = Path(out_path)
out_path.parent.mkdir(parents=True, exist_ok=True)
img = nib.Nifti1Image(mask.astype(np.uint8), affine=reference.affine)
img.header.set_zooms(tuple(float(s) for s in reference.spacing))
img.set_qform(reference.affine, code=1) # 1 = scanner_anat
img.set_sform(reference.affine, code=1)
nib.save(img, str(out_path))
return out_path