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/cross_view_fusion.py from shareefch1413/ACL-LKNet: direct link, hf CLI and curl.
- Browser
- Download file 3.86 kB
-
https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/src/models/cross_view_fusion.py
- Command line
-
hf download hf://shareefch1413/ACL-LKNet/src/models/cross_view_fusion.py
-
curl -L -o cross_view_fusion.py https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/src/models/cross_view_fusion.py
3.86 kB
| """ | |
| Cross-View Attention Fusion for ACL-LKNet. | |
| Fuses plane-level representations from sagittal, coronal, and axial views | |
| using multi-head attention. This allows the model to learn cross-view | |
| dependencies rather than merely concatenating independent features. | |
| Supports two modes: | |
| - 'attention': Multi-head self-attention over 3 view embeddings | |
| - 'concat': Simple concatenation + FC (ablation baseline) | |
| """ | |
| import torch | |
| import torch.nn as nn | |
| class CrossViewAttentionFusion(nn.Module): | |
| """ | |
| Fuses 3 plane embeddings using multi-head self-attention. | |
| Input: 3 view embeddings, each of shape (B, D) | |
| Process: | |
| 1. Stack into sequence of length 3: (B, 3, D) | |
| 2. Add learnable view position embeddings | |
| 3. Multi-head self-attention | |
| 4. Mean-pool the 3 output embeddings → (B, D) | |
| """ | |
| def __init__( | |
| self, | |
| feature_dim: int, | |
| num_heads: int = 2, | |
| dropout: float = 0.1, | |
| ): | |
| super().__init__() | |
| self.feature_dim = feature_dim | |
| # Learnable position embeddings for each view | |
| # (sagittal=0, coronal=1, axial=2) | |
| self.view_pos_embed = nn.Parameter(torch.zeros(1, 3, feature_dim)) | |
| nn.init.trunc_normal_(self.view_pos_embed, std=0.02) | |
| # Pre-norm | |
| self.norm = nn.LayerNorm(feature_dim) | |
| # Multi-head self-attention | |
| self.mha = nn.MultiheadAttention( | |
| embed_dim=feature_dim, | |
| num_heads=num_heads, | |
| dropout=dropout, | |
| batch_first=True, | |
| ) | |
| # Post-attention FFN | |
| self.ffn = nn.Sequential( | |
| nn.LayerNorm(feature_dim), | |
| nn.Linear(feature_dim, feature_dim * 2), | |
| nn.GELU(), | |
| nn.Dropout(dropout), | |
| nn.Linear(feature_dim * 2, feature_dim), | |
| nn.Dropout(dropout), | |
| ) | |
| def forward( | |
| self, | |
| sagittal: torch.Tensor, | |
| coronal: torch.Tensor, | |
| axial: torch.Tensor, | |
| ) -> torch.Tensor: | |
| """ | |
| Args: | |
| sagittal: (B, D) sagittal plane embedding | |
| coronal: (B, D) coronal plane embedding | |
| axial: (B, D) axial plane embedding | |
| Returns: | |
| fused: (B, D) fused representation | |
| """ | |
| # Stack into sequence: (B, 3, D) | |
| x = torch.stack([sagittal, coronal, axial], dim=1) | |
| # Add view position embeddings | |
| x = x + self.view_pos_embed | |
| # Self-attention with residual | |
| x_norm = self.norm(x) | |
| attn_out, _ = self.mha(x_norm, x_norm, x_norm) | |
| x = x + attn_out | |
| # FFN with residual | |
| x = x + self.ffn(x) | |
| # Mean-pool over the 3 views | |
| fused = x.mean(dim=1) # (B, D) | |
| return fused | |
| class ConcatFusion(nn.Module): | |
| """ | |
| Simple concatenation + FC fusion (ablation baseline). | |
| Concatenates 3 view embeddings → FC → output embedding. | |
| """ | |
| def __init__(self, feature_dim: int, dropout: float = 0.1): | |
| super().__init__() | |
| self.fusion = nn.Sequential( | |
| nn.Linear(feature_dim * 3, feature_dim), | |
| nn.GELU(), | |
| nn.Dropout(dropout), | |
| ) | |
| def forward( | |
| self, | |
| sagittal: torch.Tensor, | |
| coronal: torch.Tensor, | |
| axial: torch.Tensor, | |
| ) -> torch.Tensor: | |
| x = torch.cat([sagittal, coronal, axial], dim=-1) # (B, 3D) | |
| return self.fusion(x) # (B, D) | |
| def create_fusion( | |
| fusion_type: str, | |
| feature_dim: int, | |
| num_heads: int = 2, | |
| dropout: float = 0.1, | |
| ) -> nn.Module: | |
| """Factory for fusion modules.""" | |
| if fusion_type == "attention": | |
| return CrossViewAttentionFusion(feature_dim, num_heads, dropout) | |
| elif fusion_type == "concat": | |
| return ConcatFusion(feature_dim, dropout) | |
| else: | |
| raise ValueError(f"Unknown fusion type: {fusion_type}") | |