# Codeblock 7 class SEResNeXt(nn.Module): def __init__(self): super().__init__() # conv1 stage self.resnext_conv1 = nn.Conv2d(in_channels=NUM_CHANNELS[0], out_channels=NUM_CHANNELS[1], kernel_size=7, stride=2, padding=3, bias=False) nn.init.kaiming_normal_(self.resnext_conv1.weight, nonlinearity='relu') self.resnext_bn1 = nn.BatchNorm2d(num_features=NUM_CHANNELS[1]) self.relu = nn.ReLU() self.resnext_maxpool1 = nn.MaxPool2d(kernel_size=3, stride=2, padding=1) # conv2 stage self.resnext_conv2 = nn.ModuleList([ Block(in_channels=NUM_CHANNELS[1], add_channel=True, channel_multiplier=4, downsample=False) ]) for _ in range(NUM_BLOCKS[0]-1): self.resnext_conv2.append(Block(in_channels=NUM_CHANNELS[2])) # conv3 stage self.resnext_conv3 = nn.ModuleList([Block(in_channels=NUM_CHANNELS[2], add_channel=True, downsample=True)]) for _ in range(NUM_BLOCKS[1]-1): self.resnext_conv3.append(Block(in_channels=NUM_CHANNELS[3])) # conv4 stage self.resnext_conv4 = nn.ModuleList([Block(in_channels=NUM_CHANNELS[3], add_channel=True, downsample=True)]) for _ in range(NUM_BLOCKS[2]-1): self.resnext_conv4.append(Block(in_channels=NUM_CHANNELS[4])) # conv5 stage self.resnext_conv5 = nn.ModuleList([Block(in_channels=NUM_CHANNELS[4], add_channel=True, downsample=True)]) for _ in range(NUM_BLOCKS[3]-1): self.resnext_conv5.append(Block(in_channels=NUM_CHANNELS[5])) self.avgpool = nn.AdaptiveAvgPool2d(output_size=(1,1)) self.fc = nn.Linear(in_features=NUM_CHANNELS[5], out_features=NUM_CLASSES) def forward(self, x): print(f'original\t\t: {x.size()}') x = self.relu(self.resnext_bn1(self.resnext_conv1(x))) print(f'after resnext_conv1\t: {x.size()}') x = self.resnext_maxpool1(x) print(f'after resnext_maxpool1\t: {x.size()}') for i, block in enumerate(self.resnext_conv2): x = block(x) print(f'after resnext_conv2 #{i}\t: {x.size()}') for i, block in enumerate(self.resnext_conv3): x = block(x) print(f'after resnext_conv3 #{i}\t: {x.size()}') for i, block in enumerate(self.resnext_conv4): x = block(x) print(f'after resnext_conv4 #{i}\t: {x.size()}') for i, block in enumerate(self.resnext_conv5): x = block(x) print(f'after resnext_conv5 #{i}\t: {x.size()}') x = self.avgpool(x) print(f'after avgpool\t\t: {x.size()}') x = torch.flatten(x, start_dim=1) print(f'after flatten\t\t: {x.size()}') x = self.fc(x) print(f'after fc\t\t: {x.size()}') return x