# Codeblock 5 class SAM(nn.Module): def __init__(self): super().__init__() self.conv = nn.Conv2d(in_channels=2, #(1) out_channels=1, kernel_size=7, padding=3, bias=False) self.sigmoid = nn.Sigmoid() #(2) def forward(self, x): original = x #(3) print(f'original\t\t: {x.size()}\n') x_max, _ = torch.max(x, dim=1, keepdim=True) #(4) print(f'x after maxpool (x_max)\t: {x_max.size()}') x_avg = torch.mean(x, dim=1, keepdim=True) #(5) print(f'x after avgpool (x_avg)\t: {x_avg.size()}\n') x = torch.cat([x_max, x_avg], dim=1) #(6) print(f'after concatenate\t: {x.size()}') x = self.conv(x) #(7) print(f'after conv\t\t: {x.size()}') x = self.sigmoid(x) #(8) print(f'after sigmoid\t\t: {x.size()}') x = x * original #(9) print(f'after multiply\t\t: {x.size()}') return x