diff --git a/monai/networks/nets/daf3d.py b/monai/networks/nets/daf3d.py index 4e47d79a2d..8e7773f6fc 100644 --- a/monai/networks/nets/daf3d.py +++ b/monai/networks/nets/daf3d.py @@ -173,12 +173,20 @@ class Daf3dResNetBottleneck(ResNetBottleneck): stride: stride to use for second conv layer. downsample: which downsample layer to use. norm: which normalization layer to use. Defaults to group. + act: activation type and arguments. Defaults to relu. """ expansion = 2 def __init__( - self, in_planes, planes, spatial_dims=3, stride=1, downsample=None, norm=("group", {"num_groups": 32}) + self, + in_planes, + planes, + spatial_dims=3, + stride=1, + downsample=None, + norm=("group", {"num_groups": 32}), + act=("relu", {"inplace": True}), ): conv_type: Callable = Conv[Conv.CONV, spatial_dims] @@ -191,7 +199,7 @@ def __init__( norm_layer(channels=planes * self.expansion), ) - super().__init__(in_planes, planes, spatial_dims, stride, downsample) + super().__init__(in_planes, planes, spatial_dims, stride, downsample, act=act) # change norm from batch to group norm self.bn1 = norm_layer(channels=planes) @@ -216,12 +224,21 @@ class Daf3dResNetDilatedBottleneck(Daf3dResNetBottleneck): spatial_dims: number of spatial dimensions of the input image. stride: stride to use for second conv layer. downsample: which downsample layer to use. + norm: which normalization layer to use. Defaults to group. + act: activation type and arguments. Defaults to relu. """ def __init__( - self, in_planes, planes, spatial_dims=3, stride=1, downsample=None, norm=("group", {"num_groups": 32}) + self, + in_planes, + planes, + spatial_dims=3, + stride=1, + downsample=None, + norm=("group", {"num_groups": 32}), + act=("relu", {"inplace": True}), ): - super().__init__(in_planes, planes, spatial_dims, stride, downsample, norm) + super().__init__(in_planes, planes, spatial_dims, stride, downsample, norm, act) # add dilation in second convolution conv_type: Callable = Conv[Conv.CONV, spatial_dims] diff --git a/monai/networks/nets/resnet.py b/monai/networks/nets/resnet.py index 9142116ae4..3b30f57b8d 100644 --- a/monai/networks/nets/resnet.py +++ b/monai/networks/nets/resnet.py @@ -211,6 +211,9 @@ class ResNet(nn.Module): bias_downsample: whether to use bias term in the downsampling block when `shortcut_type` is 'B', default to `True`. act: activation type and arguments. Defaults to relu. norm: feature normalization type and arguments. Defaults to batch norm. + block_norm_act: whether the residual blocks and their downsampling layers also use `act` and `norm`. + Defaults to `False`, in which case only the first convolution uses them and the blocks keep + ReLU and batch norm, matching previous versions. """ @@ -231,6 +234,7 @@ def __init__( bias_downsample: bool = True, # for backwards compatibility (also see PR #5477) act: str | tuple = ("relu", {"inplace": True}), norm: str | tuple = "batch", + block_norm_act: bool = False, ) -> None: super().__init__() @@ -271,10 +275,17 @@ def __init__( self.bn1 = norm_layer self.act = get_act_layer(name=act) self.maxpool = pool_type(kernel_size=3, stride=2, padding=1) - self.layer1 = self._make_layer(block, block_inplanes[0], layers[0], spatial_dims, shortcut_type) - self.layer2 = self._make_layer(block, block_inplanes[1], layers[1], spatial_dims, shortcut_type, stride=2) - self.layer3 = self._make_layer(block, block_inplanes[2], layers[2], spatial_dims, shortcut_type, stride=2) - self.layer4 = self._make_layer(block, block_inplanes[3], layers[3], spatial_dims, shortcut_type, stride=2) + block_kwargs: dict = {"norm": norm, "act": act} if block_norm_act else {} + self.layer1 = self._make_layer(block, block_inplanes[0], layers[0], spatial_dims, shortcut_type, **block_kwargs) + self.layer2 = self._make_layer( + block, block_inplanes[1], layers[1], spatial_dims, shortcut_type, stride=2, **block_kwargs + ) + self.layer3 = self._make_layer( + block, block_inplanes[2], layers[2], spatial_dims, shortcut_type, stride=2, **block_kwargs + ) + self.layer4 = self._make_layer( + block, block_inplanes[3], layers[3], spatial_dims, shortcut_type, stride=2, **block_kwargs + ) self.avgpool = avgp_type(block_avgpool[spatial_dims]) self.fc = nn.Linear(block_inplanes[3] * block.expansion, num_classes) if feed_forward else None @@ -282,8 +293,11 @@ def __init__( if isinstance(m, conv_type): nn.init.kaiming_normal_(torch.as_tensor(m.weight), mode="fan_out", nonlinearity="relu") elif isinstance(m, type(norm_layer)): - nn.init.constant_(torch.as_tensor(m.weight), 1) - nn.init.constant_(torch.as_tensor(m.bias), 0) + # non-affine norm layers (e.g. instance/layer norm defaults) have no weight/bias + if m.weight is not None: + nn.init.constant_(torch.as_tensor(m.weight), 1) + if m.bias is not None: + nn.init.constant_(torch.as_tensor(m.bias), 0) elif isinstance(m, nn.Linear): nn.init.constant_(torch.as_tensor(m.bias), 0) @@ -302,6 +316,7 @@ def _make_layer( shortcut_type: str, stride: int = 1, norm: str | tuple = "batch", + act: str | tuple = ("relu", {"inplace": True}), ) -> nn.Sequential: conv_type: Callable = Conv[Conv.CONV, spatial_dims] @@ -333,13 +348,14 @@ def _make_layer( spatial_dims=spatial_dims, stride=stride, downsample=downsample, + act=act, norm=norm, ) ] self.in_planes = planes * block.expansion for _i in range(1, blocks): - layers.append(block(self.in_planes, planes, spatial_dims=spatial_dims, norm=norm)) + layers.append(block(self.in_planes, planes, spatial_dims=spatial_dims, act=act, norm=norm)) return nn.Sequential(*layers) diff --git a/tests/networks/nets/test_resnet.py b/tests/networks/nets/test_resnet.py index 241f57c78d..464a6d4257 100644 --- a/tests/networks/nets/test_resnet.py +++ b/tests/networks/nets/test_resnet.py @@ -232,6 +232,16 @@ [model, *TEST_CASE_1] for model in [resnet10, resnet18, resnet34, resnet50, resnet101, resnet152, resnet200] ] +# small 2D net used by the norm/act threading tests +TEST_CASE_NORM_ACT = { + "block": "basic", + "layers": [1, 1, 1, 1], + "block_inplanes": [8, 16, 32, 64], + "spatial_dims": 2, + "n_input_channels": 1, + "num_classes": 2, +} + CASE_EXTRACT_FEATURES = [ ( {"model_name": "resnet10", "pretrained": True, "spatial_dims": 3, "in_channels": 1}, @@ -316,6 +326,49 @@ def test_script(self, model, input_param, input_shape, expected_shape): test_data = torch.randn(input_shape) test_script_save(net, test_data) + def test_norm_act_reach_blocks(self): + """With `block_norm_act=True`, `norm`/`act` are used by the residual blocks and downsamples.""" + net = ResNet( + **TEST_CASE_NORM_ACT, norm=("instance", {"affine": True}), act=("leakyrelu", {}), block_norm_act=True + ) + for layer in (net.layer1, net.layer2, net.layer3, net.layer4): + for block in layer: + self.assertIsInstance(block.bn1, torch.nn.InstanceNorm2d) + self.assertIsInstance(block.bn2, torch.nn.InstanceNorm2d) + self.assertIsInstance(block.act, torch.nn.LeakyReLU) + if block.downsample is not None: + self.assertIsInstance(block.downsample[1], torch.nn.InstanceNorm2d) + + def test_non_affine_norm_init(self): + """Non-affine norms have no weight/bias, the init loop must not choke on them.""" + for block_norm_act in (False, True): + net = ResNet(**TEST_CASE_NORM_ACT, norm="instance", block_norm_act=block_norm_act) + self.assertIsInstance(net.bn1, torch.nn.InstanceNorm2d) + with eval_mode(net): + net.forward(torch.randn(1, 1, 32, 32)) + + def test_block_norm_act_off_by_default(self): + """Without `block_norm_act`, `norm`/`act` only reach the stem, as in previous versions.""" + net = ResNet(**TEST_CASE_NORM_ACT, norm=("instance", {"affine": True}), act=("leakyrelu", {})) + self.assertIsInstance(net.bn1, torch.nn.InstanceNorm2d) + self.assertIsInstance(net.layer1[0].bn1, torch.nn.BatchNorm2d) + self.assertIsInstance(net.layer1[0].act, torch.nn.ReLU) + with eval_mode(net): + net.forward(torch.randn(1, 1, 32, 32)) + + def test_default_norm_act_unchanged(self): + """Default construction must keep the same module tree as before the norm/act threading.""" + net = ResNet(**TEST_CASE_NORM_ACT) + for layer in (net.layer1, net.layer2, net.layer3, net.layer4): + for block in layer: + self.assertIsInstance(block.bn1, torch.nn.BatchNorm2d) + self.assertIsInstance(block.bn2, torch.nn.BatchNorm2d) + self.assertIsInstance(block.act, torch.nn.ReLU) + self.assertTrue(block.act.inplace) + # BatchNorm weights/biases are still initialised to 1/0 by the guarded init loop + self.assertTrue(torch.equal(net.layer1[0].bn1.weight, torch.ones_like(net.layer1[0].bn1.weight))) + self.assertTrue(torch.equal(net.layer1[0].bn1.bias, torch.zeros_like(net.layer1[0].bn1.bias))) + @SkipIfNoModule("hf_hub_download") class TestExtractFeatures(unittest.TestCase):