# Codeblock 8 class InvResidualS1(nn.Module): def __init__(self, in_channels, out_channels, t): super().__init__() in_channels = int(in_channels*WIDTH_MULTIPLIER) #(1) out_channels = int(out_channels*WIDTH_MULTIPLIER) #(2) self.in_channels = in_channels self.out_channels = out_channels self.pwconv0 = nn.Conv2d(in_channels=in_channels, out_channels=in_channels*t, kernel_size=1, stride=1, bias=False) self.bn_pwconv0 = nn.BatchNorm2d(num_features=in_channels*t) self.dwconv = nn.Conv2d(in_channels=in_channels*t, out_channels=in_channels*t, kernel_size=3, stride=1, #(3) padding=1, groups=in_channels*t, bias=False) self.bn_dwconv = nn.BatchNorm2d(num_features=in_channels*t) self.pwconv1 = nn.Conv2d(in_channels=in_channels*t, out_channels=out_channels, kernel_size=1, stride=1, bias=False) self.bn_pwconv1 = nn.BatchNorm2d(num_features=out_channels) self.relu6 = nn.ReLU6() def forward(self, x): if self.in_channels == self.out_channels: #(4) residual = x #(5) print(f'residual\t\t: {residual.size()}') x = self.pwconv0(x) print('after pwconv0\t\t:', x.shape) x = self.bn_pwconv0(x) print('after bn_pwconv0\t:', x.shape) x = self.relu6(x) print('after relu\t\t:', x.shape) x = self.dwconv(x) print('after dwconv\t\t:', x.shape) x = self.bn_dwconv(x) print('after bn_dwconv\t\t:', x.shape) x = self.relu6(x) print('after relu\t\t:', x.shape) x = self.pwconv1(x) print('after pwconv1\t\t:', x.shape) x = self.bn_pwconv1(x) print('after bn_pwconv1\t:', x.shape) if self.in_channels == self.out_channels: x = x + residual #(6) print('after summation\t\t:', x.shape) return x