From 928f3af616a339ef58b531eafbe09c0c8d915c94 Mon Sep 17 00:00:00 2001 From: Soumya Snigdha Kundu Date: Tue, 28 Jul 2026 09:44:33 +0100 Subject: [PATCH 1/2] fix(nets): thread ResNet norm/act into residual blocks and guard non-affine init The documented `norm`/`act` constructor arguments only reached the stem: `_make_layer` was called without `norm` and had no `act` parameter, so every residual block and 'B'-shortcut downsample hard-defaulted to BatchNorm+ReLU. The weight-init loop also assumed affine norms, making `norm="instance"` (non-affine by default) raise during construction. Forward `norm`/`act` from `__init__` through `_make_layer` into the blocks and downsample, and skip `None` weight/bias in the init loop. The DAF3D blocks, which reuse the inherited `_make_layer`, gain a matching `act` parameter. Default construction (norm="batch", act="relu") is unchanged. Signed-off-by: Soumya Snigdha Kundu --- monai/networks/nets/daf3d.py | 25 ++++++++++++++--- monai/networks/nets/resnet.py | 27 +++++++++++++----- tests/networks/nets/test_resnet.py | 45 ++++++++++++++++++++++++++++-- 3 files changed, 84 insertions(+), 13 deletions(-) diff --git a/monai/networks/nets/daf3d.py b/monai/networks/nets/daf3d.py index 4e47d79a2dc..9d4397319b8 100644 --- a/monai/networks/nets/daf3d.py +++ b/monai/networks/nets/daf3d.py @@ -172,13 +172,21 @@ class Daf3dResNetBottleneck(ResNetBottleneck): spatial_dims: number of spatial dimensions of the input image. stride: stride to use for second conv layer. downsample: which downsample layer to use. + act: activation type and arguments. Defaults to relu. norm: which normalization layer to use. Defaults to group. """ 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, + act=("relu", {"inplace": True}), + norm=("group", {"num_groups": 32}), ): 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) # 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. + act: activation type and arguments. Defaults to relu. + norm: which normalization layer to use. Defaults to group. """ 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, + act=("relu", {"inplace": True}), + norm=("group", {"num_groups": 32}), ): - super().__init__(in_planes, planes, spatial_dims, stride, downsample, norm) + super().__init__(in_planes, planes, spatial_dims, stride, downsample, act, norm) # 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 9142116ae4d..b99d8908348 100644 --- a/monai/networks/nets/resnet.py +++ b/monai/networks/nets/resnet.py @@ -271,10 +271,18 @@ 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) + self.layer1 = self._make_layer( + block, block_inplanes[0], layers[0], spatial_dims, shortcut_type, act=act, norm=norm + ) + self.layer2 = self._make_layer( + block, block_inplanes[1], layers[1], spatial_dims, shortcut_type, stride=2, act=act, norm=norm + ) + self.layer3 = self._make_layer( + block, block_inplanes[2], layers[2], spatial_dims, shortcut_type, stride=2, act=act, norm=norm + ) + self.layer4 = self._make_layer( + block, block_inplanes[3], layers[3], spatial_dims, shortcut_type, stride=2, act=act, norm=norm + ) 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 +290,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) @@ -301,6 +312,7 @@ def _make_layer( spatial_dims: int, shortcut_type: str, stride: int = 1, + act: str | tuple = ("relu", {"inplace": True}), norm: str | tuple = "batch", ) -> nn.Sequential: conv_type: Callable = Conv[Conv.CONV, spatial_dims] @@ -333,13 +345,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 241f57c78dc..e792c9d81ee 100644 --- a/tests/networks/nets/test_resnet.py +++ b/tests/networks/nets/test_resnet.py @@ -202,7 +202,7 @@ (1, 3), ] -TEST_CASE_9 = [ # Layer norm +TEST_CASE_9 = [ # Group norm { "block": ResNetBlock, "layers": [3, 4, 6, 3], @@ -213,7 +213,7 @@ "conv1_t_size": [3], "conv1_t_stride": 1, "act": ("relu", {"inplace": False}), - "norm": ("layer", {"normalized_shape": (64, 32)}), + "norm": ("group", {"num_groups": 8}), }, (1, 2, 32), (1, 3), @@ -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,37 @@ 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): + """`norm`/`act` given to the constructor must be used by the residual blocks and downsamples.""" + net = ResNet(**{**TEST_CASE_NORM_ACT, "norm": ("instance", {"affine": True}), "act": ("leakyrelu", {})}) + 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.""" + net = ResNet(**{**TEST_CASE_NORM_ACT, "norm": "instance"}) + self.assertIsInstance(net.layer1[0].bn1, torch.nn.InstanceNorm2d) + 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): From 5f3e2b036c4b6ddf7ac8a7fbaed91f7e3ea49daf Mon Sep 17 00:00:00 2001 From: Soumya Snigdha Kundu Date: Wed, 7 Oct 2026 14:32:47 +0100 Subject: [PATCH 2/2] fix(nets): make ResNet block norm/act threading opt-in for backward compatibility Address review: passing ResNet's norm/act into the residual blocks changed the module tree of existing configurations (e.g. act='prelu' added parameters to every block). Add block_norm_act=False as the last ResNet argument so blocks keep ReLU and batch norm unless it is set, and move the new act argument after norm in _make_layer and the DAF3D bottlenecks so positional calls keep working. Assisted-by: Claude Opus 5.5 Signed-off-by: Soumya Snigdha Kundu --- monai/networks/nets/daf3d.py | 12 ++++++------ monai/networks/nets/resnet.py | 17 ++++++++++------- tests/networks/nets/test_resnet.py | 24 ++++++++++++++++++------ 3 files changed, 34 insertions(+), 19 deletions(-) diff --git a/monai/networks/nets/daf3d.py b/monai/networks/nets/daf3d.py index 9d4397319b8..8e7773f6fc8 100644 --- a/monai/networks/nets/daf3d.py +++ b/monai/networks/nets/daf3d.py @@ -172,8 +172,8 @@ class Daf3dResNetBottleneck(ResNetBottleneck): spatial_dims: number of spatial dimensions of the input image. stride: stride to use for second conv layer. downsample: which downsample layer to use. - act: activation type and arguments. Defaults to relu. norm: which normalization layer to use. Defaults to group. + act: activation type and arguments. Defaults to relu. """ expansion = 2 @@ -185,8 +185,8 @@ def __init__( spatial_dims=3, stride=1, downsample=None, - act=("relu", {"inplace": True}), norm=("group", {"num_groups": 32}), + act=("relu", {"inplace": True}), ): conv_type: Callable = Conv[Conv.CONV, spatial_dims] @@ -199,7 +199,7 @@ def __init__( norm_layer(channels=planes * self.expansion), ) - super().__init__(in_planes, planes, spatial_dims, stride, downsample, act) + super().__init__(in_planes, planes, spatial_dims, stride, downsample, act=act) # change norm from batch to group norm self.bn1 = norm_layer(channels=planes) @@ -224,8 +224,8 @@ 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. - act: activation type and arguments. Defaults to relu. norm: which normalization layer to use. Defaults to group. + act: activation type and arguments. Defaults to relu. """ def __init__( @@ -235,10 +235,10 @@ def __init__( spatial_dims=3, stride=1, downsample=None, - act=("relu", {"inplace": True}), norm=("group", {"num_groups": 32}), + act=("relu", {"inplace": True}), ): - super().__init__(in_planes, planes, spatial_dims, stride, downsample, act, 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 b99d8908348..3b30f57b8da 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,17 +275,16 @@ 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, act=act, norm=norm - ) + 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, act=act, norm=norm + 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, act=act, norm=norm + 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, act=act, norm=norm + 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 @@ -312,8 +315,8 @@ def _make_layer( spatial_dims: int, shortcut_type: str, stride: int = 1, - act: str | tuple = ("relu", {"inplace": True}), norm: str | tuple = "batch", + act: str | tuple = ("relu", {"inplace": True}), ) -> nn.Sequential: conv_type: Callable = Conv[Conv.CONV, spatial_dims] diff --git a/tests/networks/nets/test_resnet.py b/tests/networks/nets/test_resnet.py index e792c9d81ee..464a6d42578 100644 --- a/tests/networks/nets/test_resnet.py +++ b/tests/networks/nets/test_resnet.py @@ -202,7 +202,7 @@ (1, 3), ] -TEST_CASE_9 = [ # Group norm +TEST_CASE_9 = [ # Layer norm { "block": ResNetBlock, "layers": [3, 4, 6, 3], @@ -213,7 +213,7 @@ "conv1_t_size": [3], "conv1_t_stride": 1, "act": ("relu", {"inplace": False}), - "norm": ("group", {"num_groups": 8}), + "norm": ("layer", {"normalized_shape": (64, 32)}), }, (1, 2, 32), (1, 3), @@ -327,8 +327,10 @@ def test_script(self, model, input_param, input_shape, expected_shape): test_script_save(net, test_data) def test_norm_act_reach_blocks(self): - """`norm`/`act` given to the constructor must be used by the residual blocks and downsamples.""" - net = ResNet(**{**TEST_CASE_NORM_ACT, "norm": ("instance", {"affine": True}), "act": ("leakyrelu", {})}) + """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) @@ -339,8 +341,18 @@ def test_norm_act_reach_blocks(self): def test_non_affine_norm_init(self): """Non-affine norms have no weight/bias, the init loop must not choke on them.""" - net = ResNet(**{**TEST_CASE_NORM_ACT, "norm": "instance"}) - self.assertIsInstance(net.layer1[0].bn1, torch.nn.InstanceNorm2d) + 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))