JerMa88's picture
Upload folder using huggingface_hub
3ce19a2 verified
Raw
History Blame Contribute Delete
8.09 kB
import torch
from torch import nn
from torch.nn import functional as F
from mapping_network import AdaptiveInstanceNorm, MappingNetowrk
from helpers.imle_helpers import get_1x1
from collections import defaultdict
from rtm_core import RTMMappingNetwork
import numpy as np
import itertools
def parse_layer_string(s):
layers = []
for ss in s.split(','):
if 'x' in ss:
res, num = ss.split('x')
count = int(num)
layers += [(int(res), None) for _ in range(count)]
elif 'm' in ss:
res, mixin = [int(a) for a in ss.split('m')]
layers.append((res, mixin))
elif 'd' in ss:
res, down_rate = [int(a) for a in ss.split('d')]
layers.append((res, down_rate))
else:
res = int(ss)
layers.append((res, None))
return layers
def get_width_settings(width, s):
mapping = defaultdict(lambda: width)
if s:
s = s.split(',')
for ss in s:
k, v = ss.split(':')
mapping[int(k)] = int(v)
return mapping
class SEBlock(nn.Module):
def __init__(self, channels, reduction=16):
super().__init__()
self.avg_pool = nn.AdaptiveAvgPool2d(1)
self.fc = nn.Sequential(
nn.Linear(channels, channels // reduction, bias=False),
nn.ReLU(inplace=True),
nn.Linear(channels // reduction, channels, bias=False),
nn.Sigmoid()
)
def forward(self, x):
b, c, _, _ = x.size()
y = self.avg_pool(x).view(b, c)
y = self.fc(y).view(b, c, 1, 1)
return x * y.expand_as(x)
class ConvNeXtBlock(nn.Module):
def __init__(self, dim, H, expansion=4, kernel_size=7, use_se=True, reduction=16, dropout=0.0):
super().__init__()
self.dw_conv = nn.Conv2d(dim, dim, kernel_size=kernel_size, padding=kernel_size//2, groups=dim)
if(H.convnext_norm == 'layernorm'):
self.norm = nn.LayerNorm(dim, eps=H.convnext_norm_eps)
elif(H.convnext_norm == 'rmsnorm'):
self.norm = nn.RMSNorm(dim, eps=H.convnext_norm_eps)
self.pw_conv1 = nn.Linear(dim, expansion * dim)
self.gelu = nn.GELU()
self.pw_conv2 = nn.Linear(expansion * dim, dim)
## single parameter for residual ratio
self.use_se = use_se
if use_se:
self.se = SEBlock(dim, reduction=reduction)
else:
# Indentity layer if SE is not used
self.se = nn.Identity()
self.apply(self._init_weights)
def _init_weights(self, m):
if isinstance(m, (nn.Conv2d, nn.Linear)):
# trunc_normal_(m.weight, std=.02)
if m.bias is not None:
nn.init.constant_(m.bias, 0)
def forward(self, x):
# Depthwise convolution with larger kernel
x = self.dw_conv(x)
# Permute to channels-last for LayerNorm
x = x.permute(0, 2, 3, 1)
x = self.norm(x)
x = self.pw_conv1(x)
x = self.gelu(x)
x = self.pw_conv2(x)
x = x.permute(0, 3, 1, 2)
x = self.se(x)
return x
class DecBlock(nn.Module):
def __init__(self, H, res, mixin, n_blocks):
super().__init__()
self.base = res
self.mixin = mixin
self.H = H
self.widths = get_width_settings(H.width, H.custom_width_str)
width = self.widths[res]
if mixin is not None and self.widths[mixin] != width:
self.proj = get_1x1(self.widths[mixin], width)
else:
self.proj = nn.Identity()
self.adaIN = AdaptiveInstanceNorm(width, H.latent_dim)
self.resnet = ConvNeXtBlock(width, H, kernel_size=7,
expansion=H.convnext_expansion,
use_se=H.use_se,
reduction=H.se_reduction,
dropout=H.dropout_p)
self.residual_ratio = nn.Parameter(torch.tensor(H.residual_ratio))
self.residual_type = H.residual_type # 'normal' or 'convex'
self.sigmoid = nn.Sigmoid()
def forward(self, x, w):
if self.mixin is not None:
x = F.interpolate(x, scale_factor=self.base / self.mixin, mode='bicubic')
x = self.proj(x)
residual = x
x = self.adaIN(x, w)
x = self.resnet(x)
if self.residual_type == 'normal':
return x * self.sigmoid(self.residual_ratio) + residual
elif self.residual_type == 'convex':
return x * self.sigmoid(self.residual_ratio) + residual * (1 - self.sigmoid(self.residual_ratio))
def stopgrad_keep_graph(x):
return x.detach() + 0.0 * x
class Decoder(nn.Module):
def __init__(self, H):
super().__init__()
self.H = H
self.use_rtm = getattr(H, "use_rtm", False)
if self.use_rtm:
self.mapping_network = RTMMappingNetwork(
code_dim=H.latent_dim,
num_tokens=getattr(H, "num_tokens", 1),
H_cycles=getattr(H, "H_cycles", 1),
L_cycles=getattr(H, "L_cycles", 1),
H_layers=getattr(H, "H_layers", 2),
L_layers=getattr(H, "L_layers", 2),
hidden_size=getattr(H, "rtm_hidden_size", 256),
expansion=getattr(H, "rtm_expansion", 4.0),
refinement_steps=getattr(H, "refinement_steps", 1),
with_grad=getattr(H, "rtm_with_grad", False),
cycle_noise_std=getattr(H, "rtm_cycle_noise_std", 0.0),
)
else:
self.mapping_network = MappingNetowrk(
code_dim=H.latent_dim, n_mlp=H.n_mpl,
mapping_lr_multiplier=getattr(H, "mapping_lr_multiplier", 1.0))
resos = set()
dec_blocks = []
self.widths = get_width_settings(H.width, H.custom_width_str)
blocks = parse_layer_string(H.dec_blocks)
for idx, (res, mixin) in enumerate(blocks):
dec_blocks.append(DecBlock(H, res, mixin, n_blocks=len(blocks)))
resos.add(res)
self.resolutions = sorted(resos)
self.dec_blocks = nn.ModuleList(dec_blocks)
first_res = self.resolutions[0]
last_res = self.resolutions[-1]
self.constant = nn.Parameter(torch.randn(1, self.widths[first_res], first_res, first_res))
resnets = {}
for res in self.resolutions:
key = str(res)
if res < 8:
resnets[key] = nn.Identity()
else:
resnets[key] = get_1x1(self.widths[res], H.image_channels)
self.resnets = nn.ModuleDict(resnets)
self.gains = nn.Parameter(torch.ones(1, H.image_channels, 1, 1))
self.biases = nn.Parameter(torch.zeros(1, H.image_channels, 1, 1))
def forward(self, latent_code, spatial_noise=None, input_is_w=False, train=False):
if not input_is_w:
w = self.mapping_network(latent_code)
if isinstance(w, tuple) or isinstance(w, list):
w = w[0]
else:
w = latent_code
targets = []
x = self.constant.repeat(latent_code.shape[0], 1, 1, 1)
for idx, block in enumerate(self.dec_blocks):
if(block.mixin is not None):
intermediate = self.resnets[str(block.mixin)](x)
targets.append(intermediate)
if(block.mixin >= 8 and self.H.use_stopgrad_for_intermediate):
x = x.detach()
x = block(x, w)
x = self.resnets[str(self.resolutions[-1])](x)
x = self.gains * x + self.biases
targets.append(x)
if(train):
return targets
else:
return targets[-1]
class IMLE(nn.Module):
def __init__(self, H):
super().__init__()
self.decoder = Decoder(H)
def forward(self, latents, spatial_noise=None, input_is_w=False, train=False):
return self.decoder.forward(latents, spatial_noise, input_is_w, train)