diff --git a/model.py b/model.py index 2149fd9..25a4ee0 100644 --- a/model.py +++ b/model.py @@ -8,7 +8,7 @@ import clip # from pixel2style2pixel.models.psp import pSp from argparse import Namespace from utils import normalize -from stylegan.stylegan2_generator import StyleGAN2Generator +from stylegan.model import Generator import hydra from omegaconf import DictConfig, OmegaConf import sys @@ -27,11 +27,9 @@ def get_prompt(cfg): return prompt def get_stylegan_generator(cfg): - # model, preprocess = clip.load("RN50", device=device) - resolution=cfg.resolution - generator=StyleGAN2Generator(resolution=resolution) + generator=Generator(1024, 512, 8) checkpoint = torch.load(cfg.paths.stylegan, map_location=device) - generator.load_state_dict(checkpoint['generator']) + generator.load_state_dict(checkpoint['g_ema']) generator.to(device) generator.eval() return generator @@ -72,8 +70,8 @@ class GanAttack(nn.Module): nn.ReLU(inplace=True), nn.Linear(4096, 512) ) - basecode_layer = int(np.log2(cfg.basecode_spatial_size) - 2) * 2 - self.basecode_layer=basecode_layer = f'x{basecode_layer-1:02d}' + # basecode_layer = int(np.log2(cfg.basecode_spatial_size) - 2) * 2 + # self.basecode_layer=basecode_layer = f'x{basecode_layer-1:02d}' @@ -84,6 +82,7 @@ class GanAttack(nn.Module): x_prompt=torch.cat([detailcode,prompt],dim=2) x_prompt=self.mlp(x_prompt) x=x_prompt+x + noises=self.generator.make_noise() result_images=self.generator.synthesis(detailcode,randomize_noise=False, basecode_layer=self.basecode_layer,basecode=x)['image'] diff --git a/stylegan/__init__.py b/stylegan/__init__.py new file mode 100755 index 0000000..e69de29 diff --git a/stylegan/model.py b/stylegan/model.py new file mode 100755 index 0000000..13acc7f --- /dev/null +++ b/stylegan/model.py @@ -0,0 +1,711 @@ +import math +import random + +import torch +from torch import nn +from torch.nn import functional as F + +from models.stylegan2.op import FusedLeakyReLU, fused_leaky_relu, upfirdn2d +import numpy as np + +torch.manual_seed(0) +torch.backends.cudnn.deterministic = True +torch.backends.cudnn.benchmark = False +np.random.seed(0) + + +class PixelNorm(nn.Module): + def __init__(self): + super().__init__() + + def forward(self, input): + return input * torch.rsqrt(torch.mean(input ** 2, dim=1, keepdim=True) + 1e-8) + + +def make_kernel(k): + k = torch.tensor(k, dtype=torch.float32) + + if k.ndim == 1: + k = k[None, :] * k[:, None] + + k /= k.sum() + + return k + + +class Upsample(nn.Module): + def __init__(self, kernel, factor=2): + super().__init__() + + self.factor = factor + kernel = make_kernel(kernel) * (factor ** 2) + self.register_buffer('kernel', kernel) + + p = kernel.shape[0] - factor + + pad0 = (p + 1) // 2 + factor - 1 + pad1 = p // 2 + + self.pad = (pad0, pad1) + + def forward(self, input): + out = upfirdn2d(input, self.kernel, up=self.factor, down=1, pad=self.pad) + + return out + + +class Downsample(nn.Module): + def __init__(self, kernel, factor=2): + super().__init__() + + self.factor = factor + kernel = make_kernel(kernel) + self.register_buffer('kernel', kernel) + + p = kernel.shape[0] - factor + + pad0 = (p + 1) // 2 + pad1 = p // 2 + + self.pad = (pad0, pad1) + + def forward(self, input): + out = upfirdn2d(input, self.kernel, up=1, down=self.factor, pad=self.pad) + + return out + + +class Blur(nn.Module): + def __init__(self, kernel, pad, upsample_factor=1): + super().__init__() + + kernel = make_kernel(kernel) + + if upsample_factor > 1: + kernel = kernel * (upsample_factor ** 2) + + self.register_buffer('kernel', kernel) + + self.pad = pad + + def forward(self, input): + out = upfirdn2d(input, self.kernel, pad=self.pad) + + return out + + +class EqualConv2d(nn.Module): + def __init__( + self, in_channel, out_channel, kernel_size, stride=1, padding=0, bias=True + ): + super().__init__() + + self.weight = nn.Parameter( + torch.randn(out_channel, in_channel, kernel_size, kernel_size) + ) + self.scale = 1 / math.sqrt(in_channel * kernel_size ** 2) + + self.stride = stride + self.padding = padding + + if bias: + self.bias = nn.Parameter(torch.zeros(out_channel)) + + else: + self.bias = None + + def forward(self, input): + out = F.conv2d( + input, + self.weight * self.scale, + bias=self.bias, + stride=self.stride, + padding=self.padding, + ) + + return out + + def __repr__(self): + return ( + f'{self.__class__.__name__}({self.weight.shape[1]}, {self.weight.shape[0]},' + f' {self.weight.shape[2]}, stride={self.stride}, padding={self.padding})' + ) + + +class EqualLinear(nn.Module): + def __init__( + self, in_dim, out_dim, bias=True, bias_init=0, lr_mul=1, activation=None + ): + super().__init__() + + self.weight = nn.Parameter(torch.randn(out_dim, in_dim).div_(lr_mul)) + + if bias: + self.bias = nn.Parameter(torch.zeros(out_dim).fill_(bias_init)) + + else: + self.bias = None + + self.activation = activation + + self.scale = (1 / math.sqrt(in_dim)) * lr_mul + self.lr_mul = lr_mul + + def forward(self, input): + if self.activation: + out = F.linear(input, self.weight * self.scale) + out = fused_leaky_relu(out, self.bias * self.lr_mul) + + else: + out = F.linear( + input, self.weight * self.scale, bias=self.bias * self.lr_mul + ) + + return out + + def __repr__(self): + return ( + f'{self.__class__.__name__}({self.weight.shape[1]}, {self.weight.shape[0]})' + ) + + +class ScaledLeakyReLU(nn.Module): + def __init__(self, negative_slope=0.2): + super().__init__() + + self.negative_slope = negative_slope + + def forward(self, input): + out = F.leaky_relu(input, negative_slope=self.negative_slope) + + return out * math.sqrt(2) + + +class ModulatedConv2d(nn.Module): + def __init__( + self, + in_channel, + out_channel, + kernel_size, + style_dim, + demodulate=True, + upsample=False, + downsample=False, + blur_kernel=[1, 3, 3, 1], + ): + super().__init__() + + self.eps = 1e-8 + self.kernel_size = kernel_size + self.in_channel = in_channel + self.out_channel = out_channel + self.upsample = upsample + self.downsample = downsample + + if upsample: + factor = 2 + p = (len(blur_kernel) - factor) - (kernel_size - 1) + pad0 = (p + 1) // 2 + factor - 1 + pad1 = p // 2 + 1 + + self.blur = Blur(blur_kernel, pad=(pad0, pad1), upsample_factor=factor) + + if downsample: + factor = 2 + p = (len(blur_kernel) - factor) + (kernel_size - 1) + pad0 = (p + 1) // 2 + pad1 = p // 2 + + self.blur = Blur(blur_kernel, pad=(pad0, pad1)) + + fan_in = in_channel * kernel_size ** 2 + self.scale = 1 / math.sqrt(fan_in) + self.padding = kernel_size // 2 + + self.weight = nn.Parameter( + torch.randn(1, out_channel, in_channel, kernel_size, kernel_size) + ) + + self.modulation = EqualLinear(style_dim, in_channel, bias_init=1) + + self.demodulate = demodulate + + def __repr__(self): + return ( + f'{self.__class__.__name__}({self.in_channel}, {self.out_channel}, {self.kernel_size}, ' + f'upsample={self.upsample}, downsample={self.downsample})' + ) + + def forward(self, input, style, input_is_stylespace=False): + batch, in_channel, height, width = input.shape + + if not input_is_stylespace: + style = self.modulation(style).view(batch, 1, in_channel, 1, 1) + weight = self.scale * self.weight * style + + if self.demodulate: + demod = torch.rsqrt(weight.pow(2).sum([2, 3, 4]) + 1e-8) + weight = weight * demod.view(batch, self.out_channel, 1, 1, 1) + + weight = weight.view( + batch * self.out_channel, in_channel, self.kernel_size, self.kernel_size + ) + + if self.upsample: + input = input.view(1, batch * in_channel, height, width) + weight = weight.view( + batch, self.out_channel, in_channel, self.kernel_size, self.kernel_size + ) + weight = weight.transpose(1, 2).reshape( + batch * in_channel, self.out_channel, self.kernel_size, self.kernel_size + ) + out = F.conv_transpose2d(input, weight, padding=0, stride=2, groups=batch) + _, _, height, width = out.shape + out = out.view(batch, self.out_channel, height, width) + out = self.blur(out) + + elif self.downsample: + input = self.blur(input) + _, _, height, width = input.shape + input = input.view(1, batch * in_channel, height, width) + out = F.conv2d(input, weight, padding=0, stride=2, groups=batch) + _, _, height, width = out.shape + out = out.view(batch, self.out_channel, height, width) + + else: + input = input.view(1, batch * in_channel, height, width) + out = F.conv2d(input, weight, padding=self.padding, groups=batch) + _, _, height, width = out.shape + out = out.view(batch, self.out_channel, height, width) + + return out, style + + +class NoiseInjection(nn.Module): + def __init__(self): + super().__init__() + + self.weight = nn.Parameter(torch.zeros(1)) + + def forward(self, image, noise=None): + if noise is None: + batch, _, height, width = image.shape + noise = image.new_empty(batch, 1, height, width).normal_() + + return image + self.weight * noise + + +class ConstantInput(nn.Module): + def __init__(self, channel, size=4): + super().__init__() + + self.input = nn.Parameter(torch.randn(1, channel, size, size)) + + def forward(self, input): + batch = input.shape[0] + out = self.input.repeat(batch, 1, 1, 1) + + return out + + +class StyledConv(nn.Module): + def __init__( + self, + in_channel, + out_channel, + kernel_size, + style_dim, + upsample=False, + blur_kernel=[1, 3, 3, 1], + demodulate=True, + ): + super().__init__() + + self.conv = ModulatedConv2d( + in_channel, + out_channel, + kernel_size, + style_dim, + upsample=upsample, + blur_kernel=blur_kernel, + demodulate=demodulate, + ) + + self.noise = NoiseInjection() + # self.bias = nn.Parameter(torch.zeros(1, out_channel, 1, 1)) + # self.activate = ScaledLeakyReLU(0.2) + self.activate = FusedLeakyReLU(out_channel) + + def forward(self, input, style, noise=None, input_is_stylespace=False): + out, style = self.conv(input, style, input_is_stylespace=input_is_stylespace) + out = self.noise(out, noise=noise) + # out = out + self.bias + out = self.activate(out) + + return out, style + + +class ToRGB(nn.Module): + def __init__(self, in_channel, style_dim, upsample=True, blur_kernel=[1, 3, 3, 1]): + super().__init__() + + if upsample: + self.upsample = Upsample(blur_kernel) + + self.conv = ModulatedConv2d(in_channel, 3, 1, style_dim, demodulate=False) + self.bias = nn.Parameter(torch.zeros(1, 3, 1, 1)) + + def forward(self, input, style, skip=None, input_is_stylespace=False): + out, style = self.conv(input, style, input_is_stylespace=input_is_stylespace) + out = out + self.bias + + if skip is not None: + skip = self.upsample(skip) + + out = out + skip + + return out, style + + +class Generator(nn.Module): + def __init__( + self, + size, + style_dim, + n_mlp, + channel_multiplier=2, + blur_kernel=[1, 3, 3, 1], + lr_mlp=0.01, + ): + super().__init__() + + self.size = size + + self.style_dim = style_dim + + layers = [PixelNorm()] + + for i in range(n_mlp): + layers.append( + EqualLinear( + style_dim, style_dim, lr_mul=lr_mlp, activation='fused_lrelu' + ) + ) + + self.style = nn.Sequential(*layers) + + self.channels = { + 4: 512, + 8: 512, + 16: 512, + 32: 512, + 64: 256 * channel_multiplier, + 128: 128 * channel_multiplier, + 256: 64 * channel_multiplier, + 512: 32 * channel_multiplier, + 1024: 16 * channel_multiplier, + } + + self.input = ConstantInput(self.channels[4]) + self.conv1 = StyledConv( + self.channels[4], self.channels[4], 3, style_dim, blur_kernel=blur_kernel + ) + self.to_rgb1 = ToRGB(self.channels[4], style_dim, upsample=False) + + self.log_size = int(math.log(size, 2)) + self.num_layers = (self.log_size - 2) * 2 + 1 + + self.convs = nn.ModuleList() + self.upsamples = nn.ModuleList() + self.to_rgbs = nn.ModuleList() + self.noises = nn.Module() + + in_channel = self.channels[4] + + for layer_idx in range(self.num_layers): + res = (layer_idx + 5) // 2 + shape = [1, 1, 2 ** res, 2 ** res] + self.noises.register_buffer(f'noise_{layer_idx}', torch.randn(*shape)) + + for i in range(3, self.log_size + 1): + out_channel = self.channels[2 ** i] + + self.convs.append( + StyledConv( + in_channel, + out_channel, + 3, + style_dim, + upsample=True, + blur_kernel=blur_kernel, + ) + ) + + self.convs.append( + StyledConv( + out_channel, out_channel, 3, style_dim, blur_kernel=blur_kernel + ) + ) + + self.to_rgbs.append(ToRGB(out_channel, style_dim)) + + in_channel = out_channel + + self.n_latent = self.log_size * 2 - 2 + + def make_noise(self): + device = self.input.input.device + + noises = [torch.randn(1, 1, 2 ** 2, 2 ** 2, device=device)] + + for i in range(3, self.log_size + 1): + for _ in range(2): + noises.append(torch.randn(1, 1, 2 ** i, 2 ** i, device=device)) + + return noises + + def mean_latent(self, n_latent): + latent_in = torch.randn( + n_latent, self.style_dim, device=self.input.input.device + ) + latent = self.style(latent_in).mean(0, keepdim=True) + + return latent + + def get_latent(self, input): + return self.style(input) + + def forward( + self, + styles, + return_latents=False, + inject_index=None, + truncation=1, + truncation_latent=None, + input_is_latent=False, + input_is_stylespace=False, + noise=None, + randomize_noise=True, + ): + if not input_is_latent and not input_is_stylespace: + styles = [self.style(s) for s in styles] + + if noise is None: + if randomize_noise: + noise = [None] * self.num_layers + else: + noise = [ + getattr(self.noises, f'noise_{i}') for i in range(self.num_layers) + ] + + if truncation < 1 and not input_is_stylespace: + style_t = [] + + for style in styles: + style_t.append( + truncation_latent + truncation * (style - truncation_latent) + ) + + styles = style_t + + if input_is_stylespace: + latent = styles[0] + elif len(styles) < 2: + inject_index = self.n_latent + + if styles[0].ndim < 3: + latent = styles[0].unsqueeze(1).repeat(1, inject_index, 1) + + else: + latent = styles[0] + + else: + if inject_index is None: + inject_index = random.randint(1, self.n_latent - 1) + + latent = styles[0].unsqueeze(1).repeat(1, inject_index, 1) + latent2 = styles[1].unsqueeze(1).repeat(1, self.n_latent - inject_index, 1) + + latent = torch.cat([latent, latent2], 1) + + + style_vector = [] + + if not input_is_stylespace: + out = self.input(latent) + out, out_style = self.conv1(out, latent[:, 0], noise=noise[0]) + style_vector.append(out_style) + + skip, out_style = self.to_rgb1(out, latent[:, 1]) + style_vector.append(out_style) + + i = 1 + else: + out = self.input(latent[0]) + out, out_style = self.conv1(out, latent[0], noise=noise[0], input_is_stylespace=input_is_stylespace) + style_vector.append(out_style) + + skip, out_style = self.to_rgb1(out, latent[1], input_is_stylespace=input_is_stylespace) + style_vector.append(out_style) + + i = 2 + + for conv1, conv2, noise1, noise2, to_rgb in zip( + self.convs[::2], self.convs[1::2], noise[1::2], noise[2::2], self.to_rgbs + ): + if not input_is_stylespace: + out, out_style1 = conv1(out, latent[:, i], noise=noise1) + out, out_style2 = conv2(out, latent[:, i + 1], noise=noise2) + skip, rgb_style = to_rgb(out, latent[:, i + 2], skip) + + style_vector.extend([out_style1, out_style2, rgb_style]) + + i += 2 + else: + out, out_style1 = conv1(out, latent[i], noise=noise1, input_is_stylespace=input_is_stylespace) + out, out_style2 = conv2(out, latent[i + 1], noise=noise2, input_is_stylespace=input_is_stylespace) + skip, rgb_style = to_rgb(out, latent[i + 2], skip, input_is_stylespace=input_is_stylespace) + + style_vector.extend([out_style1, out_style2, rgb_style]) + + i += 3 + + image = skip + + if return_latents: + return image, latent, style_vector + + else: + return image, None + + +class ConvLayer(nn.Sequential): + def __init__( + self, + in_channel, + out_channel, + kernel_size, + downsample=False, + blur_kernel=[1, 3, 3, 1], + bias=True, + activate=True, + ): + layers = [] + + if downsample: + factor = 2 + p = (len(blur_kernel) - factor) + (kernel_size - 1) + pad0 = (p + 1) // 2 + pad1 = p // 2 + + layers.append(Blur(blur_kernel, pad=(pad0, pad1))) + + stride = 2 + self.padding = 0 + + else: + stride = 1 + self.padding = kernel_size // 2 + + layers.append( + EqualConv2d( + in_channel, + out_channel, + kernel_size, + padding=self.padding, + stride=stride, + bias=bias and not activate, + ) + ) + + if activate: + if bias: + layers.append(FusedLeakyReLU(out_channel)) + + else: + layers.append(ScaledLeakyReLU(0.2)) + + super().__init__(*layers) + + +class ResBlock(nn.Module): + def __init__(self, in_channel, out_channel, blur_kernel=[1, 3, 3, 1]): + super().__init__() + + self.conv1 = ConvLayer(in_channel, in_channel, 3) + self.conv2 = ConvLayer(in_channel, out_channel, 3, downsample=True) + + self.skip = ConvLayer( + in_channel, out_channel, 1, downsample=True, activate=False, bias=False + ) + + def forward(self, input): + out = self.conv1(input) + out = self.conv2(out) + + skip = self.skip(input) + out = (out + skip) / math.sqrt(2) + + return out + + +class Discriminator(nn.Module): + def __init__(self, size, channel_multiplier=2, blur_kernel=[1, 3, 3, 1]): + super().__init__() + + channels = { + 4: 512, + 8: 512, + 16: 512, + 32: 512, + 64: 256 * channel_multiplier, + 128: 128 * channel_multiplier, + 256: 64 * channel_multiplier, + 512: 32 * channel_multiplier, + 1024: 16 * channel_multiplier, + } + + convs = [ConvLayer(3, channels[size], 1)] + + log_size = int(math.log(size, 2)) + + in_channel = channels[size] + + for i in range(log_size, 2, -1): + out_channel = channels[2 ** (i - 1)] + + convs.append(ResBlock(in_channel, out_channel, blur_kernel)) + + in_channel = out_channel + + self.convs = nn.Sequential(*convs) + + self.stddev_group = 4 + self.stddev_feat = 1 + + self.final_conv = ConvLayer(in_channel + 1, channels[4], 3) + self.final_linear = nn.Sequential( + EqualLinear(channels[4] * 4 * 4, channels[4], activation='fused_lrelu'), + EqualLinear(channels[4], 1), + ) + + def forward(self, input): + out = self.convs(input) + + batch, channel, height, width = out.shape + group = min(batch, self.stddev_group) + stddev = out.view( + group, -1, self.stddev_feat, channel // self.stddev_feat, height, width + ) + stddev = torch.sqrt(stddev.var(0, unbiased=False) + 1e-8) + stddev = stddev.mean([2, 3, 4], keepdims=True).squeeze(2) + stddev = stddev.repeat(group, 1, height, width) + out = torch.cat([out, stddev], 1) + + out = self.final_conv(out) + + out = out.view(batch, -1) + out = self.final_linear(out) + + return out + diff --git a/stylegan/op/__init__.py b/stylegan/op/__init__.py new file mode 100755 index 0000000..d0918d9 --- /dev/null +++ b/stylegan/op/__init__.py @@ -0,0 +1,2 @@ +from .fused_act import FusedLeakyReLU, fused_leaky_relu +from .upfirdn2d import upfirdn2d diff --git a/stylegan/op/fused_act.py b/stylegan/op/fused_act.py new file mode 100755 index 0000000..2d575bc --- /dev/null +++ b/stylegan/op/fused_act.py @@ -0,0 +1,40 @@ +import os + +import torch +from torch import nn +from torch.nn import functional as F + +module_path = os.path.dirname(__file__) + + + +class FusedLeakyReLU(nn.Module): + def __init__(self, channel, negative_slope=0.2, scale=2 ** 0.5): + super().__init__() + + self.bias = nn.Parameter(torch.zeros(channel)) + self.negative_slope = negative_slope + self.scale = scale + + def forward(self, input): + return fused_leaky_relu(input, self.bias, self.negative_slope, self.scale) + + +def fused_leaky_relu(input, bias, negative_slope=0.2, scale=2 ** 0.5): + rest_dim = [1] * (input.ndim - bias.ndim - 1) + input = input.cuda() + if input.ndim == 3: + return ( + F.leaky_relu( + input + bias.view(1, *rest_dim, bias.shape[0]), negative_slope=negative_slope + ) + * scale + ) + else: + return ( + F.leaky_relu( + input + bias.view(1, bias.shape[0], *rest_dim), negative_slope=negative_slope + ) + * scale + ) + diff --git a/stylegan/op/upfirdn2d.py b/stylegan/op/upfirdn2d.py new file mode 100755 index 0000000..02fc25a --- /dev/null +++ b/stylegan/op/upfirdn2d.py @@ -0,0 +1,60 @@ +import os + +import torch +from torch.nn import functional as F + + +module_path = os.path.dirname(__file__) + + + +def upfirdn2d(input, kernel, up=1, down=1, pad=(0, 0)): + out = upfirdn2d_native( + input, kernel, up, up, down, down, pad[0], pad[1], pad[0], pad[1] + ) + + return out + + +def upfirdn2d_native( + input, kernel, up_x, up_y, down_x, down_y, pad_x0, pad_x1, pad_y0, pad_y1 +): + _, channel, in_h, in_w = input.shape + input = input.reshape(-1, in_h, in_w, 1) + + _, in_h, in_w, minor = input.shape + kernel_h, kernel_w = kernel.shape + + out = input.view(-1, in_h, 1, in_w, 1, minor) + out = F.pad(out, [0, 0, 0, up_x - 1, 0, 0, 0, up_y - 1]) + out = out.view(-1, in_h * up_y, in_w * up_x, minor) + + out = F.pad( + out, [0, 0, max(pad_x0, 0), max(pad_x1, 0), max(pad_y0, 0), max(pad_y1, 0)] + ) + out = out[ + :, + max(-pad_y0, 0) : out.shape[1] - max(-pad_y1, 0), + max(-pad_x0, 0) : out.shape[2] - max(-pad_x1, 0), + :, + ] + + out = out.permute(0, 3, 1, 2) + out = out.reshape( + [-1, 1, in_h * up_y + pad_y0 + pad_y1, in_w * up_x + pad_x0 + pad_x1] + ) + w = torch.flip(kernel, [0, 1]).view(1, 1, kernel_h, kernel_w) + out = F.conv2d(out, w) + out = out.reshape( + -1, + minor, + in_h * up_y + pad_y0 + pad_y1 - kernel_h + 1, + in_w * up_x + pad_x0 + pad_x1 - kernel_w + 1, + ) + out = out.permute(0, 2, 3, 1) + out = out[:, ::down_y, ::down_x, :] + + out_h = (in_h * up_y + pad_y0 + pad_y1 - kernel_h) // down_y + 1 + out_w = (in_w * up_x + pad_x0 + pad_x1 - kernel_w) // down_x + 1 + + return out.view(-1, channel, out_h, out_w) \ No newline at end of file diff --git a/stylegan/stylegan2_generator.py b/stylegan/stylegan2_generator.py deleted file mode 100644 index efaeee5..0000000 --- a/stylegan/stylegan2_generator.py +++ /dev/null @@ -1,1004 +0,0 @@ -# python3.7 -"""Contains the implementation of generator described in StyleGAN2. - -Compared to that of StyleGAN, the generator in StyleGAN2 mainly introduces style -demodulation, adds skip connections, increases model size, and disables -progressive growth. This script ONLY supports config F in the original paper. - -Paper: https://arxiv.org/pdf/1912.04958.pdf - -Official TensorFlow implementation: https://github.com/NVlabs/stylegan2 -""" - -import numpy as np - -import torch -import torch.nn as nn -import torch.nn.functional as F - -from .sync_op import all_gather - -__all__ = ['StyleGAN2Generator'] - -# Resolutions allowed. -_RESOLUTIONS_ALLOWED = [8, 16, 32, 64, 128, 256, 512, 1024] - -# Initial resolution. -_INIT_RES = 4 - -# Architectures allowed. -_ARCHITECTURES_ALLOWED = ['resnet', 'skip', 'origin'] - -# Default gain factor for weight scaling. -_WSCALE_GAIN = 1.0 - - -class StyleGAN2Generator(nn.Module): - """Defines the generator network in StyleGAN2. - - NOTE: The synthesized images are with `RGB` channel order and pixel range - [-1, 1]. - - Settings for the mapping network: - - (1) z_space_dim: Dimension of the input latent space, Z. (default: 512) - (2) w_space_dim: Dimension of the outout latent space, W. (default: 512) - (3) label_size: Size of the additional label for conditional generation. - (default: 0) - (4)mapping_layers: Number of layers of the mapping network. (default: 8) - (5) mapping_fmaps: Number of hidden channels of the mapping network. - (default: 512) - (6) mapping_lr_mul: Learning rate multiplier for the mapping network. - (default: 0.01) - (7) repeat_w: Repeat w-code for different layers. - - Settings for the synthesis network: - - (1) resolution: The resolution of the output image. - (2) image_channels: Number of channels of the output image. (default: 3) - (3) final_tanh: Whether to use `tanh` to control the final pixel range. - (default: False) - (4) const_input: Whether to use a constant in the first convolutional layer. - (default: True) - (5) architecture: Type of architecture. Support `origin`, `skip`, and - `resnet`. (default: `resnet`) - (6) fused_modulate: Whether to fuse `style_modulate` and `conv2d` together. - (default: True) - (7) demodulate: Whether to perform style demodulation. (default: True) - (8) use_wscale: Whether to use weight scaling. (default: True) - (9) fmaps_base: Factor to control number of feature maps for each layer. - (default: 16 << 10) - (10) fmaps_max: Maximum number of feature maps in each layer. (default: 512) - """ - - def __init__(self, - resolution, - z_space_dim=512, - w_space_dim=512, - label_size=0, - mapping_layers=8, - mapping_fmaps=512, - mapping_lr_mul=0.01, - repeat_w=True, - image_channels=3, - final_tanh=False, - const_input=True, - architecture='skip', - fused_modulate=True, - demodulate=True, - use_wscale=True, - fmaps_base=32 << 10, - fmaps_max=512): - """Initializes with basic settings. - - Raises: - ValueError: If the `resolution` is not supported, or `architecture` - is not supported. - """ - super().__init__() - - if resolution not in _RESOLUTIONS_ALLOWED: - raise ValueError(f'Invalid resolution: `{resolution}`!\n' - f'Resolutions allowed: {_RESOLUTIONS_ALLOWED}.') - if architecture not in _ARCHITECTURES_ALLOWED: - raise ValueError(f'Invalid architecture: `{architecture}`!\n' - f'Architectures allowed: ' - f'{_ARCHITECTURES_ALLOWED}.') - - self.init_res = _INIT_RES - self.resolution = resolution - self.z_space_dim = z_space_dim - self.w_space_dim = w_space_dim - self.label_size = label_size - self.mapping_layers = mapping_layers - self.mapping_fmaps = mapping_fmaps - self.mapping_lr_mul = mapping_lr_mul - self.repeat_w = repeat_w - self.image_channels = image_channels - self.final_tanh = final_tanh - self.const_input = const_input - self.architecture = architecture - self.fused_modulate = fused_modulate - self.demodulate = demodulate - self.use_wscale = use_wscale - self.fmaps_base = fmaps_base - self.fmaps_max = fmaps_max - - self.num_layers = int(np.log2(self.resolution // self.init_res * 2)) * 2 - - if self.repeat_w: - self.mapping_space_dim = self.w_space_dim - else: - self.mapping_space_dim = self.w_space_dim * self.num_layers - self.mapping = MappingModule(input_space_dim=self.z_space_dim, - hidden_space_dim=self.mapping_fmaps, - final_space_dim=self.mapping_space_dim, - label_size=self.label_size, - num_layers=self.mapping_layers, - use_wscale=self.use_wscale, - lr_mul=self.mapping_lr_mul) - - self.truncation = TruncationModule(w_space_dim=self.w_space_dim, - num_layers=self.num_layers, - repeat_w=self.repeat_w) - - self.synthesis = SynthesisModule(resolution=self.resolution, - init_resolution=self.init_res, - w_space_dim=self.w_space_dim, - image_channels=self.image_channels, - final_tanh=self.final_tanh, - const_input=self.const_input, - architecture=self.architecture, - fused_modulate=self.fused_modulate, - demodulate=self.demodulate, - use_wscale=self.use_wscale, - fmaps_base=self.fmaps_base, - fmaps_max=self.fmaps_max) - - self.pth_to_tf_var_mapping = {} - for key, val in self.mapping.pth_to_tf_var_mapping.items(): - self.pth_to_tf_var_mapping[f'mapping.{key}'] = val - for key, val in self.truncation.pth_to_tf_var_mapping.items(): - self.pth_to_tf_var_mapping[f'truncation.{key}'] = val - for key, val in self.synthesis.pth_to_tf_var_mapping.items(): - self.pth_to_tf_var_mapping[f'synthesis.{key}'] = val - - def forward(self, - z, - label=None, - w_moving_decay=0.995, - style_mixing_prob=0.9, - trunc_psi=None, - trunc_layers=None, - randomize_noise=False, - **_unused_kwargs): - mapping_results = self.mapping(z, label) - w = mapping_results['w'] - - if self.training and w_moving_decay < 1: - batch_w_avg = all_gather(w).mean(dim=0) - self.truncation.w_avg.copy_( - self.truncation.w_avg * w_moving_decay + - batch_w_avg * (1 - w_moving_decay)) - - if self.training and style_mixing_prob > 0: - new_z = torch.randn_like(z) - new_w = self.mapping(new_z, label)['w'] - if np.random.uniform() < style_mixing_prob: - mixing_cutoff = np.random.randint(1, self.num_layers) - w = self.truncation(w) - new_w = self.truncation(new_w) - w[:, :mixing_cutoff] = new_w[:, :mixing_cutoff] - - wp = self.truncation(w, trunc_psi, trunc_layers) - synthesis_results = self.synthesis(wp, randomize_noise) - - return {**mapping_results, **synthesis_results} - - -class MappingModule(nn.Module): - """Implements the latent space mapping module. - - Basically, this module executes several dense layers in sequence. - """ - - def __init__(self, - input_space_dim=512, - hidden_space_dim=512, - final_space_dim=512, - label_size=0, - num_layers=8, - normalize_input=True, - use_wscale=True, - lr_mul=0.01): - super().__init__() - - self.input_space_dim = input_space_dim - self.hidden_space_dim = hidden_space_dim - self.final_space_dim = final_space_dim - self.label_size = label_size - self.num_layers = num_layers - self.normalize_input = normalize_input - self.use_wscale = use_wscale - self.lr_mul = lr_mul - - self.norm = PixelNormLayer() if self.normalize_input else nn.Identity() - - self.pth_to_tf_var_mapping = {} - for i in range(num_layers): - dim_mul = 2 if label_size else 1 - in_channels = (input_space_dim * dim_mul if i == 0 else - hidden_space_dim) - out_channels = (final_space_dim if i == (num_layers - 1) else - hidden_space_dim) - self.add_module(f'dense{i}', - DenseBlock(in_channels=in_channels, - out_channels=out_channels, - use_wscale=self.use_wscale, - lr_mul=self.lr_mul)) - self.pth_to_tf_var_mapping[f'dense{i}.weight'] = f'Dense{i}/weight' - self.pth_to_tf_var_mapping[f'dense{i}.bias'] = f'Dense{i}/bias' - if label_size: - self.label_weight = nn.Parameter( - torch.randn(label_size, input_space_dim)) - self.pth_to_tf_var_mapping[f'label_weight'] = f'LabelConcat/weight' - - def forward(self, z, label=None): - if z.ndim != 2 or z.shape[1] != self.input_space_dim: - raise ValueError(f'Input latent code should be with shape ' - f'[batch_size, input_dim], where ' - f'`input_dim` equals to {self.input_space_dim}!\n' - f'But `{z.shape}` is received!') - if self.label_size: - if label is None: - raise ValueError(f'Model requires an additional label ' - f'(with size {self.label_size}) as input, ' - f'but no label is received!') - if label.ndim != 2 or label.shape != (z.shape[0], self.label_size): - raise ValueError(f'Input label should be with shape ' - f'[batch_size, label_size], where ' - f'`batch_size` equals to that of ' - f'latent codes ({z.shape[0]}) and ' - f'`label_size` equals to {self.label_size}!\n' - f'But `{label.shape}` is received!') - embedding = torch.matmul(label, self.label_weight) - z = torch.cat((z, embedding), dim=1) - - z = self.norm(z) - w = z - for i in range(self.num_layers): - w = self.__getattr__(f'dense{i}')(w) - results = { - 'z': z, - 'label': label, - 'w': w, - } - if self.label_size: - results['embedding'] = embedding - return results - - -class TruncationModule(nn.Module): - """Implements the truncation module. - - Truncation is executed as follows: - - For layers in range [0, truncation_layers), the truncated w-code is computed - as - - w_new = w_avg + (w - w_avg) * truncation_psi - - To disable truncation, please set - (1) truncation_psi = 1.0 (None) OR - (2) truncation_layers = 0 (None) - - NOTE: The returned tensor is layer-wise style codes. - """ - - def __init__(self, w_space_dim, num_layers, repeat_w=True): - super().__init__() - - self.num_layers = num_layers - self.w_space_dim = w_space_dim - self.repeat_w = repeat_w - - if self.repeat_w: - self.register_buffer('w_avg', torch.zeros(w_space_dim)) - else: - self.register_buffer('w_avg', torch.zeros(num_layers * w_space_dim)) - self.pth_to_tf_var_mapping = {'w_avg': 'dlatent_avg'} - - def forward(self, w, trunc_psi=None, trunc_layers=None): - if w.ndim == 2: - if self.repeat_w and w.shape[1] == self.w_space_dim: - w = w.view(-1, 1, self.w_space_dim) - wp = w.repeat(1, self.num_layers, 1) - else: - assert w.shape[1] == self.w_space_dim * self.num_layers - wp = w.view(-1, self.num_layers, self.w_space_dim) - else: - wp = w - assert wp.ndim == 3 - assert wp.shape[1:] == (self.num_layers, self.w_space_dim) - - trunc_psi = 1.0 if trunc_psi is None else trunc_psi - trunc_layers = 0 if trunc_layers is None else trunc_layers - if trunc_psi < 1.0 and trunc_layers > 0: - layer_idx = np.arange(self.num_layers).reshape(1, -1, 1) - coefs = np.ones_like(layer_idx, dtype=np.float32) - coefs[layer_idx < trunc_layers] *= trunc_psi - coefs = torch.from_numpy(coefs).to(wp) - w_avg = self.w_avg.view(1, -1, self.w_space_dim) - wp = w_avg + (wp - w_avg) * coefs - return wp - - -class SynthesisModule(nn.Module): - """Implements the image synthesis module. - - Basically, this module executes several convolutional layers in sequence. - """ - - def __init__(self, - resolution=1024, - init_resolution=4, - w_space_dim=512, - image_channels=3, - final_tanh=False, - const_input=True, - architecture='skip', - fused_modulate=True, - demodulate=True, - use_wscale=True, - fmaps_base=32 << 10, - fmaps_max=512): - super().__init__() - - self.init_res = init_resolution - self.init_res_log2 = int(np.log2(self.init_res)) - self.resolution = resolution - self.final_res_log2 = int(np.log2(self.resolution)) - self.w_space_dim = w_space_dim - self.image_channels = image_channels - self.final_tanh = final_tanh - self.const_input = const_input - self.architecture = architecture - self.fused_modulate = fused_modulate - self.demodulate = demodulate - self.use_wscale = use_wscale - self.fmaps_base = fmaps_base - self.fmaps_max = fmaps_max - - self.num_layers = (self.final_res_log2 - self.init_res_log2 + 1) * 2 - - self.pth_to_tf_var_mapping = {} - for res_log2 in range(self.init_res_log2, self.final_res_log2 + 1): - res = 2 ** res_log2 - block_idx = res_log2 - self.init_res_log2 - - # First convolution layer for each resolution. - if res == self.init_res: - if self.const_input: - self.add_module(f'early_layer', - InputBlock(init_resolution=self.init_res, - channels=self.get_nf(res))) - self.pth_to_tf_var_mapping[f'early_layer.const'] = ( - f'{res}x{res}/Const/const') - else: - self.add_module(f'early_layer', - DenseBlock(in_channels=self.w_space_dim, - out_channels=self.get_nf(res), - use_wscale=self.use_wscale)) - self.pth_to_tf_var_mapping[f'early_layer.weight'] = ( - f'{res}x{res}/Dense/weight') - self.pth_to_tf_var_mapping[f'early_layer.bias'] = ( - f'{res}x{res}/Dense/bias') - else: - layer_name = f'layer{2 * block_idx - 1}' - self.add_module( - layer_name, - ModulateConvBlock(in_channels=self.get_nf(res // 2), - out_channels=self.get_nf(res), - resolution=res, - w_space_dim=self.w_space_dim, - scale_factor=2, - fused_modulate=self.fused_modulate, - demodulate=self.demodulate, - use_wscale=self.use_wscale)) - self.pth_to_tf_var_mapping[f'{layer_name}.weight'] = ( - f'{res}x{res}/Conv0_up/weight') - self.pth_to_tf_var_mapping[f'{layer_name}.bias'] = ( - f'{res}x{res}/Conv0_up/bias') - self.pth_to_tf_var_mapping[f'{layer_name}.style.weight'] = ( - f'{res}x{res}/Conv0_up/mod_weight') - self.pth_to_tf_var_mapping[f'{layer_name}.style.bias'] = ( - f'{res}x{res}/Conv0_up/mod_bias') - self.pth_to_tf_var_mapping[f'{layer_name}.noise_strength'] = ( - f'{res}x{res}/Conv0_up/noise_strength') - self.pth_to_tf_var_mapping[f'{layer_name}.noise'] = ( - f'noise{2 * block_idx - 1}') - - if self.architecture == 'resnet': - layer_name = f'layer{2 * block_idx - 1}' - self.add_module( - layer_name, - ConvBlock(in_channels=self.get_nf(res // 2), - out_channels=self.get_nf(res), - kernel_size=1, - add_bias=False, - scale_factor=2, - use_wscale=self.use_wscale, - activation_type='linear')) - self.pth_to_tf_var_mapping[f'{layer_name}.weight'] = ( - f'{res}x{res}/Skip/weight') - - # Second convolution layer for each resolution. - layer_name = f'layer{2 * block_idx}' - self.add_module( - layer_name, - ModulateConvBlock(in_channels=self.get_nf(res), - out_channels=self.get_nf(res), - resolution=res, - w_space_dim=self.w_space_dim, - fused_modulate=self.fused_modulate, - demodulate=self.demodulate, - use_wscale=self.use_wscale)) - tf_layer_name = 'Conv' if res == self.init_res else 'Conv1' - self.pth_to_tf_var_mapping[f'{layer_name}.weight'] = ( - f'{res}x{res}/{tf_layer_name}/weight') - self.pth_to_tf_var_mapping[f'{layer_name}.bias'] = ( - f'{res}x{res}/{tf_layer_name}/bias') - self.pth_to_tf_var_mapping[f'{layer_name}.style.weight'] = ( - f'{res}x{res}/{tf_layer_name}/mod_weight') - self.pth_to_tf_var_mapping[f'{layer_name}.style.bias'] = ( - f'{res}x{res}/{tf_layer_name}/mod_bias') - self.pth_to_tf_var_mapping[f'{layer_name}.noise_strength'] = ( - f'{res}x{res}/{tf_layer_name}/noise_strength') - self.pth_to_tf_var_mapping[f'{layer_name}.noise'] = ( - f'noise{2 * block_idx}') - - # Output convolution layer for each resolution (if needed). - if res_log2 == self.final_res_log2 or self.architecture == 'skip': - layer_name = f'output{block_idx}' - self.add_module( - layer_name, - ModulateConvBlock(in_channels=self.get_nf(res), - out_channels=image_channels, - resolution=res, - w_space_dim=self.w_space_dim, - kernel_size=1, - fused_modulate=self.fused_modulate, - demodulate=False, - use_wscale=self.use_wscale, - add_noise=False, - activation_type='linear')) - self.pth_to_tf_var_mapping[f'{layer_name}.weight'] = ( - f'{res}x{res}/ToRGB/weight') - self.pth_to_tf_var_mapping[f'{layer_name}.bias'] = ( - f'{res}x{res}/ToRGB/bias') - self.pth_to_tf_var_mapping[f'{layer_name}.style.weight'] = ( - f'{res}x{res}/ToRGB/mod_weight') - self.pth_to_tf_var_mapping[f'{layer_name}.style.bias'] = ( - f'{res}x{res}/ToRGB/mod_bias') - - if self.architecture == 'skip': - self.upsample = UpsamplingLayer() - self.final_activate = nn.Tanh() if final_tanh else nn.Identity() - - def get_nf(self, res): - """Gets number of feature maps according to current resolution.""" - return min(self.fmaps_base // res, self.fmaps_max) - - def forward(self, wp, randomize_noise=False, - basecode_layer=None, basecode=None): - if wp.ndim != 3 or wp.shape[1:] != (self.num_layers, self.w_space_dim): - raise ValueError(f'Input tensor should be with shape ' - f'[batch_size, num_layers, w_space_dim], where ' - f'`num_layers` equals to {self.num_layers}, and ' - f'`w_space_dim` equals to {self.w_space_dim}!\n' - f'But `{wp.shape}` is received!') - - results = {'wp': wp} - x = self.early_layer(wp[:, 0]) - if self.architecture == 'origin': - for layer_idx in range(self.num_layers - 1): - x, style = self.__getattr__(f'layer{layer_idx}')( - x, wp[:, layer_idx], randomize_noise) - results[f'style{layer_idx:02d}'] = style - image, style = self.__getattr__(f'output{layer_idx // 2}')( - x, wp[:, layer_idx + 1]) - results[f'output_style{layer_idx // 2}'] = style - elif self.architecture == 'skip': - for layer_idx in range(self.num_layers - 1): - x, style = self.__getattr__(f'layer{layer_idx}')(x, wp[:, layer_idx], randomize_noise) - results[f'x{layer_idx:02d}'] = x - results[f'style{layer_idx:02d}'] = style - - if basecode_layer == f'x{layer_idx:02d}': - x = basecode.contiguous() - - if layer_idx % 2 == 0: - temp, style = self.__getattr__(f'output{layer_idx // 2}')(x, wp[:, layer_idx + 1]) - results[f'output_style{layer_idx // 2}'] = style - if layer_idx == 0: - image = temp - else: - if basecode_layer == f'x{layer_idx-1:02d}': - image = temp - else: - image = temp + self.upsample(image) - - elif self.architecture == 'resnet': - x, style = self.layer0(x) - results[f'style00'] = style - for layer_idx in range(1, self.num_layers - 1, 2): - residual = self.__getattr__(f'skip_layer{layer_idx // 2}')(x) - x, style = self.__getattr__(f'layer{layer_idx}')( - x, wp[:, layer_idx], randomize_noise) - results[f'style{layer_idx:02d}'] = style - x, style = self.__getattr__(f'layer{layer_idx + 1}')( - x, wp[:, layer_idx + 1], randomize_noise) - results[f'style{layer_idx + 1:02d}'] = style - x = (x + residual) / np.sqrt(2.0) - image, style = self.__getattr__(f'output{layer_idx // 2 + 1}')( - x, wp[:, layer_idx + 2]) - results[f'output_style{layer_idx // 2}'] = style - results['image'] = self.final_activate(image) - return results - - -class PixelNormLayer(nn.Module): - """Implements pixel-wise feature vector normalization layer.""" - - def __init__(self, dim=1, epsilon=1e-8): - super().__init__() - self.dim = dim - self.eps = epsilon - - def forward(self, x): - norm = torch.sqrt( - torch.mean(x ** 2, dim=self.dim, keepdim=True) + self.eps) - return x / norm - - -class UpsamplingLayer(nn.Module): - """Implements the upsampling layer. - - This layer can also be used as filtering by setting `scale_factor` as 1. - """ - - def __init__(self, - scale_factor=2, - kernel=(1, 3, 3, 1), - extra_padding=0, - kernel_gain=None): - super().__init__() - assert scale_factor >= 1 - self.scale_factor = scale_factor - - if extra_padding != 0: - assert scale_factor == 1 - - if kernel is None: - kernel = np.ones((scale_factor), dtype=np.float32) - else: - kernel = np.array(kernel, dtype=np.float32) - assert kernel.ndim == 1 - kernel = np.outer(kernel, kernel) - kernel = kernel / np.sum(kernel) - if kernel_gain is None: - kernel = kernel * (scale_factor ** 2) - else: - assert kernel_gain > 0 - kernel = kernel * (kernel_gain ** 2) - assert kernel.ndim == 2 - assert kernel.shape[0] == kernel.shape[1] - kernel = kernel[np.newaxis, np.newaxis] - self.register_buffer('kernel', torch.from_numpy(kernel)) - self.kernel = self.kernel.flip(0, 1) - - self.upsample_padding = (0, scale_factor - 1, # Width padding. - 0, 0, # Width. - 0, scale_factor - 1, # Height padding. - 0, 0, # Height. - 0, 0, # Channel. - 0, 0) # Batch size. - - padding = kernel.shape[2] - scale_factor + extra_padding - self.padding = ((padding + 1) // 2 + scale_factor - 1, padding // 2, - (padding + 1) // 2 + scale_factor - 1, padding // 2) - - def forward(self, x): - assert x.ndim == 4 - channels = x.shape[1] - if self.scale_factor > 1: - x = x.view(-1, channels, x.shape[2], 1, x.shape[3], 1) - x = F.pad(x, self.upsample_padding, mode='constant', value=0) - x = x.view(-1, channels, x.shape[2] * self.scale_factor, - x.shape[4] * self.scale_factor) - x = x.view(-1, 1, x.shape[2], x.shape[3]) - x = F.pad(x, self.padding, mode='constant', value=0) - x = F.conv2d(x, self.kernel, stride=1) - x = x.view(-1, channels, x.shape[2], x.shape[3]) - return x - - -class InputBlock(nn.Module): - """Implements the input block. - - Basically, this block starts from a const input, which is with shape - `(channels, init_resolution, init_resolution)`. - """ - - def __init__(self, init_resolution, channels): - super().__init__() - self.const = nn.Parameter( - torch.randn(1, channels, init_resolution, init_resolution)) - - def forward(self, w): - x = self.const.repeat(w.shape[0], 1, 1, 1) - return x - - -class ConvBlock(nn.Module): - """Implements the convolutional block (no style modulation). - - Basically, this block executes, convolutional layer, filtering layer (if - needed), and activation layer in sequence. - - NOTE: This block is particularly used for skip-connection branch in the - `resnet` structure. - """ - - def __init__(self, - in_channels, - out_channels, - kernel_size=3, - add_bias=True, - scale_factor=1, - filtering_kernel=(1, 3, 3, 1), - use_wscale=True, - wscale_gain=_WSCALE_GAIN, - lr_mul=1.0, - activation_type='lrelu'): - """Initializes with block settings. - - Args: - in_channels: Number of channels of the input tensor. - out_channels: Number of channels of the output tensor. - kernel_size: Size of the convolutional kernels. (default: 3) - add_bias: Whether to add bias onto the convolutional result. - (default: True) - scale_factor: Scale factor for upsampling. `1` means skip - upsampling. (default: 1) - filtering_kernel: Kernel used for filtering after upsampling. - (default: (1, 3, 3, 1)) - use_wscale: Whether to use weight scaling. (default: True) - wscale_gain: Gain factor for weight scaling. (default: _WSCALE_GAIN) - lr_mul: Learning multiplier for both weight and bias. (default: 1.0) - activation_type: Type of activation. Support `linear` and `lrelu`. - (default: `lrelu`) - - Raises: - NotImplementedError: If the `activation_type` is not supported. - """ - super().__init__() - - if scale_factor > 1: - self.use_conv2d_transpose = True - extra_padding = scale_factor - kernel_size - self.filter = UpsamplingLayer(scale_factor=1, - kernel=filtering_kernel, - extra_padding=extra_padding, - kernel_gain=scale_factor) - self.stride = scale_factor - self.padding = 0 # Padding is done in `UpsamplingLayer`. - else: - self.use_conv2d_transpose = False - assert kernel_size % 2 == 1 - self.stride = 1 - self.padding = kernel_size // 2 - - weight_shape = (out_channels, in_channels, kernel_size, kernel_size) - fan_in = kernel_size * kernel_size * in_channels - wscale = wscale_gain / np.sqrt(fan_in) - if use_wscale: - self.weight = nn.Parameter(torch.randn(*weight_shape) / lr_mul) - self.wscale = wscale * lr_mul - else: - self.weight = nn.Parameter( - torch.randn(*weight_shape) * wscale / lr_mul) - self.wscale = lr_mul - - if add_bias: - self.bias = nn.Parameter(torch.zeros(out_channels)) - else: - self.bias = None - self.bscale = lr_mul - - if activation_type == 'linear': - self.activate = nn.Identity() - self.activate_scale = 1.0 - elif activation_type == 'lrelu': - self.activate = nn.LeakyReLU(negative_slope=0.2, inplace=True) - self.activate_scale = np.sqrt(2.0) - else: - raise NotImplementedError(f'Not implemented activation function: ' - f'`{activation_type}`!') - - def forward(self, x): - weight = self.weight * self.wscale - bias = self.bias * self.bscale if self.bias is not None else None - if self.use_conv2d_transpose: - weight = weight.permute(1, 0, 2, 3).flip(2, 3) - x = F.conv_transpose2d(x, - weight=weight, - bias=bias, - stride=self.scale_factor, - padding=self.padding) - x = self.filter(x) - else: - x = F.conv2d(x, - weight=weight, - bias=bias, - stride=self.stride, - padding=self.padding) - x = self.activate(x) * self.activate_scale - return x - - -class ModulateConvBlock(nn.Module): - """Implements the convolutional block with style modulation.""" - - def __init__(self, - in_channels, - out_channels, - resolution, - w_space_dim, - kernel_size=3, - add_bias=True, - scale_factor=1, - filtering_kernel=(1, 3, 3, 1), - fused_modulate=True, - demodulate=True, - use_wscale=True, - wscale_gain=_WSCALE_GAIN, - lr_mul=1.0, - add_noise=True, - activation_type='lrelu', - epsilon=1e-8): - """Initializes with block settings. - - Args: - in_channels: Number of channels of the input tensor. - out_channels: Number of channels of the output tensor. - resolution: Resolution of the output tensor. - w_space_dim: Dimension of W space for style modulation. - kernel_size: Size of the convolutional kernels. (default: 3) - add_bias: Whether to add bias onto the convolutional result. - (default: True) - scale_factor: Scale factor for upsampling. `1` means skip - upsampling. (default: 1) - filtering_kernel: Kernel used for filtering after upsampling. - (default: (1, 3, 3, 1)) - fused_modulate: Whether to fuse `style_modulate` and `conv2d` - together. (default: True) - demodulate: Whether to perform style demodulation. (default: True) - use_wscale: Whether to use weight scaling. (default: True) - wscale_gain: Gain factor for weight scaling. (default: _WSCALE_GAIN) - lr_mul: Learning multiplier for both weight and bias. (default: 1.0) - add_noise: Whether to add noise onto the output tensor. (default: - True) - activation_type: Type of activation. Support `linear` and `lrelu`. - (default: `lrelu`) - epsilon: Small number to avoid `divide by zero`. (default: 1e-8) - - Raises: - NotImplementedError: If the `activation_type` is not supported. - """ - super().__init__() - - self.res = resolution - self.in_c = in_channels - self.out_c = out_channels - self.ksize = kernel_size - self.eps = epsilon - - if scale_factor > 1: - self.use_conv2d_transpose = True - extra_padding = scale_factor - kernel_size - self.filter = UpsamplingLayer(scale_factor=1, - kernel=filtering_kernel, - extra_padding=extra_padding, - kernel_gain=scale_factor) - self.stride = scale_factor - self.padding = 0 # Padding is done in `UpsamplingLayer`. - else: - self.use_conv2d_transpose = False - assert kernel_size % 2 == 1 - self.stride = 1 - self.padding = kernel_size // 2 - - weight_shape = (out_channels, in_channels, kernel_size, kernel_size) - fan_in = kernel_size * kernel_size * in_channels - wscale = wscale_gain / np.sqrt(fan_in) - if use_wscale: - self.weight = nn.Parameter(torch.randn(*weight_shape) / lr_mul) - self.wscale = wscale * lr_mul - else: - self.weight = nn.Parameter( - torch.randn(*weight_shape) * wscale / lr_mul) - self.wscale = lr_mul - - self.style = DenseBlock(in_channels=w_space_dim, - out_channels=in_channels, - additional_bias=1.0, - use_wscale=use_wscale, - activation_type='linear') - - self.fused_modulate = fused_modulate - self.demodulate = demodulate - - if add_bias: - self.bias = nn.Parameter(torch.zeros(out_channels)) - else: - self.bias = None - self.bscale = lr_mul - - if activation_type == 'linear': - self.activate = nn.Identity() - self.activate_scale = 1.0 - elif activation_type == 'lrelu': - self.activate = nn.LeakyReLU(negative_slope=0.2, inplace=True) - self.activate_scale = np.sqrt(2.0) - else: - raise NotImplementedError(f'Not implemented activation function: ' - f'`{activation_type}`!') - - self.add_noise = add_noise - if self.add_noise: - self.register_buffer('noise', torch.randn(1, 1, self.res, self.res)) - self.noise_strength = nn.Parameter(torch.zeros(())) - - def forward(self, x, w, randomize_noise=False): - batch = x.shape[0] - - weight = self.weight * self.wscale - weight = weight.permute(2, 3, 1, 0) - - # Style modulation. - style = self.style(w) - _weight = weight.view(1, self.ksize, self.ksize, self.in_c, self.out_c) - _weight = _weight * style.view(batch, 1, 1, self.in_c, 1) - - # Style demodulation. - if self.demodulate: - _weight_norm = torch.sqrt( - torch.sum(_weight ** 2, dim=[1, 2, 3]) + self.eps) - _weight = _weight / _weight_norm.view(batch, 1, 1, 1, self.out_c) - - if self.fused_modulate: - x = x.view(1, batch * self.in_c, x.shape[2], x.shape[3]) - weight = _weight.permute(1, 2, 3, 0, 4).reshape( - self.ksize, self.ksize, self.in_c, batch * self.out_c) - else: - x = x * style.view(batch, self.in_c, 1, 1) - - if self.use_conv2d_transpose: - weight = weight.flip(0, 1) - if self.fused_modulate: - weight = weight.view( - self.ksize, self.ksize, self.in_c, batch, self.out_c) - weight = weight.permute(0, 1, 4, 3, 2) - weight = weight.reshape( - self.ksize, self.ksize, self.out_c, batch * self.in_c) - weight = weight.permute(3, 2, 0, 1) - else: - weight = weight.permute(2, 3, 0, 1) - x = F.conv_transpose2d(x, - weight=weight, - bias=None, - stride=self.stride, - padding=self.padding, - groups=(batch if self.fused_modulate else 1)) - x = self.filter(x) - else: - weight = weight.permute(3, 2, 0, 1) - x = F.conv2d(x, - weight=weight, - bias=None, - stride=self.stride, - padding=self.padding, - groups=(batch if self.fused_modulate else 1)) - - if self.fused_modulate: - x = x.view(batch, self.out_c, self.res, self.res) - elif self.demodulate: - x = x / _weight_norm.view(batch, self.out_c, 1, 1) - - if self.add_noise: - if randomize_noise: - noise = torch.randn(x.shape[0], 1, self.res, self.res).to(x) - else: - noise = self.noise - x = x + noise * self.noise_strength.view(1, 1, 1, 1) - - bias = self.bias * self.bscale if self.bias is not None else None - if bias is not None: - x = x + bias.view(1, -1, 1, 1) - x = self.activate(x) * self.activate_scale - return x, style - - -class DenseBlock(nn.Module): - """Implements the dense block. - - Basically, this block executes fully-connected layer and activation layer. - - NOTE: This layer supports adding an additional bias beyond the trainable - bias parameter. This is specially used for the mapping from the w code to - the style code. - """ - - def __init__(self, - in_channels, - out_channels, - add_bias=True, - additional_bias=0, - use_wscale=True, - wscale_gain=_WSCALE_GAIN, - lr_mul=1.0, - activation_type='lrelu'): - """Initializes with block settings. - - Args: - in_channels: Number of channels of the input tensor. - out_channels: Number of channels of the output tensor. - add_bias: Whether to add bias onto the fully-connected result. - (default: True) - additional_bias: The additional bias, which is independent from the - bias parameter. (default: 0.0) - use_wscale: Whether to use weight scaling. (default: True) - wscale_gain: Gain factor for weight scaling. (default: _WSCALE_GAIN) - lr_mul: Learning multiplier for both weight and bias. (default: 1.0) - activation_type: Type of activation. Support `linear` and `lrelu`. - (default: `lrelu`) - - Raises: - NotImplementedError: If the `activation_type` is not supported. - """ - super().__init__() - weight_shape = (out_channels, in_channels) - wscale = wscale_gain / np.sqrt(in_channels) - if use_wscale: - self.weight = nn.Parameter(torch.randn(*weight_shape) / lr_mul) - self.wscale = wscale * lr_mul - else: - self.weight = nn.Parameter( - torch.randn(*weight_shape) * wscale / lr_mul) - self.wscale = lr_mul - - if add_bias: - self.bias = nn.Parameter(torch.zeros(out_channels)) - else: - self.bias = None - self.bscale = lr_mul - self.additional_bias = additional_bias - - if activation_type == 'linear': - self.activate = nn.Identity() - self.activate_scale = 1.0 - elif activation_type == 'lrelu': - self.activate = nn.LeakyReLU(negative_slope=0.2, inplace=True) - self.activate_scale = np.sqrt(2.0) - else: - raise NotImplementedError(f'Not implemented activation function: ' - f'`{activation_type}`!') - - def forward(self, x): - if x.ndim != 2: - x = x.view(x.shape[0], -1) - bias = self.bias * self.bscale if self.bias is not None else None - x = F.linear(x, weight=self.weight * self.wscale, bias=bias) - x = self.activate(x + self.additional_bias) * self.activate_scale - return x \ No newline at end of file diff --git a/stylegan/sync_op.py b/stylegan/sync_op.py deleted file mode 100644 index 153ca56..0000000 --- a/stylegan/sync_op.py +++ /dev/null @@ -1,18 +0,0 @@ -# python3.7 -"""Contains the synchronizing operator.""" - -import torch -import torch.distributed as dist - -__all__ = ['all_gather'] - - -def all_gather(tensor): - """Gathers tensor from all devices and does averaging.""" - if not dist.is_initialized(): - return tensor - - world_size = dist.get_world_size() - tensor_list = [torch.ones_like(tensor) for _ in range(world_size)] - dist.all_gather(tensor_list, tensor, async_op=False) - return torch.mean(torch.stack(tensor_list, dim=0), dim=0) \ No newline at end of file