Image Classification
timm
English
medical-imaging
knee-mri
acl-tear-detection
deep-learning
convnext
self-attention
masked-slice-modeling
radiology
orthopedics
Eval Results (legacy)
Instructions to use shareefch1413/ACL-LKNet with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- timm
How to use shareefch1413/ACL-LKNet with timm:
import timm model = timm.create_model("hf-hub:shareefch1413/ACL-LKNet", pretrained=True) - Notebooks
- Google Colab
- Kaggle
Download src/models/backbone.py from shareefch1413/ACL-LKNet: direct link, hf CLI and curl.
- Browser
- Download file 9.34 kB
-
https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/src/models/backbone.py
- Command line
-
hf download hf://shareefch1413/ACL-LKNet/src/models/backbone.py
-
curl -L -o backbone.py https://huggingface.co/shareefch1413/ACL-LKNet/resolve/main/src/models/backbone.py
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") | |