ACL-LKNet / src /models /slice_attention.py
shareefch1413's picture
Upload folder using huggingface_hub
00801a0 verified
Raw History Blame Contribute Delete
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