""" 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