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