JerMa88's picture
Upload folder using huggingface_hub
3ce19a2 verified
Raw
History Blame Contribute Delete
4.38 kB
from math import sqrt
import torch
from torch import nn
import numpy as np
class PixelNorm(nn.Module):
def __init__(self):
super().__init__()
def forward(self, input):
return input / torch.sqrt(torch.mean(input ** 2, dim=1, keepdim=True) + 1e-6)
class FullyConnectedLayer(nn.Module):
def __init__(self, in_features, out_features, bias=True,
activation='linear', lr_multiplier=1, bias_init=0):
super().__init__()
self.activation = activation
self.weight = nn.Parameter(torch.randn([out_features, in_features]) / lr_multiplier)
self.bias = nn.Parameter(torch.full([out_features], np.float32(bias_init))) if bias else None
self.weight_gain = lr_multiplier / np.sqrt(in_features)
self.bias_gain = lr_multiplier
def forward(self, x):
w = self.weight.to(x.dtype) * self.weight_gain
b = self.bias
if b is not None:
b = b.to(x.dtype)
if self.bias_gain != 1:
b = b * self.bias_gain
x = torch.addmm(b.unsqueeze(0), x, w.t())
return x
class EqualLinear(nn.Module):
def __init__(self, in_dim, out_dim):
super().__init__()
linear = nn.Linear(in_dim, out_dim)
linear.bias.data.zero_()
self.linear = linear
def forward(self, input):
return self.linear(input)
def normalize_2nd_moment(x, dim=1, eps=1e-8):
return x * (x.square().mean(dim=dim, keepdim=True) + eps).rsqrt()
class MappingNetowrk(nn.Module):
def __init__(self, code_dim=512, n_mlp=8, mapping_lr_multiplier=1.0):
super().__init__()
layers = [PixelNorm()]
for i in range(n_mlp):
layers.append(FullyConnectedLayer(code_dim, code_dim,
lr_multiplier=mapping_lr_multiplier))
layers.append(nn.LeakyReLU(0.2))
self.style = nn.Sequential(*layers)
def forward(
self,
input,
noise=None,
step=0,
alpha=-1,
mean_style=None,
style_weight=0,
mixing_range=(-1, -1),
):
styles = []
# input = normalize_2nd_moment(input)
if type(input) not in (list, tuple):
input = [input]
for i in input:
x = self.style(i)
styles.append(x)
# batch = input[0].shape[0]
#
# if noise is None:
# noise = []
#
# for i in range(step + 1):
# size = 4 * 2 ** i
# noise.append(torch.randn(batch, 1, size, size, device=input[0].device))
# if mean_style is not None:
# styles_norm = []
#
# for style in styles:
# styles_norm.append(mean_style + style_weight * (style - mean_style))
#
# styles = styles_norm
return styles
def forward_w_trajectory(self, z):
"""One style tensor per MLP block (after each FC + activation), for video/strip viz."""
if type(z) not in (list, tuple):
z = [z]
z0 = z[0]
x = self.style[0](z0)
out_list = []
i = 1
while i < len(self.style):
x = self.style[i](x)
i += 1
if i < len(self.style):
x = self.style[i](x)
i += 1
out_list.append(x)
return out_list
# def mean_style(self, input):
# style = self.style(input).mean(0, keepdim=True)
#
# return style
class AdaptiveInstanceNorm(nn.Module):
def __init__(self, in_channel, style_dim):
super().__init__()
self.norm = nn.InstanceNorm2d(in_channel, eps=1e-3)
self.style = EqualLinear(style_dim, in_channel * 2)
nn.init.zeros_(self.style.linear.bias)
def forward(self, input, style):
style = self.style(style).unsqueeze(2).unsqueeze(3)
gamma, beta = style.chunk(2, 1)
if input.shape[-1] > 1:
out = self.norm(input)
else:
out = input
out = (1 + gamma) * out + beta
return out
class NoiseInjection(nn.Module):
def __init__(self, channel):
super().__init__()
self.weight = nn.Parameter(torch.randn(1, channel, 1, 1), requires_grad=False)
def forward(self, image, spatial_noise):
return image