ACL-LKNet / src /dataset.py
shareefch1413's picture
Upload folder using huggingface_hub
00801a0 verified
Raw History Blame Contribute Delete
15.1 kB
"""
MRNet dataset loading and augmentation for ACL-LKNet.
Handles:
- Loading .npy MRI volumes from MRNet directory structure
- Anatomically justified augmentation pipeline
- Variable slice count handling (subsampling + padding)
- Class-imbalanced sampling
- SSL variant (ignores labels, returns all views for masking)
MRNet directory structure:
mrnet/
β”œβ”€β”€ train/
β”‚ β”œβ”€β”€ sagittal/ (0000.npy, 0001.npy, ...)
β”‚ β”œβ”€β”€ coronal/
β”‚ └── axial/
β”œβ”€β”€ valid/
β”‚ β”œβ”€β”€ sagittal/
β”‚ β”œβ”€β”€ coronal/
β”‚ └── axial/
β”œβ”€β”€ train-acl.csv
└── valid-acl.csv
"""
import os
from typing import Tuple, Optional, Dict, List
import numpy as np
import pandas as pd
import torch
from torch.utils.data import Dataset, DataLoader, WeightedRandomSampler
import torchvision.transforms as T
import torchvision.transforms.functional as TF
# ── Augmentation ────────────────────────────────────────────────────
class MRIAugmentation:
"""
Anatomically justified augmentation for knee MRI slices.
Every augmentation has an explicit anatomical justification:
- Rotation(Β±10Β°): patient positioning variation
- Affine(translate=0.05): minor FOV variation
- Brightness/Contrast(Β±0.15): scanner intensity variation
- GaussianBlur: resolution / slight motion simulation
- RandomErasing: small artifact simulation
EXCLUDED by default:
- HorizontalFlip: changes L/R anatomy (must be an explicit experimental decision)
- VerticalFlip: anatomically meaningless
"""
def __init__(self, config, is_train: bool = True):
self.is_train = is_train
self.img_size = config.img_size
if is_train:
self.transform = T.Compose([
T.ToPILImage(),
T.Resize((config.img_size, config.img_size)),
T.RandomRotation(degrees=config.rotation_degrees),
T.RandomAffine(
degrees=0,
translate=(config.translate_range, config.translate_range),
scale=config.scale_range,
),
T.ColorJitter(
brightness=config.brightness,
contrast=config.contrast,
),
T.RandomApply(
[T.GaussianBlur(kernel_size=3, sigma=(0.1, 1.0))],
p=config.gaussian_blur_p,
),
T.ToTensor(),
T.RandomErasing(p=config.random_erasing_p, scale=(0.02, 0.1)),
])
# Horizontal flip as SEPARATE experimental decision
self.use_hflip = config.use_horizontal_flip
else:
self.transform = T.Compose([
T.ToPILImage(),
T.Resize((config.img_size, config.img_size)),
T.ToTensor(),
])
self.use_hflip = False
def __call__(self, image: np.ndarray) -> torch.Tensor:
"""
Args:
image: (H, W) numpy array, single MRI slice
Returns:
tensor: (1, img_size, img_size) normalized tensor
"""
# Ensure uint8 for PIL
if image.dtype != np.uint8:
# Normalize to 0-255
img_min, img_max = image.min(), image.max()
if img_max > img_min:
image = ((image - img_min) / (img_max - img_min) * 255).astype(np.uint8)
else:
image = np.zeros_like(image, dtype=np.uint8)
tensor = self.transform(image) # (1, H, W)
if self.use_hflip and self.is_train and torch.rand(1).item() < 0.5:
tensor = TF.hflip(tensor)
return tensor
# ── Dataset ─────────────────────────────────────────────────────────
class MRNetDataset(Dataset):
"""
MRNet Dataset for ACL tear detection.
Each sample returns:
- 3 views (sagittal, coronal, axial) as tensors
- 3 masks indicating valid (non-padded) slices
- Label (0 or 1)
- Case ID
"""
def __init__(
self,
data_dir: str,
split: str = "train",
config=None,
task: str = "acl",
case_list: Optional[List[str]] = None,
):
"""
Args:
data_dir: Root directory of MRNet data
split: 'train' or 'valid'
config: Config object with augmentation params
task: Label type ('acl', 'meniscus', 'abnormal')
case_list: Optional explicit list of case IDs (for K-fold cross-validation)
"""
self.data_dir = data_dir
self.split = split
self.max_slices = config.max_slices if config else 24
self.planes = ["sagittal", "coronal", "axial"]
# Load labels
label_file = os.path.join(data_dir, f"{split}-{task}.csv")
if os.path.exists(label_file):
self.labels_df = pd.read_csv(
label_file, header=None, names=["case", "label"]
)
else:
raise FileNotFoundError(
f"Label file not found: {label_file}\n"
f"Expected MRNet structure at: {data_dir}"
)
# Cache labels as dict for fast lookup
self.label_dict = dict(
zip(
self.labels_df["case"].astype(str).str.zfill(4),
self.labels_df["label"],
)
)
# Get case IDs from file system or explicit list
if case_list is not None:
self.cases = [str(c).zfill(4) if str(c).isdigit() else str(c) for c in case_list]
else:
sagittal_dir = os.path.join(data_dir, split, "sagittal")
if os.path.exists(sagittal_dir):
self.cases = sorted([
f.replace(".npy", "")
for f in os.listdir(sagittal_dir)
if f.endswith(".npy")
])
else:
raise FileNotFoundError(f"Data directory not found: {sagittal_dir}")
# Augmentation
is_train = split == "train"
self.augmentation = MRIAugmentation(config, is_train=is_train) if config else None
def __len__(self) -> int:
return len(self.cases)
def _load_volume(self, case_id: str, plane: str) -> np.ndarray:
"""Load a single MRI volume."""
path = os.path.join(self.data_dir, self.split, plane, f"{case_id}.npy")
volume = np.load(path) # (S, H, W)
return volume
def _subsample_slices(self, volume: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
"""
Subsample or pad volume to fixed number of slices.
Returns:
slices: (max_slices, H, W)
mask: (max_slices,) β€” True for real slices, False for padding
"""
S = volume.shape[0]
if S >= self.max_slices:
# Uniformly subsample
indices = np.linspace(0, S - 1, self.max_slices, dtype=int)
slices = volume[indices]
mask = np.ones(self.max_slices, dtype=bool)
else:
# Pad with zeros
pad_size = self.max_slices - S
slices = np.pad(volume, ((0, pad_size), (0, 0), (0, 0)), mode="constant")
mask = np.zeros(self.max_slices, dtype=bool)
mask[:S] = True
return slices, mask
def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]:
case_id = self.cases[idx]
# Load label
label_key = str(case_id).zfill(4)
# Try numeric lookup if string doesn't match
if label_key in self.label_dict:
label = self.label_dict[label_key]
else:
# Fallback: try matching by index
label = self.labels_df.iloc[idx]["label"]
label = torch.tensor(float(label), dtype=torch.float32)
views = {}
masks = {}
for plane in self.planes:
volume = self._load_volume(case_id, plane)
slices, mask = self._subsample_slices(volume)
# Apply augmentation to each slice
if self.augmentation:
processed = []
for s in range(slices.shape[0]):
if mask[s]:
processed.append(self.augmentation(slices[s]))
else:
# Padding slice β€” just resize and tensorize
processed.append(torch.zeros(1, self.augmentation.img_size, self.augmentation.img_size))
slices_tensor = torch.cat(processed, dim=0) # (max_slices, H, W)
else:
slices_tensor = torch.from_numpy(slices).float()
views[plane] = slices_tensor
masks[plane] = torch.from_numpy(mask)
return {
"sagittal": views["sagittal"],
"coronal": views["coronal"],
"axial": views["axial"],
"sag_mask": masks["sagittal"],
"cor_mask": masks["coronal"],
"axi_mask": masks["axial"],
"label": label,
"case_id": case_id,
}
def get_labels(self) -> List[int]:
"""Return all labels (for weighted sampler)."""
labels = []
for case_id in self.cases:
label_key = str(case_id).zfill(4)
if label_key in self.label_dict:
labels.append(int(self.label_dict[label_key]))
else:
idx = self.cases.index(case_id)
labels.append(int(self.labels_df.iloc[idx]["label"]))
return labels
class MRNetSSLDataset(MRNetDataset):
"""
MRNet dataset variant for self-supervised pretraining.
Same as MRNetDataset but:
- Uses ALL data (no label filtering)
- Returns views without labels
- Minimal augmentation (we want stable features for reconstruction)
"""
def __init__(self, data_dir: str, split: str = "train", config=None, task: str = "acl"):
super().__init__(data_dir, split, config, task)
# Override augmentation to be minimal for SSL
if config:
# Only resize + normalize for SSL
self.augmentation = MRIAugmentation(config, is_train=False)
# ── Stratified K-Fold Cross-Validation ─────────────────────────────
def get_stratified_folds(
data_dir: str,
split: str = "train",
n_splits: int = 5,
seed: int = 42,
task: str = "acl",
) -> List[Dict[str, List[str]]]:
"""
Generate patient-stratified K-fold train/validation splits preserving class balance.
Args:
data_dir: Root directory of MRNet data
split: Split to partition ('train' or 'valid')
n_splits: Number of cross-validation folds (default: 5)
seed: Random seed for reproducibility
task: Diagnostic task label ('acl', 'meniscus', 'abnormal')
Returns:
List of dicts: [{'fold': i, 'train_cases': [...], 'val_cases': [...]}, ...]
"""
from sklearn.model_selection import StratifiedKFold
label_file = os.path.join(data_dir, f"{split}-{task}.csv")
if not os.path.exists(label_file):
raise FileNotFoundError(f"Label file not found: {label_file}")
df = pd.read_csv(label_file, header=None, names=["case", "label"])
sagittal_dir = os.path.join(data_dir, split, "sagittal")
disk_cases = set([
f.replace(".npy", "")
for f in os.listdir(sagittal_dir)
if f.endswith(".npy")
])
# Filter to cases existing on disk
valid_cases = []
labels = []
for _, row in df.iterrows():
cid_raw = str(row["case"])
cid_pad = cid_raw.zfill(4)
if cid_pad in disk_cases:
valid_cases.append(cid_pad)
labels.append(int(row["label"]))
elif cid_raw in disk_cases:
valid_cases.append(cid_raw)
labels.append(int(row["label"]))
valid_cases = np.array(valid_cases)
labels = np.array(labels)
skf = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=seed)
folds = []
for fold_idx, (train_idx, val_idx) in enumerate(skf.split(valid_cases, labels)):
folds.append({
"fold": fold_idx,
"train_cases": valid_cases[train_idx].tolist(),
"val_cases": valid_cases[val_idx].tolist(),
"train_pos_rate": float(np.mean(labels[train_idx])),
"val_pos_rate": float(np.mean(labels[val_idx])),
})
return folds
# ── DataLoader Creation ────────────────────────────────────────────
def create_dataloaders(
config,
ssl: bool = False,
train_cases: Optional[List[str]] = None,
val_cases: Optional[List[str]] = None,
train_split: str = "train",
val_split: str = "valid",
) -> Tuple[DataLoader, DataLoader]:
"""
Create train and validation dataloaders.
Args:
config: Config object
ssl: If True, create SSL-mode datasets (no labels, minimal augmentation)
train_cases: Optional list of case IDs for training (e.g. for cross-validation)
val_cases: Optional list of case IDs for validation (e.g. for cross-validation)
train_split: Data subfolder name for training set
val_split: Data subfolder name for validation set
Returns:
train_loader, val_loader
"""
DatasetClass = MRNetSSLDataset if ssl else MRNetDataset
train_dataset = DatasetClass(
data_dir=config.data_dir,
split=train_split,
config=config,
case_list=train_cases,
)
val_dataset = DatasetClass(
data_dir=config.data_dir,
split=val_split,
config=config,
case_list=val_cases,
)
# Weighted sampler for class imbalance (supervised only)
train_sampler = None
shuffle = True
if not ssl:
labels = train_dataset.get_labels()
class_counts = np.bincount(labels)
if len(class_counts) == 2 and class_counts[1] > 0:
weights = 1.0 / class_counts.astype(float)
sample_weights = [weights[l] for l in labels]
train_sampler = WeightedRandomSampler(
sample_weights, len(sample_weights), replacement=True
)
shuffle = False # Sampler handles shuffling
train_loader = DataLoader(
train_dataset,
batch_size=config.batch_size,
shuffle=shuffle if train_sampler is None else False,
sampler=train_sampler,
num_workers=config.num_workers,
pin_memory=config.pin_memory,
drop_last=False,
)
val_loader = DataLoader(
val_dataset,
batch_size=config.batch_size,
shuffle=False,
num_workers=config.num_workers,
pin_memory=config.pin_memory,
)
return train_loader, val_loader