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