(self, kernel_size=7)
| 166 | |
| 167 | class SpatialAttention(nn.Module): |
| 168 | def __init__(self, kernel_size=7): |
| 169 | super().__init__() |
| 170 | self.conv = nn.Conv2d(2, 1, kernel_size, padding=kernel_size // 2) |
| 171 | |
| 172 | def forward(self, x): |
| 173 | avg_out = torch.mean(x, dim=1, keepdim=True) |