""" 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