# Codeblock 7a class ConvNeXt(nn.Module): def __init__(self): super().__init__() self.stem = nn.Conv2d(in_channels=IN_CHANNELS, #(1) out_channels=OUT_CHANNELS[0], kernel_size=4, stride=4, ) self.normstem = nn.LayerNorm(normalized_shape=OUT_CHANNELS[0]) #(2) #(3) self.res2 = nn.ModuleList() for _ in range(NUM_BLOCKS[0]): self.res2.append(ConvNeXtBlock(num_channels=OUT_CHANNELS[0])) #(4) self.res3 = nn.ModuleList([ConvNeXtBlockTransition(in_channels=OUT_CHANNELS[0], out_channels=OUT_CHANNELS[1])]) for _ in range(NUM_BLOCKS[1]-1): self.res3.append(ConvNeXtBlock(num_channels=OUT_CHANNELS[1])) #(5) self.res4 = nn.ModuleList([ConvNeXtBlockTransition(in_channels=OUT_CHANNELS[1], out_channels=OUT_CHANNELS[2])]) for _ in range(NUM_BLOCKS[2]-1): self.res4.append(ConvNeXtBlock(num_channels=OUT_CHANNELS[2])) #(6) self.res5 = nn.ModuleList([ConvNeXtBlockTransition(in_channels=OUT_CHANNELS[2], out_channels=OUT_CHANNELS[3])]) for _ in range(NUM_BLOCKS[3]-1): self.res5.append(ConvNeXtBlock(num_channels=OUT_CHANNELS[3])) self.avgpool = nn.AdaptiveAvgPool2d(output_size=(1,1)) #(7) self.normpool = nn.LayerNorm(normalized_shape=OUT_CHANNELS[3]) #(8) self.fc = nn.Linear(in_features=OUT_CHANNELS[3], #(9) out_features=NUM_CLASSES) self.relu = nn.ReLU()