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
File size: 3,784 Bytes
00801a0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 | """
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
|