Skip to content

Commit fb4ef6e

Browse files
committed
fix(networks): mirror subunit chain in ResidualUnit residual for explicit padding
With subunits>1 the single projection conv only reproduced one subunit's shrinkage, so cx + res mismatched. Build a Sequential mirroring the subunit geometry when padding is explicit. Signed-off-by: Soumya Snigdha Kundu <soumya_snigdha.kundu@kcl.ac.uk>
1 parent 59568cf commit fb4ef6e

2 files changed

Lines changed: 31 additions & 9 deletions

File tree

‎monai/networks/blocks/convolutions.py‎

Lines changed: 25 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -303,16 +303,32 @@ def __init__(
303303

304304
# apply convolution to input to change number of output channels and size to match that coming from self.conv
305305
if np.prod(strides) != 1 or in_channels != out_channels or not same_pad:
306-
rkernel_size = kernel_size
307-
rpadding = padding
308-
309-
# if only adapting number of channels a 1x1 kernel is used with no padding
310-
if np.prod(strides) == 1 and same_pad:
311-
rkernel_size = 1
312-
rpadding = 0
313-
314306
conv_type = Conv[Conv.CONV, self.spatial_dims]
315-
self.residual = conv_type(in_channels, out_channels, rkernel_size, strides, rpadding, bias=bias)
307+
if not same_pad and subunits > 1:
308+
# mirror the subunit chain so the residual shrinks exactly like the main path
309+
residual = nn.Sequential()
310+
schannels = in_channels
311+
sstrides = strides
312+
for su in range(subunits):
313+
residual.add_module(
314+
f"unit{su:d}",
315+
conv_type(
316+
schannels, out_channels, kernel_size, sstrides, padding, bias=bias, dilation=dilation
317+
),
318+
)
319+
schannels = out_channels
320+
sstrides = 1
321+
self.residual = residual
322+
else:
323+
rkernel_size = kernel_size
324+
rpadding = padding
325+
326+
# if only adapting number of channels a 1x1 kernel is used with no padding
327+
if np.prod(strides) == 1 and same_pad:
328+
rkernel_size = 1
329+
rpadding = 0
330+
331+
self.residual = conv_type(in_channels, out_channels, rkernel_size, strides, rpadding, bias=bias)
316332

317333
def forward(self, x: torch.Tensor) -> torch.Tensor:
318334
res: torch.Tensor = self.residual(x) # create the additive residual from x

‎tests/networks/blocks/test_convolutions.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -171,6 +171,12 @@ def test_padding0_identity_residual(self):
171171
expected_shape = (1, 2, 6, 6)
172172
self.assertEqual(out.shape, expected_shape)
173173

174+
def test_padding0_default_subunits(self):
175+
conv = ResidualUnit(2, 2, 2, kernel_size=3, padding=0)
176+
out = conv(torch.rand(1, 2, 8, 8))
177+
expected_shape = (1, 2, 4, 4)
178+
self.assertEqual(out.shape, expected_shape)
179+
174180

175181
if __name__ == "__main__":
176182
unittest.main()

0 commit comments

Comments
 (0)