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