partial support for inversion
This commit is contained in:
@@ -3,14 +3,17 @@ import sys
|
||||
import numpy as np
|
||||
import pickle
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from tqdm import trange
|
||||
|
||||
from .slerp import slerp
|
||||
|
||||
from . import dnnlib
|
||||
from . import torch_utils
|
||||
# from . import legacy
|
||||
sys.modules["dnnlib"] = dnnlib
|
||||
sys.modules["torch_utils"] = torch_utils
|
||||
# sys.modules["legacy"] = legacy
|
||||
|
||||
import folder_paths
|
||||
from comfy.utils import PROGRESS_BAR_ENABLED, ProgressBar
|
||||
@@ -38,6 +41,8 @@ class LoadStyleGAN:
|
||||
def load_stylegan(self, stylegan_file):
|
||||
with open(folder_paths.get_full_path("stylegan", stylegan_file), 'rb') as f:
|
||||
G = pickle.load(f)['G_ema'].cuda()
|
||||
# device = torch.device('cuda:0')
|
||||
# G = legacy.load_network_pkl(f)['G_ema'].requires_grad_(False).to(device)
|
||||
return (G,)
|
||||
|
||||
class GenerateStyleGANLatent:
|
||||
@@ -76,7 +81,6 @@ class StyleGANSampler:
|
||||
"required": {
|
||||
"stylegan_model": ("STYLEGAN", ),
|
||||
"stylegan_latent": ("STYLEGAN_LATENT", ),
|
||||
# "class_label": ("INT", {"default": -1, "min": -1}),
|
||||
"noise_mode": (['const', 'random'],),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
},
|
||||
@@ -104,6 +108,71 @@ class StyleGANSampler:
|
||||
imgs = torch.cat(imgs, dim=0)
|
||||
return (imgs, )
|
||||
|
||||
class StyleGANInversion:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"stylegan_model": ("STYLEGAN", ),
|
||||
"image": ("IMAGE", ),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"num_steps": ("INT", {"default": 1000, "min": 1}),
|
||||
"w_avg_samples": ("INT", {"default": 10000, "min": 1, "max": 100000}),
|
||||
"initial_learning_rate": ("FLOAT", {"default": 0.1, "min": 0.00001, "max": 1.0, "step": 0.00001}),
|
||||
"initial_noise_factor": ("FLOAT", {"default": 0.05, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"lr_rampdown_length": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
"lr_rampup_length": ("FLOAT", {"default": 0.05, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
"noise_ramp_length": ("FLOAT", {"default": 0.75, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
"regularize_noise_weight": ("FLOAT", {"default": 1e5, "min": 0.0, "max": 1e7}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STYLEGAN_LATENT", "STYLEGAN_LATENT")
|
||||
RETURN_NAMES = ("training_latents", "final_latent")
|
||||
FUNCTION = "train_inversion"
|
||||
CATEGORY = "StyleGAN"
|
||||
|
||||
def train_inversion(
|
||||
self,
|
||||
stylegan_model,
|
||||
image,
|
||||
seed,
|
||||
num_steps,
|
||||
w_avg_samples,
|
||||
initial_learning_rate,
|
||||
initial_noise_factor,
|
||||
lr_rampdown_length,
|
||||
lr_rampup_length,
|
||||
noise_ramp_length,
|
||||
regularize_noise_weight
|
||||
):
|
||||
|
||||
device = torch.device('cuda:0')
|
||||
img_resolution = stylegan_model.img_resolution
|
||||
target_image = torch.permute(image[...,:3], (0, 3, 1, 2)) # BHWC -> BCHW, RGB only
|
||||
if target_image.shape != (stylegan_model.img_channels, img_resolution, img_resolution):
|
||||
target_image = F.interpolate(target_image, size=(img_resolution, img_resolution), mode='area')
|
||||
target_image = target_image[0] * 255
|
||||
|
||||
from .projector import project
|
||||
|
||||
projected_w_steps = project(
|
||||
stylegan_model,
|
||||
target_image,
|
||||
num_steps = num_steps,
|
||||
w_avg_samples = w_avg_samples,
|
||||
seed = seed,
|
||||
initial_learning_rate = initial_learning_rate,
|
||||
initial_noise_factor = initial_noise_factor,
|
||||
lr_rampdown_length = lr_rampdown_length,
|
||||
lr_rampup_length = lr_rampup_length,
|
||||
noise_ramp_length = noise_ramp_length,
|
||||
regularize_noise_weight = regularize_noise_weight,
|
||||
device = device,
|
||||
)
|
||||
|
||||
return (projected_w_steps, projected_w_steps[-1].unsqueeze(0))
|
||||
|
||||
class BlendStyleGANLatents:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -178,6 +247,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"BlendStyleGANLatents": BlendStyleGANLatents,
|
||||
"BatchAverageStyleGANLatents": BatchAverageStyleGANLatents,
|
||||
"StyleGANLatentFromBatch": StyleGANLatentFromBatch,
|
||||
"StyleGANInversion": StyleGANInversion,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
@@ -187,4 +257,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"BlendStyleGANLatents": "Blend StyleGAN Latents (lerp or slerp)",
|
||||
"BatchAverageStyleGANLatents": "Batch Average StyleGAN Latents",
|
||||
"StyleGANLatentFromBatch": "StyleGAN Latent From Batch",
|
||||
"StyleGANInversion": "StyleGAN Inversion",
|
||||
}
|
||||
+156
@@ -0,0 +1,156 @@
|
||||
# Modified from https://github.com/ouhenio/stylegan3-projector/blob/main/projector.py
|
||||
|
||||
import os
|
||||
import io
|
||||
import copy
|
||||
import tqdm
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from . import dnnlib
|
||||
|
||||
import folder_paths
|
||||
from comfy.utils import PROGRESS_BAR_ENABLED, ProgressBar
|
||||
|
||||
def load_vgg(device):
|
||||
# set the models directory
|
||||
if "VGG" not in folder_paths.folder_names_and_paths:
|
||||
current_paths = [os.path.join(folder_paths.models_dir, "VGG")]
|
||||
if not os.path.exists(current_paths[0]):
|
||||
os.mkdir(current_paths[0])
|
||||
else:
|
||||
current_paths, _ = folder_paths.folder_names_and_paths["VGG"]
|
||||
folder_paths.folder_names_and_paths["VGG"] = (current_paths, folder_paths.supported_pt_extensions)
|
||||
|
||||
vgg_file = None
|
||||
if "vgg16.pt" in folder_paths.get_filename_list("VGG"):
|
||||
vgg_file = folder_paths.get_full_path("VGG", "vgg16.pt")
|
||||
|
||||
if vgg_file is not None:
|
||||
with open(vgg_file, 'rb') as fv:
|
||||
vgg16 = torch.jit.load(fv).eval().to(device)
|
||||
else:
|
||||
url = 'https://nvlabs-fi-cdn.nvidia.com/stylegan2-ada-pytorch/pretrained/metrics/vgg16.pt'
|
||||
print("downloading VGG16")
|
||||
with dnnlib.util.open_url(url=url, cache=False) as fd:
|
||||
filename = os.path.join(current_paths[0], "vgg16.pt")
|
||||
with open(filename, "wb") as fv:
|
||||
print(f"saving VGG16 to {filename}")
|
||||
fv.write(fd.getvalue())
|
||||
vgg16 = torch.jit.load(fd).eval().to(device)
|
||||
|
||||
return vgg16
|
||||
|
||||
# @torch.enable_grad()
|
||||
@torch.inference_mode(mode=False)
|
||||
def project(
|
||||
G,
|
||||
target: torch.Tensor, # [C,H,W] and dynamic range [0,255], W & H must match G output resolution
|
||||
*,
|
||||
num_steps = 1000,
|
||||
w_avg_samples = 10000,
|
||||
seed = 0,
|
||||
initial_learning_rate = 0.1,
|
||||
initial_noise_factor = 0.05,
|
||||
lr_rampdown_length = 0.25,
|
||||
lr_rampup_length = 0.05,
|
||||
noise_ramp_length = 0.75,
|
||||
regularize_noise_weight = 1e5,
|
||||
device: torch.device
|
||||
):
|
||||
assert target.shape == (G.img_channels, G.img_resolution, G.img_resolution)
|
||||
|
||||
G = copy.deepcopy(G).eval().requires_grad_(False).to(device) # type: ignore
|
||||
|
||||
# Compute w stats.
|
||||
print(f'Computing W midpoint and stddev using {w_avg_samples} samples...')
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
z_samples = np.random.RandomState(seed).randn(w_avg_samples, G.z_dim)
|
||||
w_samples = G.mapping(torch.from_numpy(z_samples).to(device), None) # [N, L, C]
|
||||
w_samples = w_samples[:, :1, :].cpu().numpy().astype(np.float32) # [N, 1, C]
|
||||
w_avg = np.mean(w_samples, axis=0, keepdims=True) # [1, 1, C]
|
||||
w_std = (np.sum((w_samples - w_avg) ** 2) / w_avg_samples) ** 0.5
|
||||
|
||||
# Setup noise inputs.
|
||||
noise_bufs = { name: buf for (name, buf) in G.synthesis.named_buffers() if 'noise_const' in name }
|
||||
|
||||
vgg16 = load_vgg(device)
|
||||
|
||||
# Features for target image.
|
||||
target_images = target.unsqueeze(0).to(device).to(torch.float32)
|
||||
if target_images.shape[2] > 256:
|
||||
target_images = F.interpolate(target_images, size=(256, 256), mode='area')
|
||||
target_features = vgg16(target_images, resize_images=False, return_lpips=True)
|
||||
|
||||
w_opt = torch.tensor(w_avg, dtype=torch.float32, device=device, requires_grad=True) # pylint: disable=not-callable
|
||||
w_out = torch.zeros([num_steps] + list(w_opt.shape[1:]), dtype=torch.float32, device=device)
|
||||
optimizer = torch.optim.AdamW([w_opt] + list(noise_bufs.values()), betas=(0.9, 0.999), lr=initial_learning_rate)
|
||||
|
||||
# Init noise.
|
||||
for buf in noise_bufs.values():
|
||||
buf[:] = torch.randn_like(buf)
|
||||
buf.requires_grad = True
|
||||
|
||||
pbar = None
|
||||
if PROGRESS_BAR_ENABLED and num_steps > 1:
|
||||
pbar = ProgressBar(num_steps)
|
||||
tq_bar = tqdm.trange(num_steps, desc="Projecting")
|
||||
for step in tq_bar:
|
||||
# Learning rate schedule.
|
||||
t = step / num_steps
|
||||
w_noise_scale = w_std * initial_noise_factor * max(0.0, 1.0 - t / noise_ramp_length) ** 2
|
||||
lr_ramp = min(1.0, (1.0 - t) / lr_rampdown_length)
|
||||
lr_ramp = 0.5 - 0.5 * np.cos(lr_ramp * np.pi)
|
||||
lr_ramp = lr_ramp * min(1.0, t / lr_rampup_length)
|
||||
lr = initial_learning_rate * lr_ramp
|
||||
for param_group in optimizer.param_groups:
|
||||
param_group['lr'] = lr
|
||||
|
||||
# Synth images from opt_w.
|
||||
w_noise = torch.randn_like(w_opt) * w_noise_scale
|
||||
ws = (w_opt + w_noise).repeat([1, G.mapping.num_ws, 1])
|
||||
synth_images = G.synthesis(ws, noise_mode='const')
|
||||
|
||||
# Downsample image to 256x256 if it's larger than that. VGG was built for 224x224 images.
|
||||
synth_images = (synth_images + 1) * (255/2)
|
||||
if synth_images.shape[2] > 256:
|
||||
synth_images = F.interpolate(synth_images, size=(256, 256), mode='area')
|
||||
|
||||
# Features for synth images.
|
||||
synth_features = vgg16(synth_images, resize_images=False, return_lpips=True)
|
||||
dist = (target_features - synth_features).square().sum()
|
||||
|
||||
# Noise regularization.
|
||||
reg_loss = 0.0
|
||||
for v in noise_bufs.values():
|
||||
noise = v[None,None,:,:] # must be [1,1,H,W] for F.avg_pool2d()
|
||||
while True:
|
||||
reg_loss += (noise*torch.roll(noise, shifts=1, dims=3)).mean()**2
|
||||
reg_loss += (noise*torch.roll(noise, shifts=1, dims=2)).mean()**2
|
||||
if noise.shape[2] <= 8:
|
||||
break
|
||||
noise = F.avg_pool2d(noise, kernel_size=2)
|
||||
loss = dist + reg_loss * regularize_noise_weight
|
||||
# loss.requires_grad = True
|
||||
|
||||
# Step
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
tq_bar.set_postfix_str(f"dist: {dist:.2f} total: {float(loss):<5.2f}")
|
||||
|
||||
# Save projected W for each optimization step.
|
||||
w_out[step] = w_opt.detach()[0]
|
||||
|
||||
# Normalize noise.
|
||||
with torch.no_grad():
|
||||
for buf in noise_bufs.values():
|
||||
buf -= buf.mean()
|
||||
buf *= buf.square().mean().rsqrt()
|
||||
|
||||
if pbar is not None:
|
||||
pbar.update(1)
|
||||
|
||||
return w_out.repeat([1, G.mapping.num_ws, 1])
|
||||
Reference in New Issue
Block a user