Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 21 additions & 4 deletions monai/networks/nets/daf3d.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]

Expand All @@ -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)
Expand All @@ -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]
Expand Down
30 changes: 23 additions & 7 deletions monai/networks/nets/resnet.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

"""

Expand All @@ -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__()

Expand Down Expand Up @@ -271,19 +275,29 @@ 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

for m in self.modules():
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)

Expand All @@ -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]

Expand Down Expand Up @@ -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)

Expand Down
53 changes: 53 additions & 0 deletions tests/networks/nets/test_resnet.py
Original file line number Diff line number Diff line change
Expand Up @@ -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},
Expand Down Expand Up @@ -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):
Expand Down
Loading