""" Cross-View Attention Fusion for ACL-LKNet. Fuses plane-level representations from sagittal, coronal, and axial views using multi-head attention. This allows the model to learn cross-view dependencies rather than merely concatenating independent features. Supports two modes: - 'attention': Multi-head self-attention over 3 view embeddings - 'concat': Simple concatenation + FC (ablation baseline) """ import torch import torch.nn as nn class CrossViewAttentionFusion(nn.Module): """ Fuses 3 plane embeddings using multi-head self-attention. Input: 3 view embeddings, each of shape (B, D) Process: 1. Stack into sequence of length 3: (B, 3, D) 2. Add learnable view position embeddings 3. Multi-head self-attention 4. Mean-pool the 3 output embeddings → (B, D) """ def __init__( self, feature_dim: int, num_heads: int = 2, dropout: float = 0.1, ): super().__init__() self.feature_dim = feature_dim # Learnable position embeddings for each view # (sagittal=0, coronal=1, axial=2) self.view_pos_embed = nn.Parameter(torch.zeros(1, 3, feature_dim)) nn.init.trunc_normal_(self.view_pos_embed, std=0.02) # Pre-norm self.norm = nn.LayerNorm(feature_dim) # Multi-head self-attention self.mha = nn.MultiheadAttention( embed_dim=feature_dim, num_heads=num_heads, dropout=dropout, batch_first=True, ) # Post-attention FFN self.ffn = nn.Sequential( nn.LayerNorm(feature_dim), nn.Linear(feature_dim, feature_dim * 2), nn.GELU(), nn.Dropout(dropout), nn.Linear(feature_dim * 2, feature_dim), nn.Dropout(dropout), ) def forward( self, sagittal: torch.Tensor, coronal: torch.Tensor, axial: torch.Tensor, ) -> torch.Tensor: """ Args: sagittal: (B, D) sagittal plane embedding coronal: (B, D) coronal plane embedding axial: (B, D) axial plane embedding Returns: fused: (B, D) fused representation """ # Stack into sequence: (B, 3, D) x = torch.stack([sagittal, coronal, axial], dim=1) # Add view position embeddings x = x + self.view_pos_embed # Self-attention with residual x_norm = self.norm(x) attn_out, _ = self.mha(x_norm, x_norm, x_norm) x = x + attn_out # FFN with residual x = x + self.ffn(x) # Mean-pool over the 3 views fused = x.mean(dim=1) # (B, D) return fused class ConcatFusion(nn.Module): """ Simple concatenation + FC fusion (ablation baseline). Concatenates 3 view embeddings → FC → output embedding. """ def __init__(self, feature_dim: int, dropout: float = 0.1): super().__init__() self.fusion = nn.Sequential( nn.Linear(feature_dim * 3, feature_dim), nn.GELU(), nn.Dropout(dropout), ) def forward( self, sagittal: torch.Tensor, coronal: torch.Tensor, axial: torch.Tensor, ) -> torch.Tensor: x = torch.cat([sagittal, coronal, axial], dim=-1) # (B, 3D) return self.fusion(x) # (B, D) def create_fusion( fusion_type: str, feature_dim: int, num_heads: int = 2, dropout: float = 0.1, ) -> nn.Module: """Factory for fusion modules.""" if fusion_type == "attention": return CrossViewAttentionFusion(feature_dim, num_heads, dropout) elif fusion_type == "concat": return ConcatFusion(feature_dim, dropout) else: raise ValueError(f"Unknown fusion type: {fusion_type}")