""" Backbone feature extractor for ACL-LKNet. Uses timm for standard pretrained backbones (ConvNeXt, ResNet, EfficientNet). Includes optional large-kernel modification for kernel-size ablation. Design decisions: - ConvNeXt-Tiny is the default (28M params, 7×7 DW conv, T4-friendly) - ResNet-50 included for RadImageNet comparison (RadImageNet provides ResNet50 weights) - Large-kernel variant replaces ConvNeXt DW convolutions with larger kernels - All backbones output a single feature vector per input image via global avg pool """ import torch import torch.nn as nn import torch.nn.functional as F import timm from torch.utils.checkpoint import checkpoint as grad_checkpoint class BackboneWrapper(nn.Module): """ Wraps a timm backbone to: 1. Accept 1-channel grayscale input (converts to 3-channel) 2. Strip the classification head 3. Return a feature vector via global average pooling 4. Support gradient checkpointing for T4 memory """ def __init__( self, name: str = "convnext_tiny", pretrained: bool = True, use_grad_checkpoint: bool = True, ): super().__init__() self.name = name self.use_grad_checkpoint = use_grad_checkpoint # Create backbone from timm (no classification head) self.backbone = timm.create_model( name, pretrained=pretrained, num_classes=0, # Remove classifier → feature extractor global_pool="avg", # Global average pooling ) self.feature_dim = self.backbone.num_features # 1-channel → 3-channel adapter # We use a lightweight conv instead of simple replication so the model # can learn an optimal channel mapping for grayscale MRI self.channel_adapter = nn.Sequential( nn.Conv2d(1, 3, kernel_size=1, bias=False), nn.BatchNorm2d(3), ) # Initialize adapter to approximate channel replication nn.init.constant_(self.channel_adapter[0].weight, 1.0 / 3.0) if use_grad_checkpoint and hasattr(self.backbone, "set_grad_checkpointing"): self.backbone.set_grad_checkpointing(enable=True) def forward(self, x: torch.Tensor) -> torch.Tensor: """ Args: x: (B, 1, H, W) grayscale MRI slices Returns: features: (B, D) feature vectors """ # Grayscale → 3-channel x = self.channel_adapter(x) # Extract features features = self.backbone(x) return features class LargeKernelBlock(nn.Module): """ Large-kernel depth-wise convolution block inspired by RepLKNet. Uses depth-wise separable convolution with a large kernel, plus SE (Squeeze-and-Excitation) channel attention. For efficiency, large DW convolutions are decomposed into: depth-wise conv (large kernel) + point-wise conv (1×1) The large DW conv has very few parameters (kernel_size² × channels). """ def __init__(self, dim: int, kernel_size: int = 31, drop_path: float = 0.0): super().__init__() padding = kernel_size // 2 self.norm = nn.BatchNorm2d(dim) # Large-kernel depth-wise convolution self.dw_conv = nn.Conv2d( dim, dim, kernel_size=kernel_size, padding=padding, groups=dim, bias=False ) # SE attention self.se = nn.Sequential( nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Linear(dim, dim // 4), nn.GELU(), nn.Linear(dim // 4, dim), nn.Sigmoid(), ) # Point-wise (1×1) expansion self.pw_conv1 = nn.Conv2d(dim, dim * 4, kernel_size=1) self.act = nn.GELU() self.pw_conv2 = nn.Conv2d(dim * 4, dim, kernel_size=1) # Stochastic depth self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity() def forward(self, x: torch.Tensor) -> torch.Tensor: residual = x x = self.norm(x) x = self.dw_conv(x) # SE attention se_weight = self.se(x).unsqueeze(-1).unsqueeze(-1) x = x * se_weight # FFN x = self.pw_conv1(x) x = self.act(x) x = self.pw_conv2(x) x = self.drop_path(x) + residual return x class DropPath(nn.Module): """Stochastic depth — drops entire residual branches during training.""" def __init__(self, drop_prob: float = 0.0): super().__init__() self.drop_prob = drop_prob def forward(self, x: torch.Tensor) -> torch.Tensor: if not self.training or self.drop_prob == 0.0: return x keep_prob = 1 - self.drop_prob shape = (x.shape[0],) + (1,) * (x.ndim - 1) mask = torch.empty(shape, device=x.device).bernoulli_(keep_prob) return x * mask / keep_prob class LKNetBackbone(nn.Module): """ Custom Large-Kernel Network backbone for kernel-size ablation. 4-stage hierarchical design with configurable kernel sizes per stage. Used when we need to isolate the effect of kernel size independently of the backbone architecture (ConvNeXt vs ResNet, etc.). Stage dims: [64, 128, 256, 512] Stage depths: [2, 2, 6, 2] Default kernels: [7, 13, 21, 31] """ def __init__( self, in_channels: int = 1, dims: list = None, depths: list = None, kernel_sizes: list = None, drop_path_rate: float = 0.1, ): super().__init__() dims = dims or [64, 128, 256, 512] depths = depths or [2, 2, 6, 2] kernel_sizes = kernel_sizes or [7, 13, 21, 31] self.feature_dim = dims[-1] # Stem: 4× downsampling with small kernels (stable) self.stem = nn.Sequential( nn.Conv2d(in_channels, dims[0], kernel_size=4, stride=4), nn.BatchNorm2d(dims[0]), ) # Build stages self.stages = nn.ModuleList() dp_rates = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))] cur = 0 for i in range(4): # Downsampling between stages (except first) if i > 0: downsample = nn.Sequential( nn.BatchNorm2d(dims[i - 1]), nn.Conv2d(dims[i - 1], dims[i], kernel_size=2, stride=2), ) else: downsample = nn.Identity() # Stack of LK blocks blocks = nn.Sequential(*[ LargeKernelBlock(dims[i], kernel_sizes[i], dp_rates[cur + j]) for j in range(depths[i]) ]) cur += depths[i] self.stages.append(nn.Sequential(downsample, blocks)) self.norm = nn.LayerNorm(dims[-1]) self.pool = nn.AdaptiveAvgPool2d(1) def forward(self, x: torch.Tensor) -> torch.Tensor: """ Args: x: (B, C, H, W) Returns: features: (B, feature_dim) """ x = self.stem(x) for stage in self.stages: x = stage(x) x = self.pool(x).flatten(1) x = self.norm(x) return x def create_backbone( name: str = "convnext_tiny", pretrained: bool = True, use_grad_checkpoint: bool = True, use_large_kernels: bool = False, large_kernel_sizes: list = None, ) -> nn.Module: """ Factory function to create a backbone feature extractor. Args: name: Backbone name ('convnext_tiny', 'resnet18', 'resnet50', 'efficientnet_b0', 'lknet') pretrained: Load ImageNet pretrained weights (for timm models) use_grad_checkpoint: Enable gradient checkpointing use_large_kernels: Replace DW convolutions with larger kernels large_kernel_sizes: Kernel sizes per stage [s1, s2, s3, s4] Returns: backbone: nn.Module with .feature_dim attribute """ if name == "lknet": # Custom LK backbone (no pretrained weights — train from scratch or MSM) backbone = LKNetBackbone( in_channels=1, kernel_sizes=large_kernel_sizes or [7, 13, 21, 31], ) return backbone # Standard timm backbone backbone = BackboneWrapper( name=name, pretrained=pretrained, use_grad_checkpoint=use_grad_checkpoint, ) return backbone def load_radimagenet_weights(model: nn.Module, weights_path: str): """ Load RadImageNet pretrained weights into a ResNet/DenseNet backbone. RadImageNet provides weights for ResNet50, DenseNet121, InceptionV3. These are loaded into the backbone's internal model. Args: model: BackboneWrapper with a timm ResNet50 backbone weights_path: Path to RadImageNet .pt/.h5 weights file """ state_dict = torch.load(weights_path, map_location="cpu", weights_only=True) # RadImageNet weights may have different key names — attempt flexible loading model_dict = model.backbone.state_dict() filtered = {k: v for k, v in state_dict.items() if k in model_dict and v.shape == model_dict[k].shape} model_dict.update(filtered) model.backbone.load_state_dict(model_dict, strict=False) print(f"Loaded {len(filtered)}/{len(model_dict)} layers from RadImageNet weights")