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/slice_attention.py from shareefch1413/ACL-LKNet: direct link, hf CLI and curl.
- Browser
- Download file 3.78 kB
-
https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/src/models/slice_attention.py
- Command line
-
hf download hf://shareefch1413/ACL-LKNet/src/models/slice_attention.py
-
curl -L -o slice_attention.py https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/src/models/slice_attention.py
3.78 kB
| """ | |
| Slice Attention module for ACL-LKNet. | |
| Aggregates a variable-length sequence of slice features into a single | |
| exam-level representation using learned attention weights. | |
| Why attention over max-pool: | |
| - Max-pool discards information about WHICH slices are informative | |
| - Attention learns to weight slices showing pathology more heavily | |
| - Attention weights are interpretable (can visualize which slices matter) | |
| """ | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| class SliceAttention(nn.Module): | |
| """ | |
| Attention-based aggregation over a sequence of slice features. | |
| Given slice features [f_1, f_2, ..., f_S], computes: | |
| α_i = softmax(w^T · tanh(W_1 · f_i + b_1)) | |
| v = Σ α_i · f_i | |
| This is Bahdanau-style (additive) attention with a single attention head. | |
| """ | |
| def __init__(self, feature_dim: int, hidden_dim: int = 256): | |
| """ | |
| Args: | |
| feature_dim: Dimension of input slice features | |
| hidden_dim: Hidden dimension in attention computation | |
| """ | |
| super().__init__() | |
| self.attention = nn.Sequential( | |
| nn.Linear(feature_dim, hidden_dim), | |
| nn.Tanh(), | |
| nn.Linear(hidden_dim, 1), | |
| ) | |
| def forward( | |
| self, features: torch.Tensor, mask: torch.Tensor = None | |
| ) -> tuple: | |
| """ | |
| Args: | |
| features: (B, S, D) — batch of slice feature sequences | |
| mask: (B, S) — optional boolean mask (True = valid slice) | |
| Returns: | |
| aggregated: (B, D) — attention-weighted feature vector | |
| weights: (B, S) — attention weights (for visualization) | |
| """ | |
| # Compute attention scores | |
| scores = self.attention(features).squeeze(-1) # (B, S) | |
| # Mask invalid slices (from padding) | |
| if mask is not None: | |
| scores = scores.masked_fill(~mask, float("-inf")) | |
| # Softmax → attention weights | |
| weights = F.softmax(scores, dim=1) # (B, S) | |
| # Weighted sum | |
| aggregated = torch.bmm( | |
| weights.unsqueeze(1), features | |
| ).squeeze(1) # (B, D) | |
| return aggregated, weights | |
| class MultiHeadSliceAttention(nn.Module): | |
| """ | |
| Multi-head variant of slice attention for richer aggregation. | |
| Each head learns to attend to different aspects of the slices | |
| (e.g., one head for anatomy, another for pathology signal). | |
| Outputs are concatenated and projected. | |
| """ | |
| def __init__(self, feature_dim: int, num_heads: int = 4, hidden_dim: int = 256): | |
| super().__init__() | |
| assert feature_dim % num_heads == 0, "feature_dim must be divisible by num_heads" | |
| self.num_heads = num_heads | |
| self.head_dim = feature_dim // num_heads | |
| self.heads = nn.ModuleList([ | |
| SliceAttention(feature_dim, hidden_dim) | |
| for _ in range(num_heads) | |
| ]) | |
| self.projection = nn.Linear(feature_dim * num_heads, feature_dim) | |
| def forward( | |
| self, features: torch.Tensor, mask: torch.Tensor = None | |
| ) -> tuple: | |
| """ | |
| Args: | |
| features: (B, S, D) | |
| mask: (B, S) | |
| Returns: | |
| aggregated: (B, D) | |
| weights: (B, num_heads, S) — per-head attention weights | |
| """ | |
| head_outputs = [] | |
| all_weights = [] | |
| for head in self.heads: | |
| out, w = head(features, mask) | |
| head_outputs.append(out) | |
| all_weights.append(w) | |
| # Concat heads and project | |
| concatenated = torch.cat(head_outputs, dim=-1) # (B, D * num_heads) | |
| aggregated = self.projection(concatenated) # (B, D) | |
| weights = torch.stack(all_weights, dim=1) # (B, num_heads, S) | |
| return aggregated, weights | |