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/acl_lknet.py from shareefch1413/ACL-LKNet: direct link, hf CLI and curl.
- Browser
- Download file 8.6 kB
-
https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/src/models/acl_lknet.py
- Command line
-
hf download hf://shareefch1413/ACL-LKNet/src/models/acl_lknet.py
-
curl -L -o acl_lknet.py https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/src/models/acl_lknet.py
8.6 kB
| """ | |
| ACL-LKNet: Full model assembly. | |
| Combines all components into the complete pipeline: | |
| Backbone β Slice Attention β Cross-View Fusion β Classifier | |
| Design: | |
| - Shared backbone processes slices from all 3 views (parameter-efficient) | |
| - Slices are processed in memory-efficient chunks (for T4 GPU) | |
| - Supports switching between attention/concat fusion and pool/attention aggregation | |
| """ | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from torch.utils.checkpoint import checkpoint as grad_checkpoint | |
| from .backbone import create_backbone | |
| from .slice_attention import SliceAttention | |
| from .cross_view_fusion import create_fusion | |
| class ACLLKNet(nn.Module): | |
| """ | |
| ACL-LKNet: Hierarchical Self-Supervised Large-Kernel Network | |
| for ACL Tear Detection in Knee MRI. | |
| Architecture: | |
| 1. Shared LKNet backbone extracts per-slice features | |
| 2. Slice attention aggregates variable-length slice sequences per view | |
| 3. Cross-view attention fuses the 3 plane embeddings | |
| 4. Classification head outputs ACL tear probability | |
| """ | |
| def __init__( | |
| self, | |
| backbone_name: str = "convnext_tiny", | |
| pretrained: bool = True, | |
| feature_dim: int = 0, | |
| attn_hidden_dim: int = 256, | |
| fusion_type: str = "attention", | |
| fusion_num_heads: int = 2, | |
| fusion_dropout: float = 0.1, | |
| classifier_hidden: int = 256, | |
| classifier_dropout: float = 0.3, | |
| aggregation: str = "attention", # 'attention', 'max', 'mean' | |
| use_grad_checkpoint: bool = True, | |
| slice_chunk_size: int = 8, | |
| use_large_kernels: bool = False, | |
| large_kernel_sizes: list = None, | |
| ): | |
| super().__init__() | |
| self.slice_chunk_size = slice_chunk_size | |
| self.aggregation_type = aggregation | |
| self.use_grad_checkpoint = use_grad_checkpoint | |
| # ββ Backbone ββ | |
| self.backbone = create_backbone( | |
| name=backbone_name, | |
| pretrained=pretrained, | |
| use_grad_checkpoint=use_grad_checkpoint, | |
| use_large_kernels=use_large_kernels, | |
| large_kernel_sizes=large_kernel_sizes, | |
| ) | |
| feat_dim = feature_dim if feature_dim > 0 else self.backbone.feature_dim | |
| # ββ Slice Attention (per view) ββ | |
| if aggregation == "attention": | |
| self.slice_attn = SliceAttention(feat_dim, attn_hidden_dim) | |
| else: | |
| self.slice_attn = None | |
| # ββ Cross-View Fusion ββ | |
| self.fusion = create_fusion( | |
| fusion_type=fusion_type, | |
| feature_dim=feat_dim, | |
| num_heads=fusion_num_heads, | |
| dropout=fusion_dropout, | |
| ) | |
| # ββ Classification Head ββ | |
| self.classifier = nn.Sequential( | |
| nn.LayerNorm(feat_dim), | |
| nn.Linear(feat_dim, classifier_hidden), | |
| nn.GELU(), | |
| nn.Dropout(classifier_dropout), | |
| nn.Linear(classifier_hidden, 1), | |
| ) | |
| self._feat_dim = feat_dim | |
| def feature_dim(self): | |
| return self._feat_dim | |
| def _extract_slice_features(self, slices: torch.Tensor) -> torch.Tensor: | |
| """ | |
| Extract features from a batch of slices, processing in memory-efficient chunks. | |
| Args: | |
| slices: (B, S, H, W) β grayscale MRI slices for ONE view | |
| Returns: | |
| features: (B, S, D) β per-slice feature vectors | |
| """ | |
| B, S, H, W = slices.shape | |
| all_features = [] | |
| for i in range(0, S, self.slice_chunk_size): | |
| chunk = slices[:, i : i + self.slice_chunk_size] # (B, chunk, H, W) | |
| chunk_size = chunk.shape[1] | |
| # Reshape: (B, chunk, H, W) β (B*chunk, 1, H, W) | |
| chunk = chunk.reshape(B * chunk_size, 1, H, W) | |
| # Forward through backbone | |
| feat = self.backbone(chunk) # (B*chunk, D) | |
| # Reshape back: (B, chunk, D) | |
| feat = feat.reshape(B, chunk_size, -1) | |
| all_features.append(feat) | |
| features = torch.cat(all_features, dim=1) # (B, S, D) | |
| return features | |
| def _aggregate_slices( | |
| self, features: torch.Tensor, mask: torch.Tensor = None | |
| ) -> tuple: | |
| """ | |
| Aggregate slice features into a single view embedding. | |
| Args: | |
| features: (B, S, D) | |
| mask: (B, S) β True for valid slices | |
| Returns: | |
| embedding: (B, D) | |
| attn_weights: (B, S) or None | |
| """ | |
| if self.aggregation_type == "attention" and self.slice_attn is not None: | |
| return self.slice_attn(features, mask) | |
| elif self.aggregation_type == "max": | |
| if mask is not None: | |
| features = features.masked_fill(~mask.unsqueeze(-1), float("-inf")) | |
| return features.max(dim=1)[0], None | |
| elif self.aggregation_type == "mean": | |
| if mask is not None: | |
| features = features * mask.unsqueeze(-1).float() | |
| return features.sum(dim=1) / mask.sum(dim=1, keepdim=True).float(), None | |
| return features.mean(dim=1), None | |
| else: | |
| raise ValueError(f"Unknown aggregation: {self.aggregation_type}") | |
| def forward( | |
| self, | |
| sagittal: torch.Tensor, | |
| coronal: torch.Tensor, | |
| axial: torch.Tensor, | |
| sag_mask: torch.Tensor = None, | |
| cor_mask: torch.Tensor = None, | |
| axi_mask: torch.Tensor = None, | |
| ) -> dict: | |
| """ | |
| Full forward pass. | |
| Args: | |
| sagittal: (B, S_sag, H, W) β sagittal MRI slices | |
| coronal: (B, S_cor, H, W) β coronal MRI slices | |
| axial: (B, S_axi, H, W) β axial MRI slices | |
| sag_mask: (B, S_sag) β optional padding mask | |
| cor_mask: (B, S_cor) β optional padding mask | |
| axi_mask: (B, S_axi) β optional padding mask | |
| Returns: | |
| dict with: | |
| 'logits': (B, 1) β raw logits | |
| 'probs': (B, 1) β sigmoid probabilities | |
| 'sag_weights': (B, S_sag) β slice attention weights | |
| 'cor_weights': (B, S_cor) β slice attention weights | |
| 'axi_weights': (B, S_axi) β slice attention weights | |
| """ | |
| # Extract slice features (shared backbone) | |
| sag_feats = self._extract_slice_features(sagittal) | |
| cor_feats = self._extract_slice_features(coronal) | |
| axi_feats = self._extract_slice_features(axial) | |
| # Aggregate slices β view embeddings | |
| sag_emb, sag_w = self._aggregate_slices(sag_feats, sag_mask) | |
| cor_emb, cor_w = self._aggregate_slices(cor_feats, cor_mask) | |
| axi_emb, axi_w = self._aggregate_slices(axi_feats, axi_mask) | |
| # Cross-view fusion | |
| fused = self.fusion(sag_emb, cor_emb, axi_emb) # (B, D) | |
| # Classification | |
| logits = self.classifier(fused) # (B, 1) | |
| probs = torch.sigmoid(logits) | |
| return { | |
| "logits": logits, | |
| "probs": probs, | |
| "sag_weights": sag_w, | |
| "cor_weights": cor_w, | |
| "axi_weights": axi_w, | |
| } | |
| def get_slice_features( | |
| self, | |
| sagittal: torch.Tensor, | |
| coronal: torch.Tensor, | |
| axial: torch.Tensor, | |
| ) -> dict: | |
| """ | |
| Extract slice features only (for MSM pretraining). | |
| Does not run attention/fusion/classifier. | |
| Returns: | |
| dict with 'sagittal', 'coronal', 'axial' β each (B, S, D) | |
| """ | |
| return { | |
| "sagittal": self._extract_slice_features(sagittal), | |
| "coronal": self._extract_slice_features(coronal), | |
| "axial": self._extract_slice_features(axial), | |
| } | |
| def create_model_from_config(config) -> ACLLKNet: | |
| """Create an ACLLKNet model from a Config object.""" | |
| return ACLLKNet( | |
| backbone_name=config.backbone, | |
| pretrained=config.pretrained, | |
| feature_dim=config.feature_dim, | |
| attn_hidden_dim=config.attn_hidden_dim, | |
| fusion_type=config.fusion_type, | |
| fusion_num_heads=config.fusion_num_heads, | |
| fusion_dropout=config.fusion_dropout, | |
| classifier_hidden=config.classifier_hidden, | |
| classifier_dropout=config.classifier_dropout, | |
| aggregation=getattr(config, "aggregation", "attention"), | |
| use_grad_checkpoint=config.grad_checkpoint, | |
| slice_chunk_size=config.slice_chunk_size, | |
| use_large_kernels=config.use_large_kernels, | |
| large_kernel_sizes=config.large_kernel_sizes, | |
| ) | |