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