File size: 9,336 Bytes
00801a0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
"""
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")