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/models/msm.py from shareefch1413/ACL-LKNet: direct link, hf CLI and curl.
- Browser
- Download file 7.79 kB
-
https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/src/models/msm.py
- Command line
-
hf download hf://shareefch1413/ACL-LKNet/src/models/msm.py
-
curl -L -o msm.py https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/src/models/msm.py
7.79 kB
| """ | |
| Masked Slice Modeling (MSM) — Self-Supervised Pretraining for ACL-LKNet. | |
| This is the PRIMARY RESEARCH CONTRIBUTION. | |
| Core idea: MRI volumes have ordered slices with anatomical continuity. | |
| We exploit this structure by masking some slices and training the model | |
| to reconstruct their features from the remaining (unmasked) context. | |
| Research question: | |
| Can a self-supervised objective that explicitly models inter-slice | |
| anatomical context produce better transferable representations for | |
| ACL injury detection than established pretraining? | |
| Masking strategies (research axis): | |
| - random: Mask 50% of slices uniformly at random | |
| - contiguous: Mask contiguous blocks of 3-5 adjacent slices | |
| - structured: Preferentially mask central slices (clinically relevant) | |
| - mixed: Alternate random and contiguous per batch | |
| """ | |
| import math | |
| import random as py_random | |
| from typing import Tuple, Optional | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| class MaskedSliceModeling(nn.Module): | |
| """ | |
| Self-supervised pretraining via Masked Slice Modeling. | |
| Architecture: | |
| 1. Encode all slices through the shared backbone → slice features | |
| 2. Replace masked slice features with learnable [MASK] tokens | |
| 3. Add positional encoding (sinusoidal — respects slice ordering) | |
| 4. Pass through lightweight Transformer decoder | |
| 5. Predict the original features of masked slices | |
| Loss: MSE between predicted and actual features of masked slices | |
| """ | |
| def __init__( | |
| self, | |
| feature_dim: int, | |
| decoder_dim: int = 256, | |
| decoder_layers: int = 2, | |
| decoder_heads: int = 4, | |
| max_slices: int = 48, | |
| mask_ratio: float = 0.5, | |
| mask_strategy: str = "random", | |
| ): | |
| super().__init__() | |
| self.feature_dim = feature_dim | |
| self.decoder_dim = decoder_dim | |
| self.mask_ratio = mask_ratio | |
| self.mask_strategy = mask_strategy | |
| self.max_slices = max_slices | |
| # Learnable [MASK] token | |
| self.mask_token = nn.Parameter(torch.zeros(1, 1, feature_dim)) | |
| nn.init.trunc_normal_(self.mask_token, std=0.02) | |
| # Project encoder features → decoder dimension | |
| self.encoder_to_decoder = nn.Linear(feature_dim, decoder_dim) | |
| # Sinusoidal positional encoding (respects spatial ordering of slices) | |
| self.register_buffer( | |
| "pos_encoding", self._sinusoidal_encoding(max_slices, decoder_dim) | |
| ) | |
| # Lightweight Transformer decoder | |
| decoder_layer = nn.TransformerEncoderLayer( | |
| d_model=decoder_dim, | |
| nhead=decoder_heads, | |
| dim_feedforward=decoder_dim * 4, | |
| dropout=0.1, | |
| activation="gelu", | |
| batch_first=True, | |
| norm_first=True, | |
| ) | |
| self.decoder = nn.TransformerEncoder( | |
| decoder_layer, | |
| num_layers=decoder_layers, | |
| ) | |
| # Predict original features from decoded representations | |
| self.predictor = nn.Sequential( | |
| nn.LayerNorm(decoder_dim), | |
| nn.Linear(decoder_dim, feature_dim), | |
| ) | |
| def _sinusoidal_encoding(max_len: int, dim: int) -> torch.Tensor: | |
| """Generate sinusoidal positional encoding.""" | |
| pe = torch.zeros(max_len, dim) | |
| position = torch.arange(0, max_len).unsqueeze(1).float() | |
| div_term = torch.exp( | |
| torch.arange(0, dim, 2).float() * (-math.log(10000.0) / dim) | |
| ) | |
| pe[:, 0::2] = torch.sin(position * div_term) | |
| pe[:, 1::2] = torch.cos(position * div_term) | |
| return pe.unsqueeze(0) # (1, max_len, dim) | |
| def generate_mask( | |
| self, num_slices: int, strategy: Optional[str] = None | |
| ) -> torch.Tensor: | |
| """ | |
| Generate a binary mask indicating which slices to mask. | |
| Args: | |
| num_slices: Number of slices in the volume | |
| strategy: Override the default masking strategy | |
| Returns: | |
| mask: (num_slices,) boolean tensor. True = masked (to predict) | |
| """ | |
| strategy = strategy or self.mask_strategy | |
| num_mask = max(1, int(num_slices * self.mask_ratio)) | |
| # Always keep at least 2 slices unmasked for context | |
| num_mask = min(num_mask, num_slices - 2) | |
| mask = torch.zeros(num_slices, dtype=torch.bool) | |
| if strategy == "random": | |
| indices = torch.randperm(num_slices)[:num_mask] | |
| mask[indices] = True | |
| elif strategy == "contiguous": | |
| # Mask contiguous blocks of 3-5 slices | |
| remaining = num_mask | |
| while remaining > 0: | |
| block_size = min(py_random.randint(3, 5), remaining) | |
| max_start = num_slices - block_size | |
| if max_start <= 0: | |
| start = 0 | |
| else: | |
| start = py_random.randint(0, max_start) | |
| mask[start : start + block_size] = True | |
| remaining = num_mask - mask.sum().item() | |
| elif strategy == "structured": | |
| # Preferentially mask central slices (where ACL is typically visible) | |
| center = num_slices // 2 | |
| # Create probability distribution peaked at center | |
| positions = torch.arange(num_slices).float() | |
| probs = torch.exp(-0.5 * ((positions - center) / (num_slices / 4)) ** 2) | |
| probs = probs / probs.sum() | |
| indices = torch.multinomial(probs, num_mask, replacement=False) | |
| mask[indices] = True | |
| elif strategy == "mixed": | |
| # Randomly choose between random and contiguous per call | |
| sub_strategy = py_random.choice(["random", "contiguous"]) | |
| mask = self.generate_mask(num_slices, strategy=sub_strategy) | |
| else: | |
| raise ValueError(f"Unknown mask strategy: {strategy}") | |
| return mask | |
| def forward( | |
| self, | |
| slice_features: torch.Tensor, | |
| slice_mask: Optional[torch.Tensor] = None, | |
| ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: | |
| """ | |
| Forward pass for MSM pretraining. | |
| Args: | |
| slice_features: (B, S, D) — encoded slice features from backbone | |
| slice_mask: (B, S) — optional padding mask (True = valid) | |
| Returns: | |
| loss: scalar MSE loss on masked slices | |
| predictions: (B, S, D) — predicted features for all slices | |
| mask: (B, S) — boolean mask of which slices were masked | |
| """ | |
| B, S, D = slice_features.shape | |
| # Generate masks for each sample in the batch | |
| masks = torch.stack([self.generate_mask(S) for _ in range(B)]) # (B, S) | |
| masks = masks.to(slice_features.device) | |
| # Replace masked positions with [MASK] token | |
| mask_tokens = self.mask_token.expand(B, S, -1) # (B, S, D) | |
| masked_features = slice_features.clone() | |
| masked_features[masks] = mask_tokens[masks] | |
| # Project to decoder dimension | |
| x = self.encoder_to_decoder(masked_features) # (B, S, decoder_dim) | |
| # Add positional encoding | |
| x = x + self.pos_encoding[:, :S, :] | |
| # Transformer decoder | |
| x = self.decoder(x) # (B, S, decoder_dim) | |
| # Predict original features | |
| predictions = self.predictor(x) # (B, S, D) | |
| # Compute loss only on masked positions | |
| if masks.any(): | |
| pred_masked = predictions[masks] # (num_masked, D) | |
| target_masked = slice_features[masks] # (num_masked, D) | |
| loss = F.mse_loss(pred_masked, target_masked) | |
| else: | |
| loss = torch.tensor(0.0, device=slice_features.device) | |
| return loss, predictions, masks | |