ACL-LKNet / src /models /acl_lknet.py
shareefch1413's picture
Upload folder using huggingface_hub
00801a0 verified
Raw History Blame Contribute Delete
8.6 kB
"""
ACL-LKNet: Full model assembly.
Combines all components into the complete pipeline:
Backbone β†’ Slice Attention β†’ Cross-View Fusion β†’ Classifier
Design:
- Shared backbone processes slices from all 3 views (parameter-efficient)
- Slices are processed in memory-efficient chunks (for T4 GPU)
- Supports switching between attention/concat fusion and pool/attention aggregation
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.checkpoint import checkpoint as grad_checkpoint
from .backbone import create_backbone
from .slice_attention import SliceAttention
from .cross_view_fusion import create_fusion
class ACLLKNet(nn.Module):
"""
ACL-LKNet: Hierarchical Self-Supervised Large-Kernel Network
for ACL Tear Detection in Knee MRI.
Architecture:
1. Shared LKNet backbone extracts per-slice features
2. Slice attention aggregates variable-length slice sequences per view
3. Cross-view attention fuses the 3 plane embeddings
4. Classification head outputs ACL tear probability
"""
def __init__(
self,
backbone_name: str = "convnext_tiny",
pretrained: bool = True,
feature_dim: int = 0,
attn_hidden_dim: int = 256,
fusion_type: str = "attention",
fusion_num_heads: int = 2,
fusion_dropout: float = 0.1,
classifier_hidden: int = 256,
classifier_dropout: float = 0.3,
aggregation: str = "attention", # 'attention', 'max', 'mean'
use_grad_checkpoint: bool = True,
slice_chunk_size: int = 8,
use_large_kernels: bool = False,
large_kernel_sizes: list = None,
):
super().__init__()
self.slice_chunk_size = slice_chunk_size
self.aggregation_type = aggregation
self.use_grad_checkpoint = use_grad_checkpoint
# ── Backbone ──
self.backbone = create_backbone(
name=backbone_name,
pretrained=pretrained,
use_grad_checkpoint=use_grad_checkpoint,
use_large_kernels=use_large_kernels,
large_kernel_sizes=large_kernel_sizes,
)
feat_dim = feature_dim if feature_dim > 0 else self.backbone.feature_dim
# ── Slice Attention (per view) ──
if aggregation == "attention":
self.slice_attn = SliceAttention(feat_dim, attn_hidden_dim)
else:
self.slice_attn = None
# ── Cross-View Fusion ──
self.fusion = create_fusion(
fusion_type=fusion_type,
feature_dim=feat_dim,
num_heads=fusion_num_heads,
dropout=fusion_dropout,
)
# ── Classification Head ──
self.classifier = nn.Sequential(
nn.LayerNorm(feat_dim),
nn.Linear(feat_dim, classifier_hidden),
nn.GELU(),
nn.Dropout(classifier_dropout),
nn.Linear(classifier_hidden, 1),
)
self._feat_dim = feat_dim
@property
def feature_dim(self):
return self._feat_dim
def _extract_slice_features(self, slices: torch.Tensor) -> torch.Tensor:
"""
Extract features from a batch of slices, processing in memory-efficient chunks.
Args:
slices: (B, S, H, W) β€” grayscale MRI slices for ONE view
Returns:
features: (B, S, D) β€” per-slice feature vectors
"""
B, S, H, W = slices.shape
all_features = []
for i in range(0, S, self.slice_chunk_size):
chunk = slices[:, i : i + self.slice_chunk_size] # (B, chunk, H, W)
chunk_size = chunk.shape[1]
# Reshape: (B, chunk, H, W) β†’ (B*chunk, 1, H, W)
chunk = chunk.reshape(B * chunk_size, 1, H, W)
# Forward through backbone
feat = self.backbone(chunk) # (B*chunk, D)
# Reshape back: (B, chunk, D)
feat = feat.reshape(B, chunk_size, -1)
all_features.append(feat)
features = torch.cat(all_features, dim=1) # (B, S, D)
return features
def _aggregate_slices(
self, features: torch.Tensor, mask: torch.Tensor = None
) -> tuple:
"""
Aggregate slice features into a single view embedding.
Args:
features: (B, S, D)
mask: (B, S) β€” True for valid slices
Returns:
embedding: (B, D)
attn_weights: (B, S) or None
"""
if self.aggregation_type == "attention" and self.slice_attn is not None:
return self.slice_attn(features, mask)
elif self.aggregation_type == "max":
if mask is not None:
features = features.masked_fill(~mask.unsqueeze(-1), float("-inf"))
return features.max(dim=1)[0], None
elif self.aggregation_type == "mean":
if mask is not None:
features = features * mask.unsqueeze(-1).float()
return features.sum(dim=1) / mask.sum(dim=1, keepdim=True).float(), None
return features.mean(dim=1), None
else:
raise ValueError(f"Unknown aggregation: {self.aggregation_type}")
def forward(
self,
sagittal: torch.Tensor,
coronal: torch.Tensor,
axial: torch.Tensor,
sag_mask: torch.Tensor = None,
cor_mask: torch.Tensor = None,
axi_mask: torch.Tensor = None,
) -> dict:
"""
Full forward pass.
Args:
sagittal: (B, S_sag, H, W) β€” sagittal MRI slices
coronal: (B, S_cor, H, W) β€” coronal MRI slices
axial: (B, S_axi, H, W) β€” axial MRI slices
sag_mask: (B, S_sag) β€” optional padding mask
cor_mask: (B, S_cor) β€” optional padding mask
axi_mask: (B, S_axi) β€” optional padding mask
Returns:
dict with:
'logits': (B, 1) β€” raw logits
'probs': (B, 1) β€” sigmoid probabilities
'sag_weights': (B, S_sag) β€” slice attention weights
'cor_weights': (B, S_cor) β€” slice attention weights
'axi_weights': (B, S_axi) β€” slice attention weights
"""
# Extract slice features (shared backbone)
sag_feats = self._extract_slice_features(sagittal)
cor_feats = self._extract_slice_features(coronal)
axi_feats = self._extract_slice_features(axial)
# Aggregate slices β†’ view embeddings
sag_emb, sag_w = self._aggregate_slices(sag_feats, sag_mask)
cor_emb, cor_w = self._aggregate_slices(cor_feats, cor_mask)
axi_emb, axi_w = self._aggregate_slices(axi_feats, axi_mask)
# Cross-view fusion
fused = self.fusion(sag_emb, cor_emb, axi_emb) # (B, D)
# Classification
logits = self.classifier(fused) # (B, 1)
probs = torch.sigmoid(logits)
return {
"logits": logits,
"probs": probs,
"sag_weights": sag_w,
"cor_weights": cor_w,
"axi_weights": axi_w,
}
def get_slice_features(
self,
sagittal: torch.Tensor,
coronal: torch.Tensor,
axial: torch.Tensor,
) -> dict:
"""
Extract slice features only (for MSM pretraining).
Does not run attention/fusion/classifier.
Returns:
dict with 'sagittal', 'coronal', 'axial' β€” each (B, S, D)
"""
return {
"sagittal": self._extract_slice_features(sagittal),
"coronal": self._extract_slice_features(coronal),
"axial": self._extract_slice_features(axial),
}
def create_model_from_config(config) -> ACLLKNet:
"""Create an ACLLKNet model from a Config object."""
return ACLLKNet(
backbone_name=config.backbone,
pretrained=config.pretrained,
feature_dim=config.feature_dim,
attn_hidden_dim=config.attn_hidden_dim,
fusion_type=config.fusion_type,
fusion_num_heads=config.fusion_num_heads,
fusion_dropout=config.fusion_dropout,
classifier_hidden=config.classifier_hidden,
classifier_dropout=config.classifier_dropout,
aggregation=getattr(config, "aggregation", "attention"),
use_grad_checkpoint=config.grad_checkpoint,
slice_chunk_size=config.slice_chunk_size,
use_large_kernels=config.use_large_kernels,
large_kernel_sizes=config.large_kernel_sizes,
)