| 252 | |
| 253 | |
| 254 | class SPADE(nn.Module): |
| 255 | def __init__(self, norm_nc, label_nc): |
| 256 | super().__init__() |
| 257 | |
| 258 | self.param_free_norm = nn.InstanceNorm2d(norm_nc, affine=False) |
| 259 | nhidden = 128 |
| 260 | |
| 261 | self.mlp_shared = nn.Sequential( |
| 262 | nn.Conv2d(label_nc, nhidden, kernel_size=3, padding=1), |
| 263 | nn.ReLU()) |
| 264 | self.mlp_gamma = nn.Conv2d(nhidden, norm_nc, kernel_size=3, padding=1) |
| 265 | self.mlp_beta = nn.Conv2d(nhidden, norm_nc, kernel_size=3, padding=1) |
| 266 | |
| 267 | def forward(self, x, segmap): |
| 268 | normalized = self.param_free_norm(x) |
| 269 | segmap = F.interpolate(segmap, size=x.size()[2:], mode='nearest') |
| 270 | actv = self.mlp_shared(segmap) |
| 271 | gamma = self.mlp_gamma(actv) |
| 272 | beta = self.mlp_beta(actv) |
| 273 | out = normalized * (1 + gamma) + beta |
| 274 | return out |
| 275 | |
| 276 | |
| 277 | class SPADEResnetBlock(nn.Module): |