(self, x, img, ws, force_fp32=False, fused_modconv=None, update_emas=False, **layer_kwargs)
| 414 | resample_filter=resample_filter, channels_last=self.channels_last) |
| 415 | |
| 416 | def forward(self, x, img, ws, force_fp32=False, fused_modconv=None, update_emas=False, **layer_kwargs): |
| 417 | _ = update_emas # unused |
| 418 | misc.assert_shape(ws, [None, self.num_conv + self.num_torgb, self.w_dim]) |
| 419 | w_iter = iter(ws.unbind(dim=1)) |
| 420 | if ws.device.type != 'cuda': |
| 421 | force_fp32 = True |
| 422 | dtype = torch.float16 if self.use_fp16 and not force_fp32 else torch.float32 |
| 423 | memory_format = torch.channels_last if self.channels_last and not force_fp32 else torch.contiguous_format |
| 424 | if fused_modconv is None: |
| 425 | fused_modconv = self.fused_modconv_default |
| 426 | if fused_modconv == 'inference_only': |
| 427 | fused_modconv = (not self.training) |
| 428 | |
| 429 | # Input. |
| 430 | if self.in_channels == 0: |
| 431 | x = self.const.to(dtype=dtype, memory_format=memory_format) |
| 432 | x = x.unsqueeze(0).repeat([ws.shape[0], 1, 1, 1]) |
| 433 | else: |
| 434 | misc.assert_shape(x, [None, self.in_channels, self.resolution // 2, self.resolution // 2]) |
| 435 | x = x.to(dtype=dtype, memory_format=memory_format) |
| 436 | |
| 437 | # Main layers. |
| 438 | if self.in_channels == 0: |
| 439 | x = self.conv1(x, next(w_iter), fused_modconv=fused_modconv, **layer_kwargs) |
| 440 | elif self.architecture == 'resnet': |
| 441 | y = self.skip(x, gain=np.sqrt(0.5)) |
| 442 | x = self.conv0(x, next(w_iter), fused_modconv=fused_modconv, **layer_kwargs) |
| 443 | x = self.conv1(x, next(w_iter), fused_modconv=fused_modconv, gain=np.sqrt(0.5), **layer_kwargs) |
| 444 | x = y.add_(x) |
| 445 | else: |
| 446 | x = self.conv0(x, next(w_iter), fused_modconv=fused_modconv, **layer_kwargs) |
| 447 | x = self.conv1(x, next(w_iter), fused_modconv=fused_modconv, **layer_kwargs) |
| 448 | |
| 449 | # ToRGB. |
| 450 | if img is not None: |
| 451 | misc.assert_shape(img, [None, self.img_channels, self.resolution // 2, self.resolution // 2]) |
| 452 | img = upfirdn2d.upsample2d(img, self.resample_filter) |
| 453 | if self.is_last or self.architecture == 'skip': |
| 454 | y = self.torgb(x, next(w_iter), fused_modconv=fused_modconv) |
| 455 | y = y.to(dtype=torch.float32, memory_format=torch.contiguous_format) |
| 456 | img = img.add_(y) if img is not None else y |
| 457 | |
| 458 | assert x.dtype == dtype |
| 459 | assert img is None or img.dtype == torch.float32 |
| 460 | return x, img |
| 461 | |
| 462 | def extra_repr(self): |
| 463 | return f'resolution={self.resolution:d}, architecture={self.architecture:s}' |
nothing calls this directly
no outgoing calls
no test coverage detected