Image Classification
timm
English
medical-imaging
knee-mri
acl-tear-detection
deep-learning
convnext
self-attention
masked-slice-modeling
radiology
orthopedics
Eval Results (legacy)
Instructions to use shareefch1413/ACL-LKNet with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- timm
How to use shareefch1413/ACL-LKNet with timm:
import timm model = timm.create_model("hf-hub:shareefch1413/ACL-LKNet", pretrained=True) - Notebooks
- Google Colab
- Kaggle
Download src/dataset.py from shareefch1413/ACL-LKNet: direct link, hf CLI and curl.
- Browser
- Download file 15.1 kB
-
https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/src/dataset.py
- Command line
-
hf download hf://shareefch1413/ACL-LKNet/src/dataset.py
-
curl -L -o dataset.py https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/src/dataset.py
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 | |