248 lines
10 KiB
Python
248 lines
10 KiB
Python
"""
|
|
Load axial slices from a directory of NIfTI CT volumes.
|
|
|
|
Each ``.nii.gz`` file is treated as one patient (the filename stem is the
|
|
ground-truth patient ID). A handful of evenly-spaced axial slices are sampled
|
|
per volume, HU-windowed, and returned both as grayscale thumbnails (for the
|
|
manuscript's clustering method) and as full-resolution slices (for CNN feature
|
|
extraction).
|
|
|
|
This is used to validate the patient-clustering method on a dataset where the
|
|
true patient identity IS known — e.g. the Medical Segmentation Decathlon
|
|
Task06_Lung set — unlike IQ-OTH/NCCD, where patient IDs are unavailable.
|
|
"""
|
|
|
|
import os
|
|
import glob
|
|
import numpy as np
|
|
|
|
|
|
class NiftiSliceDataset:
|
|
"""Sample HU-windowed axial slices from a folder of NIfTI CT volumes.
|
|
|
|
Parameters
|
|
----------
|
|
volume_dir : str
|
|
Directory containing ``*.nii.gz`` volumes, one per patient.
|
|
n_slices_per_patient : int
|
|
Number of evenly-spaced axial slices to sample per volume.
|
|
thumbnail_size : tuple of int
|
|
Grayscale thumbnail size for clustering (matches the manuscript's 64x64).
|
|
hu_window : tuple of float
|
|
(low, high) Hounsfield-unit clip range applied before scaling to 8-bit.
|
|
The default (-1000, 400) keeps the body outline, soft tissue, and lung
|
|
parenchyma — the cues that identify a patient.
|
|
central_fraction : float
|
|
Fraction of the volume (centred on the middle slice) to sample from,
|
|
avoiding empty/partial slices at the superior/inferior ends.
|
|
slice_axis : int
|
|
Axis along which axial slices are indexed (-1 = last, correct for the
|
|
(H, W, Z) layout used by the Decathlon volumes).
|
|
"""
|
|
|
|
def __init__(self, volume_dir, n_slices_per_patient=10,
|
|
thumbnail_size=(64, 64), hu_window=(-1000.0, 400.0),
|
|
central_fraction=0.6, slice_axis=-1, random_slices=True,
|
|
seed=42, flip_vertical=False, rotate_deg=0):
|
|
self.volume_dir = volume_dir
|
|
self.n_slices_per_patient = n_slices_per_patient
|
|
self.thumbnail_size = thumbnail_size
|
|
self.hu_window = hu_window
|
|
self.central_fraction = central_fraction
|
|
self.slice_axis = slice_axis
|
|
self.random_slices = random_slices
|
|
self.seed = seed
|
|
self.flip_vertical = flip_vertical
|
|
self.rotate_deg = rotate_deg # 0, 90, 180, or 270
|
|
|
|
# Populated by load()
|
|
self.thumbnails = None # (n_slices, H*W) float32 in [0, 1]
|
|
self.slices_rgb = None # list of (H, W) uint8 arrays
|
|
self.patient_labels = None # (n_slices,) ground-truth patient IDs
|
|
self.filenames = None # (n_slices,) per-slice names
|
|
|
|
def load(self, verbose=True):
|
|
"""Load and slice every volume in ``volume_dir``."""
|
|
import nibabel as nib
|
|
from PIL import Image
|
|
|
|
# Prefer decompressed .nii (fast memmap random-slice access) and fall
|
|
# back to .nii.gz, skipping any .gz that has an uncompressed twin.
|
|
nii = glob.glob(os.path.join(self.volume_dir, "*.nii"))
|
|
stems = {os.path.basename(p)[:-4] for p in nii}
|
|
gz = [p for p in glob.glob(os.path.join(self.volume_dir, "*.nii.gz"))
|
|
if os.path.basename(p)[:-7] not in stems]
|
|
paths = sorted(nii + gz)
|
|
if not paths:
|
|
raise FileNotFoundError(
|
|
f"No .nii/.nii.gz volumes found in {self.volume_dir}")
|
|
|
|
thumbs, slices_rgb, patient_labels, filenames = [], [], [], []
|
|
lo, hi = self.hu_window
|
|
|
|
for path in paths:
|
|
patient_id = os.path.basename(path).replace(".nii.gz", "").replace(".nii", "")
|
|
img = nib.load(path)
|
|
n_z = img.shape[self.slice_axis]
|
|
axis = self.slice_axis % img.ndim
|
|
|
|
for k in self._slice_indices(n_z):
|
|
# Lazy slice: pull only this plane off disk instead of
|
|
# materialising the whole float64 volume with get_fdata().
|
|
slicer = [slice(None)] * img.ndim
|
|
slicer[axis] = k
|
|
sl = np.asarray(img.dataobj[tuple(slicer)], dtype=np.float32)
|
|
# HU window -> 8-bit grayscale
|
|
sl = np.clip(sl, lo, hi)
|
|
sl = (255.0 * (sl - lo) / (hi - lo)).astype(np.uint8)
|
|
if self.flip_vertical:
|
|
sl = np.flipud(sl)
|
|
if self.rotate_deg:
|
|
sl = np.rot90(sl, k=self.rotate_deg // 90)
|
|
|
|
thumb = Image.fromarray(sl).convert("L").resize(
|
|
self.thumbnail_size)
|
|
thumbs.append(np.asarray(thumb, dtype=np.float32).flatten() / 255.0)
|
|
slices_rgb.append(sl)
|
|
patient_labels.append(patient_id)
|
|
filenames.append(f"{patient_id}_slice{k:03d}")
|
|
|
|
if verbose:
|
|
print(f" {patient_id}: {n_z} slices -> "
|
|
f"sampled {self.n_slices_per_patient}")
|
|
|
|
self.thumbnails = np.array(thumbs, dtype=np.float32)
|
|
self.slices_rgb = slices_rgb
|
|
self.patient_labels = np.array(patient_labels)
|
|
self.filenames = np.array(filenames)
|
|
|
|
if verbose:
|
|
print(f"\n Loaded {len(self.filenames)} slices from "
|
|
f"{len(paths)} patients")
|
|
return self
|
|
|
|
def load_all_slices(self, stride=3, verbose=True):
|
|
"""Load ALL slices from the central fraction (not just a sample).
|
|
|
|
Every ``stride``-th slice is taken to keep feature extraction manageable
|
|
while still giving K-means enough data to cluster per patient.
|
|
|
|
Returns self (populates thumbnails, slices_rgb, patient_labels, filenames).
|
|
"""
|
|
import nibabel as nib
|
|
from PIL import Image
|
|
|
|
nii = glob.glob(os.path.join(self.volume_dir, "*.nii"))
|
|
stems = {os.path.basename(p)[:-4] for p in nii}
|
|
gz = [p for p in glob.glob(os.path.join(self.volume_dir, "*.nii.gz"))
|
|
if os.path.basename(p)[:-7] not in stems]
|
|
paths = sorted(nii + gz)
|
|
if not paths:
|
|
raise FileNotFoundError(
|
|
f"No .nii/.nii.gz volumes found in {self.volume_dir}")
|
|
|
|
thumbs, slices_rgb, patient_labels, filenames = [], [], [], []
|
|
lo, hi = self.hu_window
|
|
total_slices = 0
|
|
|
|
for path in paths:
|
|
patient_id = os.path.basename(path).replace(".nii.gz", "").replace(".nii", "")
|
|
img = nib.load(path)
|
|
n_z = img.shape[self.slice_axis]
|
|
axis = self.slice_axis % img.ndim
|
|
|
|
# All slices in central fraction, stepped
|
|
half = self.central_fraction / 2.0
|
|
start = int(n_z * (0.5 - half))
|
|
end = int(n_z * (0.5 + half))
|
|
indices = list(range(start, end, stride))
|
|
|
|
for k in indices:
|
|
slicer = [slice(None)] * img.ndim
|
|
slicer[axis] = k
|
|
sl = np.asarray(img.dataobj[tuple(slicer)], dtype=np.float32)
|
|
sl = np.clip(sl, lo, hi)
|
|
sl = (255.0 * (sl - lo) / (hi - lo)).astype(np.uint8)
|
|
if self.flip_vertical:
|
|
sl = np.flipud(sl)
|
|
if self.rotate_deg:
|
|
sl = np.rot90(sl, k=self.rotate_deg // 90)
|
|
|
|
thumb = Image.fromarray(sl).convert("L").resize(
|
|
self.thumbnail_size)
|
|
thumbs.append(np.asarray(thumb, dtype=np.float32).flatten() / 255.0)
|
|
slices_rgb.append(sl)
|
|
patient_labels.append(patient_id)
|
|
filenames.append(f"{patient_id}_slice{k:03d}")
|
|
total_slices += 1
|
|
|
|
if verbose:
|
|
print(f" {patient_id}: {n_z} slices -> {len(indices)} sampled")
|
|
|
|
self.thumbnails = np.array(thumbs, dtype=np.float32)
|
|
self.slices_rgb = slices_rgb
|
|
self.patient_labels = np.array(patient_labels)
|
|
self.filenames = np.array(filenames)
|
|
|
|
if verbose:
|
|
print(f"\n Loaded {total_slices} slices from "
|
|
f"{len(paths)} patients (stride={stride})")
|
|
return self
|
|
|
|
def export_pngs(self, cache_dir):
|
|
"""Export loaded slices as PNG files for fast reloading.
|
|
|
|
Creates ``cache_dir/`` with one PNG per slice and a manifest.json
|
|
mapping paths → patient IDs and z-indices.
|
|
|
|
Returns (paths, patient_ids, z_indices).
|
|
"""
|
|
import json
|
|
from PIL import Image as PILImage
|
|
|
|
os.makedirs(cache_dir, exist_ok=True)
|
|
manifest_path = os.path.join(cache_dir, "manifest.json")
|
|
|
|
if os.path.exists(manifest_path):
|
|
print(f" Loading cached PNGs from {cache_dir}")
|
|
with open(manifest_path) as f:
|
|
manifest = json.load(f)
|
|
return (manifest["paths"], np.array(manifest["patient_ids"]),
|
|
np.array(manifest["z_indices"]))
|
|
|
|
if self.slices_rgb is None or self.patient_labels is None:
|
|
raise RuntimeError("Call load() or load_all_slices() first.")
|
|
|
|
print(f" Exporting {len(self.slices_rgb)} PNGs to {cache_dir} ...")
|
|
paths, pids, zs = [], [], []
|
|
for i, (sl, pid, fname) in enumerate(
|
|
zip(self.slices_rgb, self.patient_labels, self.filenames)):
|
|
out_path = os.path.join(cache_dir, f"{fname}.png")
|
|
PILImage.fromarray(sl).save(out_path)
|
|
paths.append(out_path)
|
|
pids.append(pid)
|
|
zs.append(int(fname.split("_slice")[-1]))
|
|
if (i + 1) % 500 == 0:
|
|
print(f" {i + 1}/{len(self.slices_rgb)} ...", flush=True)
|
|
|
|
manifest = {"paths": paths, "patient_ids": pids, "z_indices": zs}
|
|
with open(manifest_path, "w") as f:
|
|
json.dump(manifest, f)
|
|
print(f" Exported {len(paths)} PNGs")
|
|
return paths, np.array(pids), np.array(zs)
|
|
|
|
@property
|
|
def n_patients(self):
|
|
return len(np.unique(self.patient_labels))
|
|
|
|
def _slice_indices(self, n_z):
|
|
"""Slice indices within the central fraction (even or random)."""
|
|
half = self.central_fraction / 2.0
|
|
start = int(n_z * (0.5 - half))
|
|
end = int(n_z * (0.5 + half))
|
|
if self.random_slices:
|
|
rng = np.random.default_rng(self.seed)
|
|
return np.sort(rng.integers(start, end, size=self.n_slices_per_patient))
|
|
else:
|
|
return np.linspace(start, end - 1, self.n_slices_per_patient).astype(int)
|