MCPcopy Create free account
hub / github.com/OpenGVLab/DragGAN / forward

Method forward

draggan/stylegan2/training/networks.py:416–460  ·  view source on GitHub ↗
(self, x, img, ws, force_fp32=False, fused_modconv=None, update_emas=False, **layer_kwargs)

Source from the content-addressed store, hash-verified

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}'

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected