SACFlow / sacflow /data /nifti_dataset.py
sathiiii's picture
Add SACFlow source code
d65ae7d verified
Raw
History Blame Contribute Delete
5.93 kB
from __future__ import annotations
from pathlib import Path
from typing import Dict, List, Optional, Tuple
import json
import random
import numpy as np
import nibabel as nib
import torch
from torch.utils.data import Dataset
def load_nifti(path: str | Path):
img = nib.load(str(path))
data = img.get_fdata(dtype=np.float32)
return data, img.affine, img.header
def normalize_intensity(x: np.ndarray, mode: str = "zscore_nonzero", clip=None) -> np.ndarray:
x = x.astype(np.float32)
if clip is not None:
lo, hi = clip
x = np.clip(x, lo, hi)
if mode == "zscore_nonzero":
mask = np.abs(x) > 1e-6
if mask.sum() > 10:
mu, sd = x[mask].mean(), x[mask].std()
else:
mu, sd = x.mean(), x.std()
x = (x - mu) / (sd + 1e-6)
elif mode == "zscore":
x = (x - x.mean()) / (x.std() + 1e-6)
elif mode == "minmax":
x = (x - x.min()) / (x.max() - x.min() + 1e-6)
elif mode in ("none", None):
pass
else:
raise ValueError(f"Unknown intensity mode: {mode}")
return x.astype(np.float32)
def _pad_to_shape(arr: np.ndarray, shape: Tuple[int, int, int], value=0):
pads = []
for dim, target in zip(arr.shape[-3:], shape):
total = max(0, target - dim)
before = total // 2
after = total - before
pads.append((before, after))
if arr.ndim == 4:
pad_width = [(0, 0)] + pads
else:
pad_width = pads
return np.pad(arr, pad_width, mode="constant", constant_values=value)
def random_crop_pair(img: np.ndarray, lab: Optional[np.ndarray], patch_size: Tuple[int, int, int], foreground_prob: float = 0.5):
# img: [C,H,W,D], lab: [H,W,D]
img = _pad_to_shape(img, patch_size, 0)
if lab is not None:
lab = _pad_to_shape(lab, patch_size, 0)
H, W, D = img.shape[-3:]
ph, pw, pd = patch_size
if lab is not None and random.random() < foreground_prob and (lab > 0).sum() > 0:
coords = np.argwhere(lab > 0)
center = coords[random.randrange(len(coords))]
starts = []
for c, dim, p in zip(center, (H, W, D), patch_size):
s = int(c) - p // 2
s = max(0, min(s, dim - p))
starts.append(s)
else:
starts = [random.randint(0, max(0, dim - p)) for dim, p in zip((H, W, D), patch_size)]
sh, sw, sd = starts
img_c = img[:, sh:sh+ph, sw:sw+pw, sd:sd+pd]
lab_c = None if lab is None else lab[sh:sh+ph, sw:sw+pw, sd:sd+pd]
return img_c, lab_c
def center_crop_pair(img: np.ndarray, lab: Optional[np.ndarray], patch_size: Tuple[int, int, int]):
img = _pad_to_shape(img, patch_size, 0)
if lab is not None:
lab = _pad_to_shape(lab, patch_size, 0)
H, W, D = img.shape[-3:]
ph, pw, pd = patch_size
sh, sw, sd = [(dim - p)//2 for dim, p in zip((H,W,D), patch_size)]
img_c = img[:, sh:sh+ph, sw:sw+pw, sd:sd+pd]
lab_c = None if lab is None else lab[sh:sh+ph, sw:sw+pw, sd:sd+pd]
return img_c, lab_c
def augment(img: np.ndarray, lab: Optional[np.ndarray], cfg: Dict):
if cfg.get("random_flip", False):
for ax in range(3):
if random.random() < 0.5:
img = np.flip(img, axis=ax+1).copy()
if lab is not None:
lab = np.flip(lab, axis=ax).copy()
shift = float(cfg.get("random_intensity_shift", 0.0) or 0.0)
scale = float(cfg.get("random_intensity_scale", 0.0) or 0.0)
if shift > 0:
img = img + np.random.uniform(-shift, shift)
if scale > 0:
img = img * np.random.uniform(1-scale, 1+scale)
return img, lab
class NiftiSegDataset(Dataset):
def __init__(self, manifest: str | Path, split: str, cfg: Dict, training: bool = True, require_label: bool = False):
self.manifest_path = Path(manifest)
with open(self.manifest_path, "r") as f:
man = json.load(f)
if split not in man:
raise KeyError(f"Split {split} not found in {manifest}")
self.items = man[split]
self.split = split
self.cfg = cfg
self.training = training
self.require_label = require_label
self.patch_size = tuple(cfg.get("patch_size", [96,96,96]))
self.intensity_cfg = cfg.get("intensity", {})
self.aug_cfg = cfg.get("augmentation", {})
def __len__(self):
return len(self.items)
def __getitem__(self, idx):
item = self.items[idx]
img, affine, header = load_nifti(item["image"])
if img.ndim == 3:
img = img[None]
elif img.ndim == 4:
# assume channels last if last dim small
if img.shape[-1] <= 8:
img = np.moveaxis(img, -1, 0)
else:
img = img[None, ..., 0]
img = normalize_intensity(img, self.intensity_cfg.get("normalize", "zscore_nonzero"), self.intensity_cfg.get("clip"))
lab = None
if item.get("label"):
lab, _, _ = load_nifti(item["label"])
lab = lab.astype(np.int64)
elif self.require_label:
raise ValueError(f"Item {idx} has no label in split {self.split}")
if self.training:
img, lab = random_crop_pair(img, lab, self.patch_size)
img, lab = augment(img, lab, self.aug_cfg)
else:
# keep full volume for sliding window eval; no crop
pass
batch = {
"image": torch.from_numpy(img.copy()).float(),
"case_id": item.get("id", Path(item["image"]).stem),
"image_path": item["image"],
}
if lab is not None:
batch["label"] = torch.from_numpy(lab.copy()).long()
batch["spacing"] = torch.tensor(header.get_zooms()[:3] if header is not None else (1,1,1), dtype=torch.float32)
return batch
def load_manifest(path: str | Path):
with open(path, "r") as f:
return json.load(f)