JerMa88's picture
Upload folder using huggingface_hub
3ce19a2 verified
Raw
History Blame Contribute Delete
5.01 kB
from typing import Optional
import torch
import torch.nn as nn
from models.stylegan2 import FullyConnectedLayer, normalize_2nd_moment
from models.rtm_core import RTMMappingNetwork
class RTMMappingNetworkStyleGAN2(nn.Module):
def __init__(
self,
z_dim: int,
c_dim: int,
w_dim: int,
num_ws: Optional[int],
num_layers: int = 8,
embed_features: Optional[int] = None,
layer_features: Optional[int] = None,
activation: str = "lrelu",
lr_multiplier: float = 0.01,
w_avg_beta: float = 0.998,
rtm_num_tokens: int = 4,
rtm_H_cycles: int = 4,
rtm_L_cycles: int = 1,
rtm_H_layers: int = 2,
rtm_L_layers: int = 2,
rtm_hidden_size: int = 128,
rtm_expansion: float = 4.0,
rtm_refinement_steps: int = 4,
rtm_with_grad: bool = False,
rtm_cycle_noise_std: float = 0.0,
use_rtm_equalized: bool = True,
rtm_lr_multiplier: float = 0.01,
):
super().__init__()
del layer_features
self.z_dim = z_dim
self.c_dim = c_dim
self.w_dim = w_dim
self.num_ws = num_ws
self.num_layers = num_layers
self.w_avg_beta = w_avg_beta
if embed_features is None:
embed_features = w_dim
if c_dim == 0:
embed_features = 0
if c_dim > 0:
self.embed = FullyConnectedLayer(c_dim, embed_features)
else:
self.embed = None
self.fuse = FullyConnectedLayer(
z_dim + embed_features, w_dim,
activation=activation, lr_multiplier=lr_multiplier,
)
self.rtm = RTMMappingNetwork(
code_dim=w_dim,
num_tokens=rtm_num_tokens,
H_cycles=rtm_H_cycles,
L_cycles=rtm_L_cycles,
H_layers=rtm_H_layers,
L_layers=rtm_L_layers,
hidden_size=rtm_hidden_size,
expansion=rtm_expansion,
refinement_steps=rtm_refinement_steps,
with_grad=rtm_with_grad,
cycle_noise_std=rtm_cycle_noise_std,
use_equalized=use_rtm_equalized,
lr_multiplier=rtm_lr_multiplier,
)
if num_ws is not None and w_avg_beta is not None:
self.register_buffer("w_avg", torch.zeros([w_dim]))
def forward(
self,
z: torch.Tensor,
c: Optional[torch.Tensor],
truncation_psi: float = 1.0,
truncation_cutoff: Optional[int] = None,
update_emas: bool = False,
) -> torch.Tensor:
x = None
if self.z_dim > 0:
x = normalize_2nd_moment(z.to(torch.float32))
if self.c_dim > 0:
assert c is not None and self.embed is not None
y = normalize_2nd_moment(self.embed(c.to(torch.float32)))
x = torch.cat([x, y], dim=1) if x is not None else y
x = self.fuse(x)
x = self.rtm(x)
if isinstance(x, (list, tuple)):
x = x[-1]
if update_emas and self.w_avg_beta is not None and hasattr(self, "w_avg"):
self.w_avg.copy_(
x.detach().mean(dim=0).lerp(self.w_avg, self.w_avg_beta)
)
if self.num_ws is not None:
x = x.unsqueeze(1).repeat([1, self.num_ws, 1])
if truncation_psi != 1:
assert self.w_avg_beta is not None
if self.num_ws is None or truncation_cutoff is None:
x = self.w_avg.lerp(x, truncation_psi)
else:
x[:, :truncation_cutoff] = self.w_avg.lerp(
x[:, :truncation_cutoff], truncation_psi
)
return x
def forward_ws_trajectory(
self,
z: torch.Tensor,
c: Optional[torch.Tensor],
truncation_psi: float = 1.0,
truncation_cutoff: Optional[int] = None,
):
x = None
if self.z_dim > 0:
x = normalize_2nd_moment(z.to(torch.float32))
if self.c_dim > 0:
assert c is not None and self.embed is not None
y = normalize_2nd_moment(self.embed(c.to(torch.float32)))
x = torch.cat([x, y], dim=1) if x is not None else y
x = self.fuse(x)
flat_list = self.rtm.forward_w_trajectory(x)
ws_out = []
for xf in flat_list:
if isinstance(xf, (list, tuple)):
xf = xf[-1]
xcur = xf
if self.num_ws is not None:
xcur = xcur.unsqueeze(1).repeat([1, self.num_ws, 1])
if truncation_psi != 1:
assert self.w_avg_beta is not None
if self.num_ws is None or truncation_cutoff is None:
xcur = self.w_avg.lerp(xcur, truncation_psi)
else:
xcur[:, :truncation_cutoff] = self.w_avg.lerp(
xcur[:, :truncation_cutoff], truncation_psi
)
ws_out.append(xcur)
return ws_out