VLbai-2.6AD / dataset.py
eyupipler's picture
Upload 21 files
1013007 verified
Raw
History Blame Contribute Delete
11.7 kB
"""
Vbai-2.6AD Datasets
===================
Reads the cached visit manifest at _cache/paired_visits.parquet, which your own
data-preparation step must produce. One row per visit: a volume path, the 13
biomarker columns with their per-feature masks, the label and the progression
fields.
Three dataset modes:
* mode="mri" → MRI + label (Phase 1 pretrain)
* mode="tab" → tabular features + label (Phase 2 pretrain)
* mode="multi" → MRI + tabular + label + progression (Phase 3 fusion)
Tabular feature contract: 2 * NUM_FEATURES floats per sample.
[normalized values..., missing-mask bits...]
A feature with missing-mask=0 has its value zeroed (after normalization).
"""
from __future__ import annotations
import os
import random
import numpy as np
import pandas as pd
import torch
from torch.utils.data import Dataset, DataLoader, WeightedRandomSampler
from scipy.ndimage import zoom, rotate
import config as C
try:
import nibabel as nib
HAS_NIBABEL = True
except Exception:
HAS_NIBABEL = False
# Optional pre-decoded .npy cache, searched in order. Decoding NIfTI is the
# slowest part of an epoch, so a local-disk cache pays for itself quickly.
# YOU MUST SET YOUR OWN PATH: point VBAI_NPY_CACHE at a fast local directory,
# or leave it unset to use the in-project cache.
_NPY_CACHE_DIRS = [d for d in [
os.environ.get("VBAI_NPY_CACHE"),
os.path.join(C.PROJECT_ROOT, "_cache", "volume_npy"),
] if d]
def _try_load_cached(image_id: str) -> np.ndarray | None:
if not image_id:
return None
for d in _NPY_CACHE_DIRS:
p = os.path.join(d, f"{image_id}.npy")
if os.path.exists(p):
try:
return np.load(p).astype(np.float32)
except Exception:
return None
return None
# ---------- MRI loading + augmentation ----------
def _hippocampus_crop(data: np.ndarray, dcfg: C.DataConfig) -> np.ndarray:
"""Find the brain bounding box, then crop to a hippocampus-focused sub-region."""
mask = data > 0
if not mask.any():
return data
coords = np.argwhere(mask)
mn = coords.min(axis=0); mx = coords.max(axis=0)
size = mx - mn + 1
rx, ry, rz = dcfg.hippo_x_range, dcfg.hippo_y_range, dcfg.hippo_z_range
x0, x1 = int(mn[0] + size[0] * rx[0]), int(mn[0] + size[0] * rx[1])
y0, y1 = int(mn[1] + size[1] * ry[0]), int(mn[1] + size[1] * ry[1])
z0, z1 = int(mn[2] + size[2] * rz[0]), int(mn[2] + size[2] * rz[1])
cropped = data[x0:x1+1, y0:y1+1, z0:z1+1]
return cropped
def _load_nifti(path: str, target_shape=(96, 96, 96),
hippocampus_crop: bool = False, dcfg: C.DataConfig = None) -> np.ndarray:
img = nib.load(path)
data = img.get_fdata().astype(np.float32)
if data.ndim == 4:
data = data[..., 0]
mask = data > 0
if mask.sum() > 0:
vals = data[mask]
lo, hi = np.percentile(vals, [1.0, 99.0])
data = np.clip(data, lo, hi)
m, s = vals.mean(), vals.std()
if s > 0:
data = (data - m) / s
data[~mask] = 0
# Hippocampus-focused crop: narrows the content, which raises effective resolution
if hippocampus_crop:
data = _hippocampus_crop(data, dcfg or C.DataConfig())
if data.shape != target_shape:
f = [t / s for t, s in zip(target_shape, data.shape)]
data = zoom(data, f, order=1)
return data.astype(np.float32)
class MRIAugment3D:
def __init__(self, dcfg: C.DataConfig):
self.cfg = dcfg
def __call__(self, vol: np.ndarray) -> np.ndarray:
if random.random() < 0.5:
angle = random.uniform(-self.cfg.aug_rotation_range, self.cfg.aug_rotation_range)
axes = random.choice([(0, 1), (0, 2), (1, 2)])
vol = rotate(vol, angle, axes=axes, reshape=False, order=1, mode="nearest")
for ax in range(3):
if random.random() < self.cfg.aug_flip_prob:
vol = np.flip(vol, axis=ax).copy()
if random.random() < 0.3:
vol = vol + np.random.normal(0, self.cfg.aug_noise_std, vol.shape).astype(np.float32)
if random.random() < 0.3:
g = random.uniform(*self.cfg.aug_gamma_range)
mn = vol.min(); rg = vol.max() - mn
if rg > 0:
vol = ((vol - mn) / rg) ** g * rg + mn
return vol.astype(np.float32)
# ---------- Tabular normalization ----------
class TabularNormalizer:
"""Robust z-score on observed (non-missing) values per feature, fit on training set."""
def __init__(self):
self.mean: np.ndarray | None = None
self.std: np.ndarray | None = None
def fit(self, df: pd.DataFrame):
means, stds = [], []
for f in C.FEATURE_NAMES:
v = pd.to_numeric(df[f], errors="coerce").dropna().values.astype(np.float64)
if len(v) > 1:
m = float(np.median(v))
s = float(np.median(np.abs(v - m)) * 1.4826) # MAD → std
if s < 1e-8:
s = float(v.std()) if v.std() > 1e-8 else 1.0
else:
m, s = 0.0, 1.0
means.append(m); stds.append(s)
self.mean = np.asarray(means, dtype=np.float32)
self.std = np.asarray(stds, dtype=np.float32)
def transform(self, values: np.ndarray, mask: np.ndarray) -> np.ndarray:
z = (values - self.mean) / self.std
z = np.where(mask > 0.5, z, 0.0) # zero out missing
return np.concatenate([z, mask.astype(np.float32)], axis=-1)
def state_dict(self):
return {"mean": self.mean.tolist() if self.mean is not None else None,
"std": self.std.tolist() if self.std is not None else None}
def load_state_dict(self, sd):
self.mean = np.asarray(sd["mean"], dtype=np.float32)
self.std = np.asarray(sd["std"], dtype=np.float32)
# ---------- Subject-level split (no leakage between train/val/test) ----------
def subject_split(df: pd.DataFrame, val_frac=0.15, test_frac=0.15, seed=42):
rng = np.random.RandomState(seed)
ptids = np.array(sorted(df["ptid"].unique()))
rng.shuffle(ptids)
n = len(ptids)
n_test = int(round(n * test_frac))
n_val = int(round(n * val_frac))
test_ids = set(ptids[:n_test])
val_ids = set(ptids[n_test:n_test + n_val])
train_ids = set(ptids[n_test + n_val:])
print(f"[split] subjects → train {len(train_ids)} / val {len(val_ids)} / test {len(test_ids)}")
return train_ids, val_ids, test_ids
# ---------- Core paired dataset ----------
class PairedVisitDataset(Dataset):
"""
One sample = one MRI scan with paired biomarkers + (optional) progression labels.
Setting mode controls which fields are loaded:
"mri" — only mri + label (skips biomarker columns)
"tab" — only biomarkers + label (skips MRI loading)
"multi" — both
"""
def __init__(self, df: pd.DataFrame, normalizer: TabularNormalizer,
mode: str = "multi", augment: bool = False,
dcfg: C.DataConfig = None, mcfg: C.ModelConfig = None,
train_modality_dropout: bool = False):
self.df = df.reset_index(drop=True).copy()
self.norm = normalizer
self.mode = mode
self.augment = augment
self.dcfg = dcfg or C.DataConfig()
self.mcfg = mcfg or C.ModelConfig()
self.augmenter = MRIAugment3D(self.dcfg) if augment else None
self.modality_dropout = train_modality_dropout
def __len__(self):
return len(self.df)
def _get_tab(self, row, training: bool):
vals = np.array([row[f] for f in C.FEATURE_NAMES], dtype=np.float32)
mask = np.array([row[f"feat_mask_{f}"] for f in C.FEATURE_NAMES], dtype=np.float32)
# NaN safety
vals = np.where(np.isnan(vals), 0.0, vals)
# Stochastic feature masking during training (simulate missing inputs)
if training and self.mcfg.p_feature_mask > 0:
drop = np.random.rand(len(C.FEATURE_NAMES)) < self.mcfg.p_feature_mask
mask = np.where(drop, 0.0, mask)
return self.norm.transform(vals, mask).astype(np.float32)
def _get_mri(self, row):
# Fast path: pre-decoded .npy on local disk
vol = _try_load_cached(row.get("image_id"))
if vol is None:
vol = _load_nifti(row["nifti_path"], self.dcfg.nifti_target_shape)
if self.augment and self.augmenter:
vol = self.augmenter(vol)
return torch.from_numpy(np.ascontiguousarray(vol)).unsqueeze(0).float()
def __getitem__(self, idx):
row = self.df.iloc[idx]
out = {
"label": torch.tensor(int(row["label"]), dtype=torch.long),
"has_progression": torch.tensor(bool(row["has_progression"]), dtype=torch.bool),
"will_progress": torch.tensor(float(row["will_progress"]), dtype=torch.float32),
"progression_months": torch.tensor(float(row["months_to_conversion"]), dtype=torch.float32),
"ptid": str(row["ptid"]),
}
load_mri = self.mode in ("mri", "multi")
load_tab = self.mode in ("tab", "multi")
# Modality dropout (Phase 3 only)
if self.modality_dropout and self.mode == "multi":
r = random.random()
if r < self.mcfg.p_drop_mri:
load_mri = False
elif r < self.mcfg.p_drop_mri + self.mcfg.p_drop_tab:
load_tab = False
if load_mri:
out["mri"] = self._get_mri(row)
if load_tab:
out["tab"] = torch.from_numpy(self._get_tab(row, training=self.augment))
out["has_mri"] = torch.tensor(load_mri, dtype=torch.bool)
out["has_tab"] = torch.tensor(load_tab, dtype=torch.bool)
return out
def collate_pad(batch):
"""Collate that handles optional mri/tab tensors per-sample."""
keys = ["label", "has_progression", "will_progress", "progression_months", "has_mri", "has_tab"]
out = {k: torch.stack([b[k] for b in batch]) for k in keys}
# MRI: only stack if all present (modality dropout makes mixed batches rare in practice;
# we drop unmatched samples to None at batch level to keep things simple)
if all("mri" in b for b in batch):
out["mri"] = torch.stack([b["mri"] for b in batch])
if all("tab" in b for b in batch):
out["tab"] = torch.stack([b["tab"] for b in batch])
out["ptid"] = [b["ptid"] for b in batch]
return out
# ---------- Helpers ----------
def get_class_weights(labels: np.ndarray, num_classes: int = 3) -> torch.Tensor:
counts = np.bincount(labels, minlength=num_classes).astype(np.float32)
counts[counts == 0] = 1.0
w = 1.0 / counts
w = w / w.sum() * num_classes
return torch.tensor(w, dtype=torch.float32)
def get_weighted_sampler(labels: np.ndarray) -> WeightedRandomSampler:
counts = np.bincount(labels)
sw = 1.0 / counts[labels]
return WeightedRandomSampler(torch.from_numpy(sw).float(), len(sw), replacement=True)
def load_paired() -> pd.DataFrame:
if not os.path.exists(C.PAIRED_PARQUET):
raise FileNotFoundError(
f"Visit manifest not found: {C.PAIRED_PARQUET}\n"
"Build it from your own data first. YOU MUST SET YOUR OWN PATHS "
"(see config.py)."
)
return pd.read_parquet(C.PAIRED_PARQUET)