# Codeblock 10c def forward(self, x): print(f'original\t\t\t: {x.size()}') x = self.first_conv(x) print(f'after first_conv\t\t: {x.size()}') x = self.first_pool(x) #(1) print(f'after first_pool\t\t: {x.size()}\n') ##### Stage 0 part1, part2 = self.split_channels(x) #(2) print(f'part1\t\t\t\t: {part1.size()}') print(f'part2\t\t\t\t: {part2.size()}') part2 = self.dense_block_0(part2) #(3) print(f'part2 after dense block 0\t: {part2.size()}') part2 = self.first_transition_0(part2) #(4) print(f'part2 after first trans 0\t: {part2.size()}') x = torch.cat((part1, part2), dim=1) #(5) print(f'after concatenate\t\t: {x.size()}') x = self.second_transition_0(x) #(6) print(f'after second transition 0\t: {x.size()}\n') ##### Stage 1 part1, part2 = self.split_channels(x) print(f'part1\t\t\t\t: {part1.size()}') print(f'part2\t\t\t\t: {part2.size()}') part2 = self.dense_block_1(part2) print(f'part2 after dense block 1\t: {part2.size()}') part2 = self.first_transition_1(part2) print(f'part2 after first trans 1\t: {part2.size()}') x = torch.cat((part1, part2), dim=1) print(f'after concatenate\t\t: {x.size()}') x = self.second_transition_1(x) print(f'after second transition 1\t: {x.size()}\n') ##### Stage 2 part1, part2 = self.split_channels(x) print(f'part1\t\t\t\t: {part1.size()}') print(f'part2\t\t\t\t: {part2.size()}') part2 = self.dense_block_2(part2) print(f'part2 after dense block 2\t: {part2.size()}') part2 = self.first_transition_2(part2) print(f'part2 after first trans 2\t: {part2.size()}') x = torch.cat((part1, part2), dim=1) print(f'after concatenate\t\t: {x.size()}') x = self.second_transition_2(x) print(f'after second transition 2\t: {x.size()}\n') ##### Stage 3 part1, part2 = self.split_channels(x) print(f'part1\t\t\t\t: {part1.size()}') print(f'part2\t\t\t\t: {part2.size()}') part2 = self.dense_block_3(part2) print(f'part2 after dense block 2\t: {part2.size()}') part2 = self.first_transition_3(part2) print(f'part2 after first trans 2\t: {part2.size()}') x = torch.cat((part1, part2), dim=1) print(f'after concatenate\t\t: {x.size()}\n') x = self.avgpool(x) print(f'after avgpool\t\t\t: {x.size()}') x = torch.flatten(x, start_dim=1) print(f'after flatten\t\t\t: {x.size()}') x = self.fc(x) print(f'after fc\t\t\t: {x.size()}') return x