torch_em.model.torchvision_unet

UNet variants with pretrained torchvision backbones as encoders.

Supported 2D backbones (pretrained on ImageNet): ResNet: resnet18, resnet34, resnet50, resnet101, resnet152 ResNeXt: resnext50_32x4d, resnext101_32x8d, resnext101_64x4d Wide ResNet: wide_resnet50_2, wide_resnet101_2 VGG: vgg11, vgg11_bn, vgg13, vgg13_bn, vgg16, vgg16_bn, vgg19, vgg19_bn DenseNet: densenet121, densenet161, densenet169, densenet201 MobileNet: mobilenet_v2, mobilenet_v3_small, mobilenet_v3_large EfficientNet: efficientnet_b0..b7, efficientnet_v2_s/m/l ConvNeXt: convnext_tiny, convnext_small, convnext_base, convnext_large RegNet: regnet_x_400mf/800mf/1_6gf/3_2gf/8gf/16gf/32gf, regnet_y_400mf/800mf/1_6gf/3_2gf/8gf/16gf/32gf MnasNet: mnasnet0_5, mnasnet0_75, mnasnet1_0, mnasnet1_3 GoogLeNet: googlenet ShuffleNet: shufflenet_v2_x0_5, shufflenet_v2_x1_0, shufflenet_v2_x1_5, shufflenet_v2_x2_0 Swin Transformer: swin_t, swin_s, swin_b, swin_v2_t, swin_v2_s, swin_v2_b

Supported 3D backbones (pretrained on Kinetics-400): r3d_18, r2plus1d_18, mc3_18

  1"""UNet variants with pretrained torchvision backbones as encoders.
  2
  3Supported 2D backbones (pretrained on ImageNet):
  4    ResNet: resnet18, resnet34, resnet50, resnet101, resnet152
  5    ResNeXt: resnext50_32x4d, resnext101_32x8d, resnext101_64x4d
  6    Wide ResNet: wide_resnet50_2, wide_resnet101_2
  7    VGG: vgg11, vgg11_bn, vgg13, vgg13_bn, vgg16, vgg16_bn, vgg19, vgg19_bn
  8    DenseNet: densenet121, densenet161, densenet169, densenet201
  9    MobileNet: mobilenet_v2, mobilenet_v3_small, mobilenet_v3_large
 10    EfficientNet: efficientnet_b0..b7, efficientnet_v2_s/m/l
 11    ConvNeXt: convnext_tiny, convnext_small, convnext_base, convnext_large
 12    RegNet: regnet_x_400mf/800mf/1_6gf/3_2gf/8gf/16gf/32gf, regnet_y_400mf/800mf/1_6gf/3_2gf/8gf/16gf/32gf
 13    MnasNet: mnasnet0_5, mnasnet0_75, mnasnet1_0, mnasnet1_3
 14    GoogLeNet: googlenet
 15    ShuffleNet: shufflenet_v2_x0_5, shufflenet_v2_x1_0, shufflenet_v2_x1_5, shufflenet_v2_x2_0
 16    Swin Transformer: swin_t, swin_s, swin_b, swin_v2_t, swin_v2_s, swin_v2_b
 17
 18Supported 3D backbones (pretrained on Kinetics-400):
 19    r3d_18, r2plus1d_18, mc3_18
 20"""
 21
 22import importlib
 23from typing import List, Optional, Union
 24
 25import torch
 26import torch.nn as nn
 27import torch.nn.functional as F
 28
 29from torchvision.models.feature_extraction import create_feature_extractor
 30
 31from .unet import UNetBase, Decoder, ConvBlock2d, ConvBlock3d, Upsampler2d, Upsampler3d
 32
 33
 34# Registry entry: (module_name, fn_name, node_names, channels, scale_factors, pre_skip_factor)
 35#
 36# node_names: all feature nodes [skip0, ..., skip_{N-1}, bottleneck], shallowest first
 37# channels: channel count at each node (same order as node_names)
 38# scale_factors: spatial downsampling between consecutive nodes (len = len(nodes) - 1)
 39# pre_skip_factor: spatial downsampling from the raw input to skip0 (int for 2D, (D,H,W) tuple for 3D)
 40#
 41# Backbones with pre_skip_factor=2 start at H/2; those with pre_skip_factor=4 start at H/4.
 42# The final_upsample in TorchvisionUNet2d restores output to input resolution.
 43
 44TV = "torchvision.models"
 45
 46# Node name lists shared across families
 47N_RESNET = ["relu", "layer1", "layer2", "layer3", "layer4"]
 48N_VGG11 = ["features.4", "features.9", "features.14", "features.19", "features.20"]
 49N_VGG11_BN = ["features.6", "features.13", "features.20", "features.27", "features.28"]
 50N_VGG13 = ["features.8", "features.13", "features.18", "features.23", "features.24"]
 51N_VGG13_BN = ["features.12", "features.19", "features.26", "features.33", "features.34"]
 52N_VGG16 = ["features.8", "features.15", "features.22", "features.29", "features.30"]
 53N_VGG16_BN = ["features.12", "features.22", "features.32", "features.42", "features.43"]
 54N_VGG19 = ["features.8", "features.17", "features.26", "features.35", "features.36"]
 55N_VGG19_BN = ["features.12", "features.25", "features.38", "features.51", "features.52"]
 56N_DENSENET = [
 57    "features.relu0", "features.transition1.conv",
 58    "features.transition2.conv", "features.transition3.conv", "features.norm5",
 59]
 60N_MOBILENET_V2 = ["features.1", "features.3", "features.6", "features.13", "features.18"]
 61N_MOBILENET_V3_S = ["features.1", "features.2", "features.4", "features.9"]
 62N_MOBILENET_V3_L = ["features.2", "features.4", "features.7", "features.16"]
 63N_EFF_B = ["features.1", "features.2", "features.3", "features.5", "features.7"]
 64N_EFF_V2 = ["features.1", "features.2", "features.3", "features.5", "features.6"]
 65N_CONVNEXT = ["features.1", "features.3", "features.5", "features.7"]
 66N_REGNET = ["stem", "trunk_output.block1", "trunk_output.block2", "trunk_output.block3", "trunk_output.block4"]
 67N_MNASNET = ["layers.8", "layers.9", "layers.10", "layers.12"]
 68N_GOOGLENET = ["conv3.relu", "inception3b.cat", "inception4e.cat", "inception5b.cat"]
 69N_SHUFFLENET = ["maxpool", "stage2.3.view_1", "stage3.7.view_1", "conv5.2"]
 70# Swin stage3 has 6 blocks for tiny and 18 for small/base, so two node lists are needed.
 71N_SWIN_T = [
 72    "features.1.1.add_1", "features.3.1.add_1", "features.5.5.add_1", "features.7.1.add_1",
 73]
 74N_SWIN_SB = [
 75    "features.1.1.add_1", "features.3.1.add_1", "features.5.17.add_1", "features.7.1.add_1",
 76]
 77
 78# Common channel and scale-factor lists
 79C_VGG = [128, 256, 512, 512, 512]
 80SF4 = [2, 2, 2, 2]  # depth=4 (pre_skip=2), scale factor 2 at each of 4 inter-node steps
 81SF3 = [2, 2, 2]  # depth=3 (pre_skip=2 or 4)
 82
 83# 3D backbone helpers
 84NODES_3D = ["layer1", "layer2", "layer3", "layer4"]
 85SF_3D_ISO = [(2, 2, 2), (2, 2, 2), (2, 2, 2)]
 86SF_3D_MC3 = [(1, 2, 2), (1, 2, 2), (1, 2, 2)]
 87
 88BACKBONE_REGISTRY_2D = {
 89    # ResNet (depth=4, pre_skip=2)
 90    "resnet18": (TV, "resnet18", N_RESNET, [64, 64, 128, 256, 512], SF4, 2),
 91    "resnet34": (TV, "resnet34", N_RESNET, [64, 64, 128, 256, 512], SF4, 2),
 92    "resnet50": (TV, "resnet50", N_RESNET, [64, 256, 512, 1024, 2048], SF4, 2),
 93    "resnet101": (TV, "resnet101", N_RESNET, [64, 256, 512, 1024, 2048], SF4, 2),
 94    "resnet152": (TV, "resnet152", N_RESNET, [64, 256, 512, 1024, 2048], SF4, 2),
 95    # ResNeXt (depth=4, pre_skip=2; same node structure as ResNet)
 96    "resnext50_32x4d": (TV, "resnext50_32x4d", N_RESNET, [64, 256, 512, 1024, 2048], SF4, 2),
 97    "resnext101_32x8d": (TV, "resnext101_32x8d", N_RESNET, [64, 256, 512, 1024, 2048], SF4, 2),
 98    "resnext101_64x4d": (TV, "resnext101_64x4d", N_RESNET, [64, 256, 512, 1024, 2048], SF4, 2),
 99    # Wide ResNet (depth=4, pre_skip=2; same node structure as ResNet)
100    "wide_resnet50_2": (TV, "wide_resnet50_2", N_RESNET, [64, 256, 512, 1024, 2048], SF4, 2),
101    "wide_resnet101_2": (TV, "wide_resnet101_2", N_RESNET, [64, 256, 512, 1024, 2048], SF4, 2),
102    # VGG (depth=4, pre_skip=2; nodes are last-ReLU before each pool, then the final MaxPool)
103    "vgg11": (TV, "vgg11", N_VGG11, C_VGG, SF4, 2),
104    "vgg11_bn": (TV, "vgg11_bn", N_VGG11_BN, C_VGG, SF4, 2),
105    "vgg13": (TV, "vgg13", N_VGG13, C_VGG, SF4, 2),
106    "vgg13_bn": (TV, "vgg13_bn", N_VGG13_BN, C_VGG, SF4, 2),
107    "vgg16": (TV, "vgg16", N_VGG16, C_VGG, SF4, 2),
108    "vgg16_bn": (TV, "vgg16_bn", N_VGG16_BN, C_VGG, SF4, 2),
109    "vgg19": (TV, "vgg19", N_VGG19, C_VGG, SF4, 2),
110    "vgg19_bn": (TV, "vgg19_bn", N_VGG19_BN, C_VGG, SF4, 2),
111    # DenseNet (depth=4, pre_skip=2; nodes are relu0, transition.conv layers, then norm5)
112    "densenet121": (TV, "densenet121", N_DENSENET, [64, 128, 256, 512, 1024], SF4, 2),
113    "densenet161": (TV, "densenet161", N_DENSENET, [96, 192, 384, 1056, 2208], SF4, 2),
114    "densenet169": (TV, "densenet169", N_DENSENET, [64, 128, 256, 640, 1664], SF4, 2),
115    "densenet201": (TV, "densenet201", N_DENSENET, [64, 128, 256, 896, 1920], SF4, 2),
116    # MobileNet V2 (depth=4, pre_skip=2)
117    "mobilenet_v2": (TV, "mobilenet_v2", N_MOBILENET_V2, [16, 24, 32, 96, 1280], SF4, 2),
118    # MobileNet V3 (depth=3, pre_skip=4; patchify-style stem goes H->H/4)
119    "mobilenet_v3_small": (TV, "mobilenet_v3_small", N_MOBILENET_V3_S, [16, 24, 40, 96], SF3, 4),
120    "mobilenet_v3_large": (TV, "mobilenet_v3_large", N_MOBILENET_V3_L, [24, 40, 80, 960], SF3, 4),
121    # EfficientNet B (depth=4, pre_skip=2)
122    "efficientnet_b0": (TV, "efficientnet_b0", N_EFF_B, [16, 24, 40, 112, 320], SF4, 2),
123    "efficientnet_b1": (TV, "efficientnet_b1", N_EFF_B, [16, 24, 40, 112, 320], SF4, 2),
124    "efficientnet_b2": (TV, "efficientnet_b2", N_EFF_B, [16, 24, 48, 120, 352], SF4, 2),
125    "efficientnet_b3": (TV, "efficientnet_b3", N_EFF_B, [24, 32, 48, 136, 384], SF4, 2),
126    "efficientnet_b4": (TV, "efficientnet_b4", N_EFF_B, [24, 32, 56, 160, 448], SF4, 2),
127    "efficientnet_b5": (TV, "efficientnet_b5", N_EFF_B, [24, 40, 64, 176, 512], SF4, 2),
128    "efficientnet_b6": (TV, "efficientnet_b6", N_EFF_B, [32, 40, 72, 200, 576], SF4, 2),
129    "efficientnet_b7": (TV, "efficientnet_b7", N_EFF_B, [32, 48, 80, 224, 640], SF4, 2),
130    # EfficientNet V2 (depth=4, pre_skip=2)
131    "efficientnet_v2_s": (TV, "efficientnet_v2_s", N_EFF_V2, [24, 48, 64, 160, 256], SF4, 2),
132    "efficientnet_v2_m": (TV, "efficientnet_v2_m", N_EFF_V2, [24, 48, 80, 176, 304], SF4, 2),
133    "efficientnet_v2_l": (TV, "efficientnet_v2_l", N_EFF_V2, [32, 64, 96, 224, 384], SF4, 2),
134    # ConvNeXt (depth=3, pre_skip=4; patchify stem goes H->H/4; stages are features.1/3/5/7)
135    "convnext_tiny": (TV, "convnext_tiny", N_CONVNEXT, [96, 192, 384, 768], SF3, 4),
136    "convnext_small": (TV, "convnext_small", N_CONVNEXT, [96, 192, 384, 768], SF3, 4),
137    "convnext_base": (TV, "convnext_base", N_CONVNEXT, [128, 256, 512, 1024], SF3, 4),
138    "convnext_large": (TV, "convnext_large", N_CONVNEXT, [192, 384, 768, 1536], SF3, 4),
139    # RegNet X (depth=4, pre_skip=2)
140    "regnet_x_400mf": (TV, "regnet_x_400mf", N_REGNET, [32, 32, 64, 160, 400], SF4, 2),
141    "regnet_x_800mf": (TV, "regnet_x_800mf", N_REGNET, [32, 64, 128, 288, 672], SF4, 2),
142    "regnet_x_1_6gf": (TV, "regnet_x_1_6gf", N_REGNET, [32, 72, 168, 408, 912], SF4, 2),
143    "regnet_x_3_2gf": (TV, "regnet_x_3_2gf", N_REGNET, [32, 96, 192, 432, 1008], SF4, 2),
144    "regnet_x_8gf": (TV, "regnet_x_8gf", N_REGNET, [32, 80, 240, 720, 1920], SF4, 2),
145    "regnet_x_16gf": (TV, "regnet_x_16gf", N_REGNET, [32, 256, 512, 896, 2048], SF4, 2),
146    "regnet_x_32gf": (TV, "regnet_x_32gf", N_REGNET, [32, 336, 672, 1344, 2520], SF4, 2),
147    # RegNet Y (depth=4, pre_skip=2)
148    "regnet_y_400mf": (TV, "regnet_y_400mf", N_REGNET, [32, 48, 104, 208, 440], SF4, 2),
149    "regnet_y_800mf": (TV, "regnet_y_800mf", N_REGNET, [32, 64, 144, 320, 784], SF4, 2),
150    "regnet_y_1_6gf": (TV, "regnet_y_1_6gf", N_REGNET, [32, 48, 120, 336, 888], SF4, 2),
151    "regnet_y_3_2gf": (TV, "regnet_y_3_2gf", N_REGNET, [32, 72, 216, 576, 1512], SF4, 2),
152    "regnet_y_8gf": (TV, "regnet_y_8gf", N_REGNET, [32, 224, 448, 896, 2016], SF4, 2),
153    "regnet_y_16gf": (TV, "regnet_y_16gf", N_REGNET, [32, 224, 448, 1232, 3024], SF4, 2),
154    "regnet_y_32gf": (TV, "regnet_y_32gf", N_REGNET, [32, 232, 696, 1392, 3712], SF4, 2),
155    # MnasNet (depth=3, pre_skip=4)
156    "mnasnet0_5": (TV, "mnasnet0_5", N_MNASNET, [16, 24, 40, 96], SF3, 4),
157    "mnasnet0_75": (TV, "mnasnet0_75", N_MNASNET, [24, 32, 64, 144], SF3, 4),
158    "mnasnet1_0": (TV, "mnasnet1_0", N_MNASNET, [24, 40, 80, 192], SF3, 4),
159    "mnasnet1_3": (TV, "mnasnet1_3", N_MNASNET, [32, 56, 104, 248], SF3, 4),
160    # GoogLeNet (depth=3, pre_skip=4; skips at last feature before each maxpool, bottleneck at inception5b)
161    "googlenet": (TV, "googlenet", N_GOOGLENET, [192, 480, 832, 1024], SF3, 4),
162    # ShuffleNet V2 (depth=3, pre_skip=4; skip0=maxpool, skips at end of stage2/3, bottleneck=conv5)
163    "shufflenet_v2_x0_5": (TV, "shufflenet_v2_x0_5", N_SHUFFLENET, [24, 48, 96, 1024], SF3, 4),
164    "shufflenet_v2_x1_0": (TV, "shufflenet_v2_x1_0", N_SHUFFLENET, [24, 116, 232, 1024], SF3, 4),
165    "shufflenet_v2_x1_5": (TV, "shufflenet_v2_x1_5", N_SHUFFLENET, [24, 176, 352, 1024], SF3, 4),
166    "shufflenet_v2_x2_0": (TV, "shufflenet_v2_x2_0", N_SHUFFLENET, [24, 244, 488, 2048], SF3, 4),
167    # Swin Transformer (depth=3, pre_skip=4; NHWC outputs require permute; t/s share channels, b is wider)
168    "swin_t": (TV, "swin_t", N_SWIN_T, [96, 192, 384, 768], SF3, 4, True),
169    "swin_s": (TV, "swin_s", N_SWIN_SB, [96, 192, 384, 768], SF3, 4, True),
170    "swin_b": (TV, "swin_b", N_SWIN_SB, [128, 256, 512, 1024], SF3, 4, True),
171    "swin_v2_t": (TV, "swin_v2_t", N_SWIN_T, [96, 192, 384, 768], SF3, 4, True),
172    "swin_v2_s": (TV, "swin_v2_s", N_SWIN_SB, [96, 192, 384, 768], SF3, 4, True),
173    "swin_v2_b": (TV, "swin_v2_b", N_SWIN_SB, [128, 256, 512, 1024], SF3, 4, True),
174}
175
176BACKBONE_REGISTRY_3D = {
177    "r3d_18": ("torchvision.models.video", "r3d_18", NODES_3D, [64, 128, 256, 512], SF_3D_ISO, (1, 2, 2)),
178    "r2plus1d_18": ("torchvision.models.video", "r2plus1d_18", NODES_3D, [64, 128, 256, 512], SF_3D_ISO, (1, 2, 2)),
179    "mc3_18": ("torchvision.models.video", "mc3_18", NODES_3D, [64, 128, 256, 512], SF_3D_MC3, (1, 2, 2)),
180}
181
182
183def _load_backbone(module_name, fn_name, pretrained):
184    module = importlib.import_module(module_name)
185    fn = getattr(module, fn_name)
186    return fn(weights="DEFAULT" if pretrained else None)
187
188
189class TorchvisionEncoder(nn.Module):
190    """@private"""
191
192    def __init__(
193        self,
194        backbone: nn.Module,
195        node_names: List[str],
196        skip_channels: List[int],
197        bottleneck_channels: int,
198        skip_targets: List[int],
199        bottleneck_target: int,
200        in_channels: int,
201        conv_cls,
202        nhwc: bool = False,
203    ):
204        super().__init__()
205
206        depth = len(skip_channels)
207        assert len(node_names) == depth + 1
208
209        return_nodes = {name: f"skip{i}" for i, name in enumerate(node_names[:-1])}
210        return_nodes[node_names[-1]] = "bottleneck"
211
212        self.input_proj = conv_cls(in_channels, 3, kernel_size=1) if in_channels != 3 else None
213        self.extractor = create_feature_extractor(backbone, return_nodes=return_nodes)
214        self._depth = depth
215
216        self.skip_projs = nn.ModuleList([
217            conv_cls(inc, outc, kernel_size=1) if inc != outc else nn.Identity()
218            for inc, outc in zip(skip_channels, skip_targets)
219        ])
220        self.bottleneck_proj = (
221            conv_cls(bottleneck_channels, bottleneck_target, kernel_size=1)
222            if bottleneck_channels != bottleneck_target else nn.Identity()
223        )
224
225        self.return_outputs = True
226        self._in_channels = in_channels
227        self._out_channels = bottleneck_target
228        self._nhwc = nhwc
229
230    @property
231    def in_channels(self):
232        return self._in_channels
233
234    @property
235    def out_channels(self):
236        return self._out_channels
237
238    def __len__(self):
239        return self._depth
240
241    def forward(self, x):
242        if self.input_proj is not None:
243            x = self.input_proj(x)
244
245        features = self.extractor(x)
246        if self._nhwc:
247            features = {k: v.permute(0, 3, 1, 2).contiguous() for k, v in features.items()}
248        skips = [proj(features[f"skip{i}"]) for i, proj in enumerate(self.skip_projs)]
249        bottleneck = self.bottleneck_proj(features["bottleneck"])
250
251        if self.return_outputs:
252            return bottleneck, skips
253        return bottleneck
254
255
256def _build_encoder_and_decoder(
257    backbone_name,
258    registry,
259    depth,
260    initial_features,
261    gain,
262    in_channels,
263    pretrained,
264    conv_block_impl,
265    sampler_impl,
266    conv_cls,
267    **conv_block_kwargs,
268):
269    entry = registry[backbone_name]
270    module_name, fn_name, all_nodes, all_channels, all_scale_factors, pre_skip_factor = entry[:6]
271    nhwc = entry[6] if len(entry) > 6 else False
272
273    max_depth = len(all_nodes) - 1
274    if depth > max_depth:
275        raise ValueError(f"Backbone '{backbone_name}' supports at most depth={max_depth}, got {depth}.")
276
277    used_nodes = all_nodes[:depth + 1]
278    used_channels = all_channels[:depth + 1]
279    skip_channels = used_channels[:depth]
280    bottleneck_channels = used_channels[depth]
281    scale_factors = all_scale_factors[:depth]
282
283    features_decoder = [initial_features * gain ** i for i in range(depth + 1)][::-1]
284
285    # Required projected skip channels (shallow to deep):
286    # decoder level i (deep to shallow) needs features_decoder[i] - features_decoder[i+1] channels
287    # from encoder_inputs[i] = skip at depth-1-i (after reversal in UNetBase)
288    skip_targets_deep_first = [features_decoder[i] - features_decoder[i + 1] for i in range(depth)]
289    skip_targets = list(reversed(skip_targets_deep_first))  # shallow to deep
290
291    # Base: initial_features * gain^(depth-1) -> initial_features * gain^depth
292    bottleneck_target = features_decoder[0] // gain
293
294    backbone = _load_backbone(module_name, fn_name, pretrained)
295
296    encoder = TorchvisionEncoder(
297        backbone=backbone,
298        node_names=used_nodes,
299        skip_channels=skip_channels,
300        bottleneck_channels=bottleneck_channels,
301        skip_targets=skip_targets,
302        bottleneck_target=bottleneck_target,
303        in_channels=in_channels,
304        conv_cls=conv_cls,
305        nhwc=nhwc,
306    )
307
308    base = conv_block_impl(bottleneck_target, features_decoder[0], **conv_block_kwargs)
309
310    decoder = Decoder(
311        features=features_decoder,
312        skip_channels=skip_targets_deep_first,
313        scale_factors=scale_factors[::-1],
314        conv_block_impl=conv_block_impl,
315        sampler_impl=sampler_impl,
316        **conv_block_kwargs,
317    )
318
319    return encoder, base, decoder, scale_factors, pre_skip_factor, features_decoder
320
321
322class TorchvisionUNetBase(UNetBase):
323    """@private"""
324
325    def __init__(
326        self,
327        encoder,
328        base,
329        decoder,
330        out_conv,
331        scale_factors,
332        pre_skip_factor,
333        final_activation,
334        postprocessing,
335        check_shape,
336        perform_range_checks,
337        norm_mean,
338        norm_std,
339    ):
340        super().__init__(
341            encoder=encoder,
342            base=base,
343            decoder=decoder,
344            out_conv=out_conv,
345            final_activation=final_activation,
346            postprocessing=postprocessing,
347            check_shape=check_shape,
348        )
349        self._scale_factors = scale_factors
350        self._pre_skip_factor = pre_skip_factor
351        self.perform_range_checks = perform_range_checks
352        self.register_buffer("norm_mean", norm_mean)
353        self.register_buffer("norm_std", norm_std)
354
355    def _check_shape(self, x):
356        spatial = x.shape[2:]
357        pre = self._pre_skip_factor
358        if not isinstance(pre, (list, tuple)):
359            pre = (pre,) * len(spatial)
360        for dim, (sh, pf) in enumerate(zip(spatial, pre)):
361            total = pf
362            for sf in self._scale_factors:
363                total = total * (sf[dim] if isinstance(sf, (list, tuple)) else sf)
364            if sh % total != 0:
365                raise ValueError(
366                    f"Input spatial shape {tuple(spatial)} is not compatible with this backbone and depth "
367                    f"(dim {dim} must be divisible by {total})."
368                )
369
370    def _apply_upsample(self, decoded: torch.Tensor) -> torch.Tensor:
371        raise NotImplementedError
372
373    def forward(self, x: torch.Tensor) -> torch.Tensor:
374        if self.norm_mean is not None and self.perform_range_checks:
375            actual_min, actual_max = torch.aminmax(x.detach())
376            if actual_min < 0.0 or actual_max > 1.0:
377                raise ValueError(
378                    f"Input is outside the expected [0, 1] range for pretrained normalization: "
379                    f"got [{actual_min.item():.4f}, {actual_max.item():.4f}]."
380                )
381
382        if self.norm_mean is not None:
383            x = (x - self.norm_mean) / self.norm_std
384
385        out = super().forward(x)
386        return [self._apply_upsample(t) for t in out] if isinstance(out, list) else self._apply_upsample(out)
387
388
389class TorchvisionUNet2d(TorchvisionUNetBase):
390    """A 2D U-Net that uses a pretrained torchvision backbone as the encoder.
391
392    Skip connections from the backbone are projected to match the standard
393    feature progression (initial_features * gain ** level). The decoder is
394    identical to the one used by UNet2d. A final bilinear upsample restores
395    the output to the original input resolution.
396
397    When pretrained=True and in_channels=3, inputs are automatically normalized
398    with ImageNet statistics (mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225))
399    and are expected in [0, 1].
400
401    Supported backbones (pretrained on ImageNet):
402        ResNet: resnet18, resnet34, resnet50, resnet101, resnet152
403        ResNeXt: resnext50_32x4d, resnext101_32x8d, resnext101_64x4d
404        Wide ResNet: wide_resnet50_2, wide_resnet101_2
405        VGG: vgg11, vgg11_bn, vgg13, vgg13_bn, vgg16, vgg16_bn, vgg19, vgg19_bn
406        DenseNet: densenet121, densenet161, densenet169, densenet201
407        MobileNet: mobilenet_v2, mobilenet_v3_small, mobilenet_v3_large
408        EfficientNet: efficientnet_b0..b7, efficientnet_v2_s/m/l
409        ConvNeXt: convnext_tiny, convnext_small, convnext_base, convnext_large
410        RegNet: regnet_x_400mf/800mf/1_6gf/3_2gf/8gf/16gf/32gf, regnet_y_400mf/800mf/1_6gf/3_2gf/8gf/16gf/32gf
411        MnasNet: mnasnet0_5, mnasnet0_75, mnasnet1_0, mnasnet1_3
412        GoogLeNet: googlenet
413        ShuffleNet: shufflenet_v2_x0_5, shufflenet_v2_x1_0, shufflenet_v2_x1_5, shufflenet_v2_x2_0
414        Swin Transformer: swin_t, swin_s, swin_b, swin_v2_t, swin_v2_s, swin_v2_b
415
416    Args:
417        backbone: Name of the torchvision backbone to use.
418        out_channels: Number of output channels.
419        in_channels: Number of input channels. Must be 3 when pretrained=True.
420            If != 3 and pretrained=False, a learned 1x1 projection maps the input
421            to 3 channels before the backbone.
422        depth: Number of encoder/decoder levels. Most backbones support depth=4;
423            ConvNeXt, MobileNetV3, and MnasNet support depth=3.
424        initial_features: Controls decoder channel widths: level i has
425            initial_features * gain ** i channels.
426        gain: Multiplier for decoder features per level.
427        pretrained: Whether to load ImageNet-pretrained backbone weights.
428        perform_range_checks: Whether to validate that inputs are in [0, 1] before
429            normalization. Disable to avoid GPU sync overhead during training.
430        final_activation: Activation applied after the output convolution.
431        postprocessing: Optional postprocessing module or name.
432        check_shape: Whether to validate the input shape.
433        conv_block_kwargs: Additional kwargs forwarded to ConvBlock2d (e.g. norm).
434    """
435
436    def __init__(
437        self,
438        backbone: str,
439        out_channels: int,
440        in_channels: int = 3,
441        depth: int = 4,
442        initial_features: int = 32,
443        gain: int = 2,
444        pretrained: bool = True,
445        perform_range_checks: bool = True,
446        final_activation: Optional[Union[str, nn.Module]] = None,
447        postprocessing: Optional[Union[str, nn.Module]] = None,
448        check_shape: bool = True,
449        **conv_block_kwargs,
450    ):
451        if backbone not in BACKBONE_REGISTRY_2D:
452            raise ValueError(f"Unknown 2D backbone '{backbone}'. Choose from: {list(BACKBONE_REGISTRY_2D)}")
453        if pretrained and in_channels != 3:
454            raise ValueError(
455                "pretrained=True requires in_channels=3. The backbone was pretrained on 3-channel inputs "
456                "and cannot be meaningfully initialized from pretrained weights with a different channel count."
457            )
458
459        encoder, base, decoder, scale_factors, pre_skip_factor, features_decoder = _build_encoder_and_decoder(
460            backbone_name=backbone, registry=BACKBONE_REGISTRY_2D, depth=depth,
461            initial_features=initial_features, gain=gain, in_channels=in_channels,
462            pretrained=pretrained, conv_block_impl=ConvBlock2d, sampler_impl=Upsampler2d,
463            conv_cls=nn.Conv2d, **conv_block_kwargs,
464        )
465        out_conv = None if out_channels is None else nn.Conv2d(features_decoder[-1], out_channels, kernel_size=1)
466
467        if pretrained:
468            norm_mean = torch.tensor((0.485, 0.456, 0.406)).view(1, 3, 1, 1)
469            norm_std = torch.tensor((0.229, 0.224, 0.225)).view(1, 3, 1, 1)
470        else:
471            norm_mean = norm_std = None
472
473        super().__init__(
474            encoder=encoder, base=base, decoder=decoder, out_conv=out_conv,
475            scale_factors=scale_factors, pre_skip_factor=pre_skip_factor,
476            final_activation=final_activation, postprocessing=postprocessing,
477            check_shape=check_shape, perform_range_checks=perform_range_checks,
478            norm_mean=norm_mean, norm_std=norm_std,
479        )
480
481        self.init_kwargs = dict(
482            backbone=backbone, out_channels=out_channels, in_channels=in_channels, depth=depth,
483            initial_features=initial_features, gain=gain, pretrained=pretrained,
484            perform_range_checks=perform_range_checks, final_activation=final_activation,
485            postprocessing=postprocessing, check_shape=check_shape, **conv_block_kwargs,
486        )
487
488    def _apply_upsample(self, decoded: torch.Tensor) -> torch.Tensor:
489        return F.interpolate(decoded, scale_factor=float(self._pre_skip_factor), mode="bilinear", align_corners=False)
490
491
492class TorchvisionUNet3d(TorchvisionUNetBase):
493    """A 3D U-Net that uses a pretrained torchvision video backbone as the encoder.
494
495    The video backbone (trained on Kinetics-400) uses true 3D convolutions and can
496    process volumetric inputs (B, C, D, H, W) directly. Skip connections are
497    projected to match the standard feature progression. A final trilinear upsample
498    restores the output to the input resolution (the backbone stem downsamples H and W
499    by 2 before the first skip level).
500
501    When pretrained=True and in_channels=3, inputs are automatically normalized with
502    Kinetics-400 statistics (mean=(0.43216, 0.394666, 0.37645), std=(0.22803, 0.22145, 0.21699))
503    and are expected in [0, 1].
504
505    Supported backbones (pretrained on Kinetics-400):
506        r3d_18, r2plus1d_18, mc3_18
507
508    Args:
509        backbone: Name of the torchvision video backbone to use.
510        out_channels: Number of output channels.
511        in_channels: Number of input channels. Must be 3 when pretrained=True.
512            If != 3 and pretrained=False, a learned 1x1 projection maps the input
513            to 3 channels before the backbone.
514        depth: Number of encoder/decoder levels (max 3 for supported video backbones).
515        initial_features: Controls decoder channel widths: level i has
516            initial_features * gain ** i channels.
517        gain: Multiplier for decoder features per level.
518        pretrained: Whether to load Kinetics-400-pretrained backbone weights.
519        perform_range_checks: Whether to validate that inputs are in [0, 1] before
520            normalization. Disable to avoid GPU sync overhead during training.
521        final_activation: Activation applied after the output convolution.
522        postprocessing: Optional postprocessing module or name.
523        check_shape: Whether to validate the input shape.
524        conv_block_kwargs: Additional kwargs forwarded to ConvBlock3d (e.g. norm).
525    """
526
527    def __init__(
528        self,
529        backbone: str,
530        out_channels: int,
531        in_channels: int = 3,
532        depth: int = 3,
533        initial_features: int = 32,
534        gain: int = 2,
535        pretrained: bool = True,
536        perform_range_checks: bool = True,
537        final_activation: Optional[Union[str, nn.Module]] = None,
538        postprocessing: Optional[Union[str, nn.Module]] = None,
539        check_shape: bool = True,
540        **conv_block_kwargs,
541    ):
542        if backbone not in BACKBONE_REGISTRY_3D:
543            raise ValueError(f"Unknown 3D backbone '{backbone}'. Choose from: {list(BACKBONE_REGISTRY_3D)}")
544        if pretrained and in_channels != 3:
545            raise ValueError(
546                "pretrained=True requires in_channels=3. The backbone was pretrained on 3-channel inputs "
547                "and cannot be meaningfully initialized from pretrained weights with a different channel count."
548            )
549
550        encoder, base, decoder, scale_factors, pre_skip_factor, features_decoder = _build_encoder_and_decoder(
551            backbone_name=backbone, registry=BACKBONE_REGISTRY_3D, depth=depth,
552            initial_features=initial_features, gain=gain, in_channels=in_channels,
553            pretrained=pretrained, conv_block_impl=ConvBlock3d, sampler_impl=Upsampler3d,
554            conv_cls=nn.Conv3d, **conv_block_kwargs,
555        )
556        out_conv = None if out_channels is None else nn.Conv3d(features_decoder[-1], out_channels, kernel_size=1)
557
558        if pretrained:
559            norm_mean = torch.tensor((0.43216, 0.394666, 0.37645)).view(1, 3, 1, 1, 1)
560            norm_std = torch.tensor((0.22803, 0.22145, 0.216989)).view(1, 3, 1, 1, 1)
561        else:
562            norm_mean = norm_std = None
563
564        super().__init__(
565            encoder=encoder, base=base, decoder=decoder, out_conv=out_conv,
566            scale_factors=scale_factors, pre_skip_factor=pre_skip_factor,
567            final_activation=final_activation, postprocessing=postprocessing,
568            check_shape=check_shape, perform_range_checks=perform_range_checks,
569            norm_mean=norm_mean, norm_std=norm_std,
570        )
571
572        self.init_kwargs = dict(
573            backbone=backbone, out_channels=out_channels, in_channels=in_channels, depth=depth,
574            initial_features=initial_features, gain=gain, pretrained=pretrained,
575            perform_range_checks=perform_range_checks, final_activation=final_activation,
576            postprocessing=postprocessing, check_shape=check_shape, **conv_block_kwargs,
577        )
578
579    def _apply_upsample(self, decoded: torch.Tensor) -> torch.Tensor:
580        scale = [float(f) for f in self._pre_skip_factor]
581        return F.interpolate(decoded, scale_factor=scale, mode="trilinear", align_corners=False)
TV = 'torchvision.models'
N_RESNET = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
N_VGG11 = ['features.4', 'features.9', 'features.14', 'features.19', 'features.20']
N_VGG11_BN = ['features.6', 'features.13', 'features.20', 'features.27', 'features.28']
N_VGG13 = ['features.8', 'features.13', 'features.18', 'features.23', 'features.24']
N_VGG13_BN = ['features.12', 'features.19', 'features.26', 'features.33', 'features.34']
N_VGG16 = ['features.8', 'features.15', 'features.22', 'features.29', 'features.30']
N_VGG16_BN = ['features.12', 'features.22', 'features.32', 'features.42', 'features.43']
N_VGG19 = ['features.8', 'features.17', 'features.26', 'features.35', 'features.36']
N_VGG19_BN = ['features.12', 'features.25', 'features.38', 'features.51', 'features.52']
N_DENSENET = ['features.relu0', 'features.transition1.conv', 'features.transition2.conv', 'features.transition3.conv', 'features.norm5']
N_MOBILENET_V2 = ['features.1', 'features.3', 'features.6', 'features.13', 'features.18']
N_MOBILENET_V3_S = ['features.1', 'features.2', 'features.4', 'features.9']
N_MOBILENET_V3_L = ['features.2', 'features.4', 'features.7', 'features.16']
N_EFF_B = ['features.1', 'features.2', 'features.3', 'features.5', 'features.7']
N_EFF_V2 = ['features.1', 'features.2', 'features.3', 'features.5', 'features.6']
N_CONVNEXT = ['features.1', 'features.3', 'features.5', 'features.7']
N_REGNET = ['stem', 'trunk_output.block1', 'trunk_output.block2', 'trunk_output.block3', 'trunk_output.block4']
N_MNASNET = ['layers.8', 'layers.9', 'layers.10', 'layers.12']
N_GOOGLENET = ['conv3.relu', 'inception3b.cat', 'inception4e.cat', 'inception5b.cat']
N_SHUFFLENET = ['maxpool', 'stage2.3.view_1', 'stage3.7.view_1', 'conv5.2']
N_SWIN_T = ['features.1.1.add_1', 'features.3.1.add_1', 'features.5.5.add_1', 'features.7.1.add_1']
N_SWIN_SB = ['features.1.1.add_1', 'features.3.1.add_1', 'features.5.17.add_1', 'features.7.1.add_1']
C_VGG = [128, 256, 512, 512, 512]
SF4 = [2, 2, 2, 2]
SF3 = [2, 2, 2]
NODES_3D = ['layer1', 'layer2', 'layer3', 'layer4']
SF_3D_ISO = [(2, 2, 2), (2, 2, 2), (2, 2, 2)]
SF_3D_MC3 = [(1, 2, 2), (1, 2, 2), (1, 2, 2)]
BACKBONE_REGISTRY_2D = {'resnet18': ('torchvision.models', 'resnet18', ['relu', 'layer1', 'layer2', 'layer3', 'layer4'], [64, 64, 128, 256, 512], [2, 2, 2, 2], 2), 'resnet34': ('torchvision.models', 'resnet34', ['relu', 'layer1', 'layer2', 'layer3', 'layer4'], [64, 64, 128, 256, 512], [2, 2, 2, 2], 2), 'resnet50': ('torchvision.models', 'resnet50', ['relu', 'layer1', 'layer2', 'layer3', 'layer4'], [64, 256, 512, 1024, 2048], [2, 2, 2, 2], 2), 'resnet101': ('torchvision.models', 'resnet101', ['relu', 'layer1', 'layer2', 'layer3', 'layer4'], [64, 256, 512, 1024, 2048], [2, 2, 2, 2], 2), 'resnet152': ('torchvision.models', 'resnet152', ['relu', 'layer1', 'layer2', 'layer3', 'layer4'], [64, 256, 512, 1024, 2048], [2, 2, 2, 2], 2), 'resnext50_32x4d': ('torchvision.models', 'resnext50_32x4d', ['relu', 'layer1', 'layer2', 'layer3', 'layer4'], [64, 256, 512, 1024, 2048], [2, 2, 2, 2], 2), 'resnext101_32x8d': ('torchvision.models', 'resnext101_32x8d', ['relu', 'layer1', 'layer2', 'layer3', 'layer4'], [64, 256, 512, 1024, 2048], [2, 2, 2, 2], 2), 'resnext101_64x4d': ('torchvision.models', 'resnext101_64x4d', ['relu', 'layer1', 'layer2', 'layer3', 'layer4'], [64, 256, 512, 1024, 2048], [2, 2, 2, 2], 2), 'wide_resnet50_2': ('torchvision.models', 'wide_resnet50_2', ['relu', 'layer1', 'layer2', 'layer3', 'layer4'], [64, 256, 512, 1024, 2048], [2, 2, 2, 2], 2), 'wide_resnet101_2': ('torchvision.models', 'wide_resnet101_2', ['relu', 'layer1', 'layer2', 'layer3', 'layer4'], [64, 256, 512, 1024, 2048], [2, 2, 2, 2], 2), 'vgg11': ('torchvision.models', 'vgg11', ['features.4', 'features.9', 'features.14', 'features.19', 'features.20'], [128, 256, 512, 512, 512], [2, 2, 2, 2], 2), 'vgg11_bn': ('torchvision.models', 'vgg11_bn', ['features.6', 'features.13', 'features.20', 'features.27', 'features.28'], [128, 256, 512, 512, 512], [2, 2, 2, 2], 2), 'vgg13': ('torchvision.models', 'vgg13', ['features.8', 'features.13', 'features.18', 'features.23', 'features.24'], [128, 256, 512, 512, 512], [2, 2, 2, 2], 2), 'vgg13_bn': ('torchvision.models', 'vgg13_bn', ['features.12', 'features.19', 'features.26', 'features.33', 'features.34'], [128, 256, 512, 512, 512], [2, 2, 2, 2], 2), 'vgg16': ('torchvision.models', 'vgg16', ['features.8', 'features.15', 'features.22', 'features.29', 'features.30'], [128, 256, 512, 512, 512], [2, 2, 2, 2], 2), 'vgg16_bn': ('torchvision.models', 'vgg16_bn', ['features.12', 'features.22', 'features.32', 'features.42', 'features.43'], [128, 256, 512, 512, 512], [2, 2, 2, 2], 2), 'vgg19': ('torchvision.models', 'vgg19', ['features.8', 'features.17', 'features.26', 'features.35', 'features.36'], [128, 256, 512, 512, 512], [2, 2, 2, 2], 2), 'vgg19_bn': ('torchvision.models', 'vgg19_bn', ['features.12', 'features.25', 'features.38', 'features.51', 'features.52'], [128, 256, 512, 512, 512], [2, 2, 2, 2], 2), 'densenet121': ('torchvision.models', 'densenet121', ['features.relu0', 'features.transition1.conv', 'features.transition2.conv', 'features.transition3.conv', 'features.norm5'], [64, 128, 256, 512, 1024], [2, 2, 2, 2], 2), 'densenet161': ('torchvision.models', 'densenet161', ['features.relu0', 'features.transition1.conv', 'features.transition2.conv', 'features.transition3.conv', 'features.norm5'], [96, 192, 384, 1056, 2208], [2, 2, 2, 2], 2), 'densenet169': ('torchvision.models', 'densenet169', ['features.relu0', 'features.transition1.conv', 'features.transition2.conv', 'features.transition3.conv', 'features.norm5'], [64, 128, 256, 640, 1664], [2, 2, 2, 2], 2), 'densenet201': ('torchvision.models', 'densenet201', ['features.relu0', 'features.transition1.conv', 'features.transition2.conv', 'features.transition3.conv', 'features.norm5'], [64, 128, 256, 896, 1920], [2, 2, 2, 2], 2), 'mobilenet_v2': ('torchvision.models', 'mobilenet_v2', ['features.1', 'features.3', 'features.6', 'features.13', 'features.18'], [16, 24, 32, 96, 1280], [2, 2, 2, 2], 2), 'mobilenet_v3_small': ('torchvision.models', 'mobilenet_v3_small', ['features.1', 'features.2', 'features.4', 'features.9'], [16, 24, 40, 96], [2, 2, 2], 4), 'mobilenet_v3_large': ('torchvision.models', 'mobilenet_v3_large', ['features.2', 'features.4', 'features.7', 'features.16'], [24, 40, 80, 960], [2, 2, 2], 4), 'efficientnet_b0': ('torchvision.models', 'efficientnet_b0', ['features.1', 'features.2', 'features.3', 'features.5', 'features.7'], [16, 24, 40, 112, 320], [2, 2, 2, 2], 2), 'efficientnet_b1': ('torchvision.models', 'efficientnet_b1', ['features.1', 'features.2', 'features.3', 'features.5', 'features.7'], [16, 24, 40, 112, 320], [2, 2, 2, 2], 2), 'efficientnet_b2': ('torchvision.models', 'efficientnet_b2', ['features.1', 'features.2', 'features.3', 'features.5', 'features.7'], [16, 24, 48, 120, 352], [2, 2, 2, 2], 2), 'efficientnet_b3': ('torchvision.models', 'efficientnet_b3', ['features.1', 'features.2', 'features.3', 'features.5', 'features.7'], [24, 32, 48, 136, 384], [2, 2, 2, 2], 2), 'efficientnet_b4': ('torchvision.models', 'efficientnet_b4', ['features.1', 'features.2', 'features.3', 'features.5', 'features.7'], [24, 32, 56, 160, 448], [2, 2, 2, 2], 2), 'efficientnet_b5': ('torchvision.models', 'efficientnet_b5', ['features.1', 'features.2', 'features.3', 'features.5', 'features.7'], [24, 40, 64, 176, 512], [2, 2, 2, 2], 2), 'efficientnet_b6': ('torchvision.models', 'efficientnet_b6', ['features.1', 'features.2', 'features.3', 'features.5', 'features.7'], [32, 40, 72, 200, 576], [2, 2, 2, 2], 2), 'efficientnet_b7': ('torchvision.models', 'efficientnet_b7', ['features.1', 'features.2', 'features.3', 'features.5', 'features.7'], [32, 48, 80, 224, 640], [2, 2, 2, 2], 2), 'efficientnet_v2_s': ('torchvision.models', 'efficientnet_v2_s', ['features.1', 'features.2', 'features.3', 'features.5', 'features.6'], [24, 48, 64, 160, 256], [2, 2, 2, 2], 2), 'efficientnet_v2_m': ('torchvision.models', 'efficientnet_v2_m', ['features.1', 'features.2', 'features.3', 'features.5', 'features.6'], [24, 48, 80, 176, 304], [2, 2, 2, 2], 2), 'efficientnet_v2_l': ('torchvision.models', 'efficientnet_v2_l', ['features.1', 'features.2', 'features.3', 'features.5', 'features.6'], [32, 64, 96, 224, 384], [2, 2, 2, 2], 2), 'convnext_tiny': ('torchvision.models', 'convnext_tiny', ['features.1', 'features.3', 'features.5', 'features.7'], [96, 192, 384, 768], [2, 2, 2], 4), 'convnext_small': ('torchvision.models', 'convnext_small', ['features.1', 'features.3', 'features.5', 'features.7'], [96, 192, 384, 768], [2, 2, 2], 4), 'convnext_base': ('torchvision.models', 'convnext_base', ['features.1', 'features.3', 'features.5', 'features.7'], [128, 256, 512, 1024], [2, 2, 2], 4), 'convnext_large': ('torchvision.models', 'convnext_large', ['features.1', 'features.3', 'features.5', 'features.7'], [192, 384, 768, 1536], [2, 2, 2], 4), 'regnet_x_400mf': ('torchvision.models', 'regnet_x_400mf', ['stem', 'trunk_output.block1', 'trunk_output.block2', 'trunk_output.block3', 'trunk_output.block4'], [32, 32, 64, 160, 400], [2, 2, 2, 2], 2), 'regnet_x_800mf': ('torchvision.models', 'regnet_x_800mf', ['stem', 'trunk_output.block1', 'trunk_output.block2', 'trunk_output.block3', 'trunk_output.block4'], [32, 64, 128, 288, 672], [2, 2, 2, 2], 2), 'regnet_x_1_6gf': ('torchvision.models', 'regnet_x_1_6gf', ['stem', 'trunk_output.block1', 'trunk_output.block2', 'trunk_output.block3', 'trunk_output.block4'], [32, 72, 168, 408, 912], [2, 2, 2, 2], 2), 'regnet_x_3_2gf': ('torchvision.models', 'regnet_x_3_2gf', ['stem', 'trunk_output.block1', 'trunk_output.block2', 'trunk_output.block3', 'trunk_output.block4'], [32, 96, 192, 432, 1008], [2, 2, 2, 2], 2), 'regnet_x_8gf': ('torchvision.models', 'regnet_x_8gf', ['stem', 'trunk_output.block1', 'trunk_output.block2', 'trunk_output.block3', 'trunk_output.block4'], [32, 80, 240, 720, 1920], [2, 2, 2, 2], 2), 'regnet_x_16gf': ('torchvision.models', 'regnet_x_16gf', ['stem', 'trunk_output.block1', 'trunk_output.block2', 'trunk_output.block3', 'trunk_output.block4'], [32, 256, 512, 896, 2048], [2, 2, 2, 2], 2), 'regnet_x_32gf': ('torchvision.models', 'regnet_x_32gf', ['stem', 'trunk_output.block1', 'trunk_output.block2', 'trunk_output.block3', 'trunk_output.block4'], [32, 336, 672, 1344, 2520], [2, 2, 2, 2], 2), 'regnet_y_400mf': ('torchvision.models', 'regnet_y_400mf', ['stem', 'trunk_output.block1', 'trunk_output.block2', 'trunk_output.block3', 'trunk_output.block4'], [32, 48, 104, 208, 440], [2, 2, 2, 2], 2), 'regnet_y_800mf': ('torchvision.models', 'regnet_y_800mf', ['stem', 'trunk_output.block1', 'trunk_output.block2', 'trunk_output.block3', 'trunk_output.block4'], [32, 64, 144, 320, 784], [2, 2, 2, 2], 2), 'regnet_y_1_6gf': ('torchvision.models', 'regnet_y_1_6gf', ['stem', 'trunk_output.block1', 'trunk_output.block2', 'trunk_output.block3', 'trunk_output.block4'], [32, 48, 120, 336, 888], [2, 2, 2, 2], 2), 'regnet_y_3_2gf': ('torchvision.models', 'regnet_y_3_2gf', ['stem', 'trunk_output.block1', 'trunk_output.block2', 'trunk_output.block3', 'trunk_output.block4'], [32, 72, 216, 576, 1512], [2, 2, 2, 2], 2), 'regnet_y_8gf': ('torchvision.models', 'regnet_y_8gf', ['stem', 'trunk_output.block1', 'trunk_output.block2', 'trunk_output.block3', 'trunk_output.block4'], [32, 224, 448, 896, 2016], [2, 2, 2, 2], 2), 'regnet_y_16gf': ('torchvision.models', 'regnet_y_16gf', ['stem', 'trunk_output.block1', 'trunk_output.block2', 'trunk_output.block3', 'trunk_output.block4'], [32, 224, 448, 1232, 3024], [2, 2, 2, 2], 2), 'regnet_y_32gf': ('torchvision.models', 'regnet_y_32gf', ['stem', 'trunk_output.block1', 'trunk_output.block2', 'trunk_output.block3', 'trunk_output.block4'], [32, 232, 696, 1392, 3712], [2, 2, 2, 2], 2), 'mnasnet0_5': ('torchvision.models', 'mnasnet0_5', ['layers.8', 'layers.9', 'layers.10', 'layers.12'], [16, 24, 40, 96], [2, 2, 2], 4), 'mnasnet0_75': ('torchvision.models', 'mnasnet0_75', ['layers.8', 'layers.9', 'layers.10', 'layers.12'], [24, 32, 64, 144], [2, 2, 2], 4), 'mnasnet1_0': ('torchvision.models', 'mnasnet1_0', ['layers.8', 'layers.9', 'layers.10', 'layers.12'], [24, 40, 80, 192], [2, 2, 2], 4), 'mnasnet1_3': ('torchvision.models', 'mnasnet1_3', ['layers.8', 'layers.9', 'layers.10', 'layers.12'], [32, 56, 104, 248], [2, 2, 2], 4), 'googlenet': ('torchvision.models', 'googlenet', ['conv3.relu', 'inception3b.cat', 'inception4e.cat', 'inception5b.cat'], [192, 480, 832, 1024], [2, 2, 2], 4), 'shufflenet_v2_x0_5': ('torchvision.models', 'shufflenet_v2_x0_5', ['maxpool', 'stage2.3.view_1', 'stage3.7.view_1', 'conv5.2'], [24, 48, 96, 1024], [2, 2, 2], 4), 'shufflenet_v2_x1_0': ('torchvision.models', 'shufflenet_v2_x1_0', ['maxpool', 'stage2.3.view_1', 'stage3.7.view_1', 'conv5.2'], [24, 116, 232, 1024], [2, 2, 2], 4), 'shufflenet_v2_x1_5': ('torchvision.models', 'shufflenet_v2_x1_5', ['maxpool', 'stage2.3.view_1', 'stage3.7.view_1', 'conv5.2'], [24, 176, 352, 1024], [2, 2, 2], 4), 'shufflenet_v2_x2_0': ('torchvision.models', 'shufflenet_v2_x2_0', ['maxpool', 'stage2.3.view_1', 'stage3.7.view_1', 'conv5.2'], [24, 244, 488, 2048], [2, 2, 2], 4), 'swin_t': ('torchvision.models', 'swin_t', ['features.1.1.add_1', 'features.3.1.add_1', 'features.5.5.add_1', 'features.7.1.add_1'], [96, 192, 384, 768], [2, 2, 2], 4, True), 'swin_s': ('torchvision.models', 'swin_s', ['features.1.1.add_1', 'features.3.1.add_1', 'features.5.17.add_1', 'features.7.1.add_1'], [96, 192, 384, 768], [2, 2, 2], 4, True), 'swin_b': ('torchvision.models', 'swin_b', ['features.1.1.add_1', 'features.3.1.add_1', 'features.5.17.add_1', 'features.7.1.add_1'], [128, 256, 512, 1024], [2, 2, 2], 4, True), 'swin_v2_t': ('torchvision.models', 'swin_v2_t', ['features.1.1.add_1', 'features.3.1.add_1', 'features.5.5.add_1', 'features.7.1.add_1'], [96, 192, 384, 768], [2, 2, 2], 4, True), 'swin_v2_s': ('torchvision.models', 'swin_v2_s', ['features.1.1.add_1', 'features.3.1.add_1', 'features.5.17.add_1', 'features.7.1.add_1'], [96, 192, 384, 768], [2, 2, 2], 4, True), 'swin_v2_b': ('torchvision.models', 'swin_v2_b', ['features.1.1.add_1', 'features.3.1.add_1', 'features.5.17.add_1', 'features.7.1.add_1'], [128, 256, 512, 1024], [2, 2, 2], 4, True)}
BACKBONE_REGISTRY_3D = {'r3d_18': ('torchvision.models.video', 'r3d_18', ['layer1', 'layer2', 'layer3', 'layer4'], [64, 128, 256, 512], [(2, 2, 2), (2, 2, 2), (2, 2, 2)], (1, 2, 2)), 'r2plus1d_18': ('torchvision.models.video', 'r2plus1d_18', ['layer1', 'layer2', 'layer3', 'layer4'], [64, 128, 256, 512], [(2, 2, 2), (2, 2, 2), (2, 2, 2)], (1, 2, 2)), 'mc3_18': ('torchvision.models.video', 'mc3_18', ['layer1', 'layer2', 'layer3', 'layer4'], [64, 128, 256, 512], [(1, 2, 2), (1, 2, 2), (1, 2, 2)], (1, 2, 2))}
class TorchvisionUNet2d(TorchvisionUNetBase):
390class TorchvisionUNet2d(TorchvisionUNetBase):
391    """A 2D U-Net that uses a pretrained torchvision backbone as the encoder.
392
393    Skip connections from the backbone are projected to match the standard
394    feature progression (initial_features * gain ** level). The decoder is
395    identical to the one used by UNet2d. A final bilinear upsample restores
396    the output to the original input resolution.
397
398    When pretrained=True and in_channels=3, inputs are automatically normalized
399    with ImageNet statistics (mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225))
400    and are expected in [0, 1].
401
402    Supported backbones (pretrained on ImageNet):
403        ResNet: resnet18, resnet34, resnet50, resnet101, resnet152
404        ResNeXt: resnext50_32x4d, resnext101_32x8d, resnext101_64x4d
405        Wide ResNet: wide_resnet50_2, wide_resnet101_2
406        VGG: vgg11, vgg11_bn, vgg13, vgg13_bn, vgg16, vgg16_bn, vgg19, vgg19_bn
407        DenseNet: densenet121, densenet161, densenet169, densenet201
408        MobileNet: mobilenet_v2, mobilenet_v3_small, mobilenet_v3_large
409        EfficientNet: efficientnet_b0..b7, efficientnet_v2_s/m/l
410        ConvNeXt: convnext_tiny, convnext_small, convnext_base, convnext_large
411        RegNet: regnet_x_400mf/800mf/1_6gf/3_2gf/8gf/16gf/32gf, regnet_y_400mf/800mf/1_6gf/3_2gf/8gf/16gf/32gf
412        MnasNet: mnasnet0_5, mnasnet0_75, mnasnet1_0, mnasnet1_3
413        GoogLeNet: googlenet
414        ShuffleNet: shufflenet_v2_x0_5, shufflenet_v2_x1_0, shufflenet_v2_x1_5, shufflenet_v2_x2_0
415        Swin Transformer: swin_t, swin_s, swin_b, swin_v2_t, swin_v2_s, swin_v2_b
416
417    Args:
418        backbone: Name of the torchvision backbone to use.
419        out_channels: Number of output channels.
420        in_channels: Number of input channels. Must be 3 when pretrained=True.
421            If != 3 and pretrained=False, a learned 1x1 projection maps the input
422            to 3 channels before the backbone.
423        depth: Number of encoder/decoder levels. Most backbones support depth=4;
424            ConvNeXt, MobileNetV3, and MnasNet support depth=3.
425        initial_features: Controls decoder channel widths: level i has
426            initial_features * gain ** i channels.
427        gain: Multiplier for decoder features per level.
428        pretrained: Whether to load ImageNet-pretrained backbone weights.
429        perform_range_checks: Whether to validate that inputs are in [0, 1] before
430            normalization. Disable to avoid GPU sync overhead during training.
431        final_activation: Activation applied after the output convolution.
432        postprocessing: Optional postprocessing module or name.
433        check_shape: Whether to validate the input shape.
434        conv_block_kwargs: Additional kwargs forwarded to ConvBlock2d (e.g. norm).
435    """
436
437    def __init__(
438        self,
439        backbone: str,
440        out_channels: int,
441        in_channels: int = 3,
442        depth: int = 4,
443        initial_features: int = 32,
444        gain: int = 2,
445        pretrained: bool = True,
446        perform_range_checks: bool = True,
447        final_activation: Optional[Union[str, nn.Module]] = None,
448        postprocessing: Optional[Union[str, nn.Module]] = None,
449        check_shape: bool = True,
450        **conv_block_kwargs,
451    ):
452        if backbone not in BACKBONE_REGISTRY_2D:
453            raise ValueError(f"Unknown 2D backbone '{backbone}'. Choose from: {list(BACKBONE_REGISTRY_2D)}")
454        if pretrained and in_channels != 3:
455            raise ValueError(
456                "pretrained=True requires in_channels=3. The backbone was pretrained on 3-channel inputs "
457                "and cannot be meaningfully initialized from pretrained weights with a different channel count."
458            )
459
460        encoder, base, decoder, scale_factors, pre_skip_factor, features_decoder = _build_encoder_and_decoder(
461            backbone_name=backbone, registry=BACKBONE_REGISTRY_2D, depth=depth,
462            initial_features=initial_features, gain=gain, in_channels=in_channels,
463            pretrained=pretrained, conv_block_impl=ConvBlock2d, sampler_impl=Upsampler2d,
464            conv_cls=nn.Conv2d, **conv_block_kwargs,
465        )
466        out_conv = None if out_channels is None else nn.Conv2d(features_decoder[-1], out_channels, kernel_size=1)
467
468        if pretrained:
469            norm_mean = torch.tensor((0.485, 0.456, 0.406)).view(1, 3, 1, 1)
470            norm_std = torch.tensor((0.229, 0.224, 0.225)).view(1, 3, 1, 1)
471        else:
472            norm_mean = norm_std = None
473
474        super().__init__(
475            encoder=encoder, base=base, decoder=decoder, out_conv=out_conv,
476            scale_factors=scale_factors, pre_skip_factor=pre_skip_factor,
477            final_activation=final_activation, postprocessing=postprocessing,
478            check_shape=check_shape, perform_range_checks=perform_range_checks,
479            norm_mean=norm_mean, norm_std=norm_std,
480        )
481
482        self.init_kwargs = dict(
483            backbone=backbone, out_channels=out_channels, in_channels=in_channels, depth=depth,
484            initial_features=initial_features, gain=gain, pretrained=pretrained,
485            perform_range_checks=perform_range_checks, final_activation=final_activation,
486            postprocessing=postprocessing, check_shape=check_shape, **conv_block_kwargs,
487        )
488
489    def _apply_upsample(self, decoded: torch.Tensor) -> torch.Tensor:
490        return F.interpolate(decoded, scale_factor=float(self._pre_skip_factor), mode="bilinear", align_corners=False)

A 2D U-Net that uses a pretrained torchvision backbone as the encoder.

Skip connections from the backbone are projected to match the standard feature progression (initial_features * gain ** level). The decoder is identical to the one used by UNet2d. A final bilinear upsample restores the output to the original input resolution.

When pretrained=True and in_channels=3, inputs are automatically normalized with ImageNet statistics (mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)) and are expected in [0, 1].

Supported backbones (pretrained on ImageNet): ResNet: resnet18, resnet34, resnet50, resnet101, resnet152 ResNeXt: resnext50_32x4d, resnext101_32x8d, resnext101_64x4d Wide ResNet: wide_resnet50_2, wide_resnet101_2 VGG: vgg11, vgg11_bn, vgg13, vgg13_bn, vgg16, vgg16_bn, vgg19, vgg19_bn DenseNet: densenet121, densenet161, densenet169, densenet201 MobileNet: mobilenet_v2, mobilenet_v3_small, mobilenet_v3_large EfficientNet: efficientnet_b0..b7, efficientnet_v2_s/m/l ConvNeXt: convnext_tiny, convnext_small, convnext_base, convnext_large RegNet: regnet_x_400mf/800mf/1_6gf/3_2gf/8gf/16gf/32gf, regnet_y_400mf/800mf/1_6gf/3_2gf/8gf/16gf/32gf MnasNet: mnasnet0_5, mnasnet0_75, mnasnet1_0, mnasnet1_3 GoogLeNet: googlenet ShuffleNet: shufflenet_v2_x0_5, shufflenet_v2_x1_0, shufflenet_v2_x1_5, shufflenet_v2_x2_0 Swin Transformer: swin_t, swin_s, swin_b, swin_v2_t, swin_v2_s, swin_v2_b

Arguments:
  • backbone: Name of the torchvision backbone to use.
  • out_channels: Number of output channels.
  • in_channels: Number of input channels. Must be 3 when pretrained=True. If != 3 and pretrained=False, a learned 1x1 projection maps the input to 3 channels before the backbone.
  • depth: Number of encoder/decoder levels. Most backbones support depth=4; ConvNeXt, MobileNetV3, and MnasNet support depth=3.
  • initial_features: Controls decoder channel widths: level i has initial_features * gain ** i channels.
  • gain: Multiplier for decoder features per level.
  • pretrained: Whether to load ImageNet-pretrained backbone weights.
  • perform_range_checks: Whether to validate that inputs are in [0, 1] before normalization. Disable to avoid GPU sync overhead during training.
  • final_activation: Activation applied after the output convolution.
  • postprocessing: Optional postprocessing module or name.
  • check_shape: Whether to validate the input shape.
  • conv_block_kwargs: Additional kwargs forwarded to ConvBlock2d (e.g. norm).
TorchvisionUNet2d( backbone: str, out_channels: int, in_channels: int = 3, depth: int = 4, initial_features: int = 32, gain: int = 2, pretrained: bool = True, perform_range_checks: bool = True, final_activation: Union[str, torch.nn.modules.module.Module, NoneType] = None, postprocessing: Union[str, torch.nn.modules.module.Module, NoneType] = None, check_shape: bool = True, **conv_block_kwargs)
437    def __init__(
438        self,
439        backbone: str,
440        out_channels: int,
441        in_channels: int = 3,
442        depth: int = 4,
443        initial_features: int = 32,
444        gain: int = 2,
445        pretrained: bool = True,
446        perform_range_checks: bool = True,
447        final_activation: Optional[Union[str, nn.Module]] = None,
448        postprocessing: Optional[Union[str, nn.Module]] = None,
449        check_shape: bool = True,
450        **conv_block_kwargs,
451    ):
452        if backbone not in BACKBONE_REGISTRY_2D:
453            raise ValueError(f"Unknown 2D backbone '{backbone}'. Choose from: {list(BACKBONE_REGISTRY_2D)}")
454        if pretrained and in_channels != 3:
455            raise ValueError(
456                "pretrained=True requires in_channels=3. The backbone was pretrained on 3-channel inputs "
457                "and cannot be meaningfully initialized from pretrained weights with a different channel count."
458            )
459
460        encoder, base, decoder, scale_factors, pre_skip_factor, features_decoder = _build_encoder_and_decoder(
461            backbone_name=backbone, registry=BACKBONE_REGISTRY_2D, depth=depth,
462            initial_features=initial_features, gain=gain, in_channels=in_channels,
463            pretrained=pretrained, conv_block_impl=ConvBlock2d, sampler_impl=Upsampler2d,
464            conv_cls=nn.Conv2d, **conv_block_kwargs,
465        )
466        out_conv = None if out_channels is None else nn.Conv2d(features_decoder[-1], out_channels, kernel_size=1)
467
468        if pretrained:
469            norm_mean = torch.tensor((0.485, 0.456, 0.406)).view(1, 3, 1, 1)
470            norm_std = torch.tensor((0.229, 0.224, 0.225)).view(1, 3, 1, 1)
471        else:
472            norm_mean = norm_std = None
473
474        super().__init__(
475            encoder=encoder, base=base, decoder=decoder, out_conv=out_conv,
476            scale_factors=scale_factors, pre_skip_factor=pre_skip_factor,
477            final_activation=final_activation, postprocessing=postprocessing,
478            check_shape=check_shape, perform_range_checks=perform_range_checks,
479            norm_mean=norm_mean, norm_std=norm_std,
480        )
481
482        self.init_kwargs = dict(
483            backbone=backbone, out_channels=out_channels, in_channels=in_channels, depth=depth,
484            initial_features=initial_features, gain=gain, pretrained=pretrained,
485            perform_range_checks=perform_range_checks, final_activation=final_activation,
486            postprocessing=postprocessing, check_shape=check_shape, **conv_block_kwargs,
487        )

Initialize internal Module state, shared by both nn.Module and ScriptModule.

init_kwargs
class TorchvisionUNet3d(TorchvisionUNetBase):
493class TorchvisionUNet3d(TorchvisionUNetBase):
494    """A 3D U-Net that uses a pretrained torchvision video backbone as the encoder.
495
496    The video backbone (trained on Kinetics-400) uses true 3D convolutions and can
497    process volumetric inputs (B, C, D, H, W) directly. Skip connections are
498    projected to match the standard feature progression. A final trilinear upsample
499    restores the output to the input resolution (the backbone stem downsamples H and W
500    by 2 before the first skip level).
501
502    When pretrained=True and in_channels=3, inputs are automatically normalized with
503    Kinetics-400 statistics (mean=(0.43216, 0.394666, 0.37645), std=(0.22803, 0.22145, 0.21699))
504    and are expected in [0, 1].
505
506    Supported backbones (pretrained on Kinetics-400):
507        r3d_18, r2plus1d_18, mc3_18
508
509    Args:
510        backbone: Name of the torchvision video backbone to use.
511        out_channels: Number of output channels.
512        in_channels: Number of input channels. Must be 3 when pretrained=True.
513            If != 3 and pretrained=False, a learned 1x1 projection maps the input
514            to 3 channels before the backbone.
515        depth: Number of encoder/decoder levels (max 3 for supported video backbones).
516        initial_features: Controls decoder channel widths: level i has
517            initial_features * gain ** i channels.
518        gain: Multiplier for decoder features per level.
519        pretrained: Whether to load Kinetics-400-pretrained backbone weights.
520        perform_range_checks: Whether to validate that inputs are in [0, 1] before
521            normalization. Disable to avoid GPU sync overhead during training.
522        final_activation: Activation applied after the output convolution.
523        postprocessing: Optional postprocessing module or name.
524        check_shape: Whether to validate the input shape.
525        conv_block_kwargs: Additional kwargs forwarded to ConvBlock3d (e.g. norm).
526    """
527
528    def __init__(
529        self,
530        backbone: str,
531        out_channels: int,
532        in_channels: int = 3,
533        depth: int = 3,
534        initial_features: int = 32,
535        gain: int = 2,
536        pretrained: bool = True,
537        perform_range_checks: bool = True,
538        final_activation: Optional[Union[str, nn.Module]] = None,
539        postprocessing: Optional[Union[str, nn.Module]] = None,
540        check_shape: bool = True,
541        **conv_block_kwargs,
542    ):
543        if backbone not in BACKBONE_REGISTRY_3D:
544            raise ValueError(f"Unknown 3D backbone '{backbone}'. Choose from: {list(BACKBONE_REGISTRY_3D)}")
545        if pretrained and in_channels != 3:
546            raise ValueError(
547                "pretrained=True requires in_channels=3. The backbone was pretrained on 3-channel inputs "
548                "and cannot be meaningfully initialized from pretrained weights with a different channel count."
549            )
550
551        encoder, base, decoder, scale_factors, pre_skip_factor, features_decoder = _build_encoder_and_decoder(
552            backbone_name=backbone, registry=BACKBONE_REGISTRY_3D, depth=depth,
553            initial_features=initial_features, gain=gain, in_channels=in_channels,
554            pretrained=pretrained, conv_block_impl=ConvBlock3d, sampler_impl=Upsampler3d,
555            conv_cls=nn.Conv3d, **conv_block_kwargs,
556        )
557        out_conv = None if out_channels is None else nn.Conv3d(features_decoder[-1], out_channels, kernel_size=1)
558
559        if pretrained:
560            norm_mean = torch.tensor((0.43216, 0.394666, 0.37645)).view(1, 3, 1, 1, 1)
561            norm_std = torch.tensor((0.22803, 0.22145, 0.216989)).view(1, 3, 1, 1, 1)
562        else:
563            norm_mean = norm_std = None
564
565        super().__init__(
566            encoder=encoder, base=base, decoder=decoder, out_conv=out_conv,
567            scale_factors=scale_factors, pre_skip_factor=pre_skip_factor,
568            final_activation=final_activation, postprocessing=postprocessing,
569            check_shape=check_shape, perform_range_checks=perform_range_checks,
570            norm_mean=norm_mean, norm_std=norm_std,
571        )
572
573        self.init_kwargs = dict(
574            backbone=backbone, out_channels=out_channels, in_channels=in_channels, depth=depth,
575            initial_features=initial_features, gain=gain, pretrained=pretrained,
576            perform_range_checks=perform_range_checks, final_activation=final_activation,
577            postprocessing=postprocessing, check_shape=check_shape, **conv_block_kwargs,
578        )
579
580    def _apply_upsample(self, decoded: torch.Tensor) -> torch.Tensor:
581        scale = [float(f) for f in self._pre_skip_factor]
582        return F.interpolate(decoded, scale_factor=scale, mode="trilinear", align_corners=False)

A 3D U-Net that uses a pretrained torchvision video backbone as the encoder.

The video backbone (trained on Kinetics-400) uses true 3D convolutions and can process volumetric inputs (B, C, D, H, W) directly. Skip connections are projected to match the standard feature progression. A final trilinear upsample restores the output to the input resolution (the backbone stem downsamples H and W by 2 before the first skip level).

When pretrained=True and in_channels=3, inputs are automatically normalized with Kinetics-400 statistics (mean=(0.43216, 0.394666, 0.37645), std=(0.22803, 0.22145, 0.21699)) and are expected in [0, 1].

Supported backbones (pretrained on Kinetics-400): r3d_18, r2plus1d_18, mc3_18

Arguments:
  • backbone: Name of the torchvision video backbone to use.
  • out_channels: Number of output channels.
  • in_channels: Number of input channels. Must be 3 when pretrained=True. If != 3 and pretrained=False, a learned 1x1 projection maps the input to 3 channels before the backbone.
  • depth: Number of encoder/decoder levels (max 3 for supported video backbones).
  • initial_features: Controls decoder channel widths: level i has initial_features * gain ** i channels.
  • gain: Multiplier for decoder features per level.
  • pretrained: Whether to load Kinetics-400-pretrained backbone weights.
  • perform_range_checks: Whether to validate that inputs are in [0, 1] before normalization. Disable to avoid GPU sync overhead during training.
  • final_activation: Activation applied after the output convolution.
  • postprocessing: Optional postprocessing module or name.
  • check_shape: Whether to validate the input shape.
  • conv_block_kwargs: Additional kwargs forwarded to ConvBlock3d (e.g. norm).
TorchvisionUNet3d( backbone: str, out_channels: int, in_channels: int = 3, depth: int = 3, initial_features: int = 32, gain: int = 2, pretrained: bool = True, perform_range_checks: bool = True, final_activation: Union[str, torch.nn.modules.module.Module, NoneType] = None, postprocessing: Union[str, torch.nn.modules.module.Module, NoneType] = None, check_shape: bool = True, **conv_block_kwargs)
528    def __init__(
529        self,
530        backbone: str,
531        out_channels: int,
532        in_channels: int = 3,
533        depth: int = 3,
534        initial_features: int = 32,
535        gain: int = 2,
536        pretrained: bool = True,
537        perform_range_checks: bool = True,
538        final_activation: Optional[Union[str, nn.Module]] = None,
539        postprocessing: Optional[Union[str, nn.Module]] = None,
540        check_shape: bool = True,
541        **conv_block_kwargs,
542    ):
543        if backbone not in BACKBONE_REGISTRY_3D:
544            raise ValueError(f"Unknown 3D backbone '{backbone}'. Choose from: {list(BACKBONE_REGISTRY_3D)}")
545        if pretrained and in_channels != 3:
546            raise ValueError(
547                "pretrained=True requires in_channels=3. The backbone was pretrained on 3-channel inputs "
548                "and cannot be meaningfully initialized from pretrained weights with a different channel count."
549            )
550
551        encoder, base, decoder, scale_factors, pre_skip_factor, features_decoder = _build_encoder_and_decoder(
552            backbone_name=backbone, registry=BACKBONE_REGISTRY_3D, depth=depth,
553            initial_features=initial_features, gain=gain, in_channels=in_channels,
554            pretrained=pretrained, conv_block_impl=ConvBlock3d, sampler_impl=Upsampler3d,
555            conv_cls=nn.Conv3d, **conv_block_kwargs,
556        )
557        out_conv = None if out_channels is None else nn.Conv3d(features_decoder[-1], out_channels, kernel_size=1)
558
559        if pretrained:
560            norm_mean = torch.tensor((0.43216, 0.394666, 0.37645)).view(1, 3, 1, 1, 1)
561            norm_std = torch.tensor((0.22803, 0.22145, 0.216989)).view(1, 3, 1, 1, 1)
562        else:
563            norm_mean = norm_std = None
564
565        super().__init__(
566            encoder=encoder, base=base, decoder=decoder, out_conv=out_conv,
567            scale_factors=scale_factors, pre_skip_factor=pre_skip_factor,
568            final_activation=final_activation, postprocessing=postprocessing,
569            check_shape=check_shape, perform_range_checks=perform_range_checks,
570            norm_mean=norm_mean, norm_std=norm_std,
571        )
572
573        self.init_kwargs = dict(
574            backbone=backbone, out_channels=out_channels, in_channels=in_channels, depth=depth,
575            initial_features=initial_features, gain=gain, pretrained=pretrained,
576            perform_range_checks=perform_range_checks, final_activation=final_activation,
577            postprocessing=postprocessing, check_shape=check_shape, **conv_block_kwargs,
578        )

Initialize internal Module state, shared by both nn.Module and ScriptModule.

init_kwargs