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)
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).
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.
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).
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.