USER
i,have,import argparse
import typing as tp
from collections import defaultdict
from functools import wraps
from pathlib import Path
import numpy as np
import torch
import torchvision.transforms.functional as F
from PIL import Image
from torchvision.io import read_image, ImageReadMode
from models.Alignment import Alignment
from models.Blending import Blending
from models.Embedding import Embedding
from models.Net import Net
from utils.image_utils import equal_replacer
from utils.seed import seed_setter
from utils.shape_predictor import align_face
from utils.time import bench_session
TImage = tp.TypeVar('TImage', torch.Tensor, Image.Image, np.ndarray)
TPath = tp.TypeVar('TPath', Path, str)
TReturn = tp.TypeVar('TReturn', torch.Tensor, tuple[torch.Tensor, ...])
class HairFast:
"""
HairFast implementation with hairstyle transfer interface
"""
def __init__(self, args):
self.args = args
self.net = Net(self.args)
self.embed = Embedding(args, net=self.net)
self.align = Alignment(args, self.embed.get_e4e_embed, net=self.net)
self.blend = Blending(args, net=self.net)
@seed_setter
@bench_session
def __swap_from_tensors(self, face: torch.Tensor, shape: torch.Tensor, color: torch.Tensor,
**kwargs) -> torch.Tensor:
images_to_name = defaultdict(list)
for image, name in zip((face, shape, color), ('face', 'shape', 'color')):
images_to_name[image].append(name)
# Embedding stage
name_to_embed = self.embed.embedding_images(images_to_name, **kwargs)
# Alignment stage
align_shape = self.align.align_images('face', 'shape', name_to_embed, **kwargs)
# Shape Module stage for blending
if shape is not color:
align_color = self.align.shape_module('face', 'color', name_to_embed, **kwargs)
else:
align_color = align_shape
# Blending and Post Process stage
final_image = self.blend.blend_images(align_shape, align_color, name_to_embed, **kwargs)
return final_image
def swap(self, face_img: TImage | TPath, shape_img: TImage | TPath, color_img: TImage | TPath,
benchmark=False, align=False, seed=None, exp_name=None, **kwargs) -> TReturn:
"""
Run HairFast on the input images to transfer hair shape and color to the desired images.
:param face_img: face image in Tensor, PIL Image, array or file path format
:param shape_img: shape image in Tensor, PIL Image, array or file path format
:param color_img: color image in Tensor, PIL Image, array or file path format
:param benchmark: starts counting the speed of the session
:param align: for arbitrary photos crops images to faces
:param seed: fixes seed for reproducibility, default 3407
:param exp_name: used as a folder name when 'save_all' model is enabled
:return: returns the final image as a Tensor
"""
images: list[torch.Tensor] = []
path_to_images: dict[TPath, torch.Tensor] = {}
for img in (face_img, shape_img, color_img):
if isinstance(img, (torch.Tensor, Image.Image, np.ndarray)):
if not isinstance(img, torch.Tensor):
img = F.to_tensor(img)
elif isinstance(img, (Path, str)):
path_img = img
if path_img not in path_to_images:
path_to_images[path_img] = read_image(str(path_img), mode=ImageReadMode.RGB)
img = path_to_images[path_img]
else:
raise TypeError(f'Unsupported image format {type(img)}')
images.append(img)
if align:
images = align_face(images)
images = equal_replacer(images)
final_image = self.__swap_from_tensors(*images, seed=seed, benchmark=benchmark, exp_name=exp_name, **kwargs)
if align:
return final_image, *images
return final_image
@wraps(swap)
def __call__(self, *args, **kwargs):
return self.swap(*args, **kwargs)
def get_parser():
parser = argparse.ArgumentParser(description='HairFast')
# I/O arguments
parser.add_argument('--save_all_dir', type=Path, default=Path('output'),
help='the directory to save the latent codes and inversion images')
# StyleGAN2 setting
parser.add_argument('--size', type=int, default=1024)
parser.add_argument('--ckpt', type=str, default="pretrained_models/StyleGAN/ffhq.pt")
parser.add_argument('--channel_multiplier', type=int, default=2)
parser.add_argument('--latent', type=int, default=512)
parser.add_argument('--n_mlp', type=int, default=8)
# Arguments
parser.add_argument('--device', type=str, default='cuda')
parser.add_argument('--batch_size', type=int, default=3, help='batch size for encoding images')
parser.add_argument('--save_all', action='store_true', help='save and print mode information')
# HairFast setting
parser.add_argument('--mixing', type=float, default=0.95, help='hair blending in alignment')
parser.add_argument('--smooth', type=int, default=5, help='dilation and erosion parameter')
parser.add_argument('--rotate_checkpoint', type=str, default='pretrained_models/Rotate/rotate_best.pth')
parser.add_argument('--blending_checkpoint', type=str, default='pretrained_models/Blending/checkpoint.pth')
parser.add_argument('--pp_checkpoint', type=str, default='pretrained_models/PostProcess/pp_model.pth')
return parser
if __name__ == '__main__':
model_args = get_parser()
args = model_args.parse_args()
hair_fast = HairFast(args)
import torch
import torch.nn.functional as F
import torchvision.transforms as T
from torch import nn
from models.CtrlHair.shape_branch.config import cfg as cfg_mask
from models.CtrlHair.shape_branch.solver import get_hair_face_code, get_new_shape, Solver as SolverMask
from models.Encoders import RotateModel
from models.Net import Net, get_segmentation
from models.sean_codes.models.pix2pix_model import Pix2PixModel, SEAN_OPT, encode_sean, decode_sean
from utils.image_utils import DilateErosion
from utils.save_utils import save_vis_mask, save_gen_image, save_latents
class Alignment(nn.Module):
"""
Module for transferring the desired hair shape
"""
def __init__(self, opts, latent_encoder=None, net=None):
super().__init__()
self.opts = opts
self.latent_encoder = latent_encoder
if not net:
self.net = Net(self.opts)
else:
self.net = net
self.sean_model = Pix2PixModel(SEAN_OPT)
self.sean_model.eval()
solver_mask = SolverMask(cfg_mask, device=self.opts.device, local_rank=-1, training=False)
self.mask_generator = solver_mask.gen
self.mask_generator.load_state_dict(torch.load('pretrained_models/ShapeAdaptor/mask_generator.pth'))
self.rotate_model = RotateModel()
self.rotate_model.load_state_dict(torch.load(self.opts.rotate_checkpoint)['model_state_dict'])
self.rotate_model.to(self.opts.device).eval()
self.dilate_erosion = DilateErosion(dilate_erosion=self.opts.smooth, device=self.opts.device)
self.to_bisenet = T.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))
@torch.inference_mode()
def shape_module(self, im_name1: str, im_name2: str, name_to_embed, only_target=True, **kwargs):
device = self.opts.device
# load images
img1_in = name_to_embed[im_name1]['image_256']
img2_in = name_to_embed[im_name2]['image_256']
# load latents
latent_W_1 = name_to_embed[im_name1]["W"]
latent_W_2 = name_to_embed[im_name2]["W"]
# load masks
inp_mask1 = name_to_embed[im_name1]['mask']
inp_mask2 = name_to_embed[im_name2]['mask']
# Rotate stage
if img1_in is not img2_in:
rotate_to = self.rotate_model(latent_W_2[:, :6], latent_W_1[:, :6])
rotate_to = torch.cat((rotate_to, latent_W_2[:, 6:]), dim=1)
I_rot, _ = self.net.generator([rotate_to], input_is_latent=True, return_latents=False)
I_rot_to_seg = ((I_rot + 1) / 2).clip(0, 1)
I_rot_to_seg = self.to_bisenet(I_rot_to_seg)
rot_mask = get_segmentation(I_rot_to_seg)
else:
I_rot = None
rot_mask = inp_mask2
# Shape Adaptor
if img1_in is not img2_in:
face_1, hair_1 = get_hair_face_code(self.mask_generator, inp_mask1[0, 0, ...])
face_2, hair_2 = get_hair_face_code(self.mask_generator, rot_mask[0, 0, ...])
target_mask = get_new_shape(self.mask_generator, face_1, hair_2)[None, None]
else:
target_mask = inp_mask1
# Hair mask
hair_mask_target = torch.where(target_mask == 13, torch.ones_like(target_mask, device=device),
torch.zeros_like(target_mask, device=device))
if self.opts.save_all:
exp_name = exp_name if (exp_name := kwargs.get('exp_name')) is not None else ""
output_dir = self.opts.save_all_dir / exp_name
if I_rot is not None:
save_gen_image(output_dir, 'Shape', f'{im_name2}_rotate_to_{im_name1}.png', I_rot)
save_vis_mask(output_dir, 'Shape', f'mask_{im_name1}.png', inp_mask1)
save_vis_mask(output_dir, 'Shape', f'mask_{im_name2}.png', inp_mask2)
save_vis_mask(output_dir, 'Shape', f'mask_{im_name2}_rotate_to_{im_name1}.png', rot_mask)
save_vis_mask(output_dir, 'Shape', f'mask_{im_name1}_{im_name2}_target.png', target_mask)
if only_target:
return {'HM_X': hair_mask_target}
else:
hair_mask1 = torch.where(inp_mask1 == 13, torch.ones_like(inp_mask1, device=device),
torch.zeros_like(inp_mask1, device=device))
hair_mask2 = torch.where(inp_mask2 == 13, torch.ones_like(inp_mask2, device=device),
torch.zeros_like(inp_mask2, device=device))
return inp_mask1, hair_mask1, inp_mask2, hair_mask2, target_mask, hair_mask_target
@torch.inference_mode()
def align_images(self, im_name1, im_name2, name_to_embed, **kwargs):
# load images
img1_in = name_to_embed[im_name1]['image_256']
img2_in = name_to_embed[im_name2]['image_256']
# load latents
latent_S_1, latent_F_1 = name_to_embed[im_name1]["S"], name_to_embed[im_name1]["F"]
latent_S_2, latent_F_2 = name_to_embed[im_name2]["S"], name_to_embed[im_name2]["F"]
# Shape Module
if img1_in is img2_in:
hair_mask_target = self.shape_module(im_name1, im_name2, name_to_embed, only_target=True, **kwargs)['HM_X']
return {'latent_F_align': latent_F_1, 'HM_X': hair_mask_target}
inp_mask1, hair_mask1, inp_mask2, hair_mask2, target_mask, hair_mask_target = (
self.shape_module(im_name1, im_name2, name_to_embed, only_target=False, **kwargs)
)
images = torch.cat([img1_in, img2_in], dim=0)
labels = torch.cat([inp_mask1, inp_mask2], dim=0)
# SEAN for inpaint
img1_code, img2_code = encode_sean(self.sean_model, images, labels)
gen1_sean = decode_sean(self.sean_model, img1_code.unsqueeze(0), target_mask)
gen2_sean = decode_sean(self.sean_model, img2_code.unsqueeze(0), target_mask)
# Encoding result in F from E4E
enc_imgs = self.latent_encoder([gen1_sean, gen2_sean])
intermediate_align, latent_inter = enc_imgs["F"][0].unsqueeze(0), enc_imgs["W"][0].unsqueeze(0)
latent_F_out_new, latent_out = enc_imgs["F"][1].unsqueeze(0), enc_imgs["W"][1].unsqueeze(0)
# Alignment of F space
masks = [
1 - (1 - hair_mask1) * (1 - hair_mask_target),
hair_mask_target,
hair_mask2 * hair_mask_target
]
masks = torch.cat(masks, dim=0)
# masks = T.functional.resize(masks, (1024, 1024), interpolation=T.InterpolationMode.NEAREST)
dilate, erosion = self.dilate_erosion.mask(masks)
free_mask = [
dilate[0],
erosion[1],
erosion[2]
]
free_mask = torch.stack(free_mask, dim=0)
free_mask_down_32 = F.interpolate(free_mask.float(), size=(32, 32), mode='bicubic')
interpolation_low = 1 - free_mask_down_32
latent_F_align = intermediate_align + interpolation_low[0] * (latent_F_1 - intermediate_align)
latent_F_align = latent_F_out_new + interpolation_low[1] * (latent_F_align - latent_F_out_new)
latent_F_align = latent_F_2 + interpolation_low[2] * (latent_F_align - latent_F_2)
if self.opts.save_all:
exp_name = exp_name if (exp_name := kwargs.get('exp_name')) is not None else ""
output_dir = self.opts.save_all_dir / exp_name
save_gen_image(output_dir, 'Align', f'{im_name1}_{im_name2}_SEAN.png', gen1_sean)
save_gen_image(output_dir, 'Align', f'{im_name2}_{im_name1}_SEAN.png', gen2_sean)
img1_e4e = self.net.generator([latent_inter], input_is_latent=True, return_latents=False, start_layer=4,
end_layer=8, layer_in=intermediate_align)[0]
img2_e4e = self.net.generator([latent_out], input_is_latent=True, return_latents=False, start_layer=4,
end_layer=8, layer_in=latent_F_out_new)[0]
save_gen_image(output_dir, 'Align', f'{im_name1}_{im_name2}_e4e.png', img1_e4e)
save_gen_image(output_dir, 'Align', f'{im_name2}_{im_name1}_e4e.png', img2_e4e)
gen_im, _ = self.net.generator([latent_S_1], input_is_latent=True, return_latents=False, start_layer=4,
end_layer=8, layer_in=latent_F_align)
save_gen_image(output_dir, 'Align', f'{im_name1}_{im_name2}_output.png', gen_im)
save_latents(output_dir, 'Align', f'{im_name1}_{im_name2}_F.npz', latent_F_align=latent_F_align)
return {'latent_F_align': latent_F_align, 'HM_X': hair_mask_target}
import torch
from torch import nn
from models.Encoders import ClipBlendingModel, PostProcessModel
from models.Net import Net
from utils.bicubic import BicubicDownSample
from utils.image_utils import DilateErosion
from utils.save_utils import save_gen_image, save_latents
class Blending(nn.Module):
"""
Module for transferring the desired hair color and post processing
"""
def __init__(self, opts, net=None):
super().__init__()
self.opts = opts
if net is None:
self.net = Net(self.opts)
else:
self.net = net
blending_checkpoint = torch.load(self.opts.blending_checkpoint)
self.blending_encoder = ClipBlendingModel(blending_checkpoint.get('clip', "ViT-B/32"))
self.blending_encoder.load_state_dict(blending_checkpoint['model_state_dict'], strict=False)
self.blending_encoder.to(self.opts.device).eval()
self.post_process = PostProcessModel().to(self.opts.device).eval()
self.post_process.load_state_dict(torch.load(self.opts.pp_checkpoint)['model_state_dict'])
self.dilate_erosion = DilateErosion(dilate_erosion=self.opts.smooth, device=self.opts.device)
self.downsample_256 = BicubicDownSample(factor=4)
@torch.inference_mode()
def blend_images(self, align_shape, align_color, name_to_embed, **kwargs):
I_1 = name_to_embed['face']['image_norm_256']
I_2 = name_to_embed['shape']['image_norm_256']
I_3 = name_to_embed['color']['image_norm_256']
mask_de = self.dilate_erosion.hair_from_mask(
torch.cat([name_to_embed[x]['mask'] for x in ['face', 'color']], dim=0)
)
HM_1D, _ = mask_de[0][0].unsqueeze(0), mask_de[1][0].unsqueeze(0)
HM_3D, HM_3E = mask_de[0][1].unsqueeze(0), mask_de[1][1].unsqueeze(0)
latent_S_1, latent_F_align = name_to_embed['face']['S'], align_shape['latent_F_align']
HM_X = align_color['HM_X']
latent_S_3 = name_to_embed['color']["S"]
HM_XD, _ = self.dilate_erosion.mask(HM_X)
target_mask = (1 - HM_1D) * (1 - HM_3D) * (1 - HM_XD)
# Blending
if I_1 is not I_3 or I_1 is not I_2:
S_blend_6_18 = self.blending_encoder(latent_S_1[:, 6:], latent_S_3[:, 6:], I_1 * target_mask, I_3 * HM_3E)
S_blend = torch.cat((latent_S_1[:, :6], S_blend_6_18), dim=1)
else:
S_blend = latent_S_1
I_blend, _ = self.net.generator([S_blend], input_is_latent=True, return_latents=False, start_layer=4,
end_layer=8, layer_in=latent_F_align)
I_blend_256 = self.downsample_256(I_blend)
# Post Process
S_final, F_final = self.post_process(I_1, I_blend_256)
I_final, _ = self.net.generator([S_final], input_is_latent=True, return_latents=False,
start_layer=5, end_layer=8, layer_in=F_final)
if self.opts.save_all:
exp_name = exp_name if (exp_name := kwargs.get('exp_name')) is not None else ""
output_dir = self.opts.save_all_dir / exp_name
save_gen_image(output_dir, 'Blending', 'blending.png', I_blend)
save_latents(output_dir, 'Blending', 'blending.npz', S_blend=S_blend)
save_gen_image(output_dir, 'Final', 'final.png', I_final)
save_latents(output_dir, 'Final', 'final.npz', S_final=S_final, F_final=F_final)
final_image = ((I_final[0] + 1) / 2).clip(0, 1)
return final_image
from collections import defaultdict
import torch
import torch.nn.functional as F
import torchvision.transforms as T
from torch import nn
from torch.utils.data import DataLoader
from datasets.image_dataset import ImagesDataset, image_collate
from models.FeatureStyleEncoder import FSencoder
from models.Net import Net, get_segmentation
from models.encoder4editing.utils.model_utils import setup_model, get_latents
from utils.bicubic import BicubicDownSample
from utils.save_utils import save_gen_image, save_latents
class Embedding(nn.Module):
"""
Module for image embedding
"""
def __init__(self, opts, net=None):
super().__init__()
self.opts = opts
if net is None:
self.net = Net(self.opts)
else:
self.net = net
self.encoder = FSencoder.get_trainer(self.opts.device)
self.e4e, _ = setup_model('pretrained_models/encoder4editing/e4e_ffhq_encode.pt', self.opts.device)
self.normalize = T.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
self.to_bisenet = T.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))
self.downsample_512 = BicubicDownSample(factor=2)
self.downsample_256 = BicubicDownSample(factor=4)
def setup_dataloader(self, images: dict[torch.Tensor, list[str]] | list[torch.Tensor], batch_size=None):
self.dataset = ImagesDataset(images)
self.dataloader = DataLoader(self.dataset, collate_fn=image_collate, shuffle=False,
batch_size=batch_size or self.opts.batch_size)
@torch.inference_mode()
def get_e4e_embed(self, images: list[torch.Tensor]) -> dict[str, torch.Tensor]:
device = self.opts.device
self.setup_dataloader(images, batch_size=len(images))
for image, _ in self.dataloader:
image = image.to(device)
latent_W = get_latents(self.e4e, image)
latent_F, _ = self.net.generator([latent_W], input_is_latent=True, return_latents=False,
start_layer=0, end_layer=3)
return {"F": latent_F, "W": latent_W}
@torch.inference_mode()
def embedding_images(self, images_to_name: dict[torch.Tensor, list[str]], **kwargs) -> dict[
str, dict[str, torch.Tensor]]:
device = self.opts.device
self.setup_dataloader(images_to_name)
name_to_embed = defaultdict(dict)
for image, names in self.dataloader:
image = image.to(device)
im_512 = self.downsample_512(image)
im_256 = self.downsample_256(image)
im_256_norm = self.normalize(im_256)
# E4E
latent_W = get_latents(self.e4e, im_256_norm)
# FS encoder
output = self.encoder.test(img=self.normalize(image), return_latent=True)
latent = output.pop() # [bs, 512, 16, 16]
latent_S = output.pop() # [bs, 18, 512]
latent_F, _ = self.net.generator([latent_S], input_is_latent=True, return_latents=False,
start_layer=3, end_layer=3, layer_in=latent) # [bs, 512, 32, 32]
# BiSeNet
masks = torch.cat([get_segmentation(image.unsqueeze(0)) for image in self.to_bisenet(im_512)])
# Mixing if we change the color or shape
if len(images_to_name) > 1:
hair_mask = torch.where(masks == 13, torch.ones_like(masks, device=device),
torch.zeros_like(masks, device=device))
hair_mask = F.interpolate(hair_mask.float(), size=(32, 32), mode='bicubic')
latent_F_from_W = self.net.generator([latent_W], input_is_latent=True, return_latents=False,
start_layer=0, end_layer=3)[0]
latent_F = latent_F + self.opts.mixing * hair_mask * (latent_F_from_W - latent_F)
for k, names in enumerate(names):
for name in names:
name_to_embed[name]['W'] = latent_W[k].unsqueeze(0)
name_to_embed[name]['F'] = latent_F[k].unsqueeze(0)
name_to_embed[name]['S'] = latent_S[k].unsqueeze(0)
name_to_embed[name]['mask'] = masks[k].unsqueeze(0)
name_to_embed[name]['image_256'] = im_256[k].unsqueeze(0)
name_to_embed[name]['image_norm_256'] = im_256_norm[k].unsqueeze(0)
if self.opts.save_all:
gen_W_im, _ = self.net.generator([latent_W], input_is_latent=True, return_latents=False)
gen_FS_im, _ = self.net.generator([latent_S], input_is_latent=True, return_latents=False,
start_layer=4, end_layer=8, layer_in=latent_F)
exp_name = exp_name if (exp_name := kwargs.get('exp_name')) is not None else ""
output_dir = self.opts.save_all_dir / exp_name
for name, im_W, lat_W in zip(names, gen_W_im, latent_W):
save_gen_image(output_dir, 'W+', f'{name}.png', im_W)
save_latents(output_dir, 'W+', f'{name}.npz', latent_W=lat_W)
for name, im_F, lat_S, lat_F in zip(names, gen_FS_im, latent_S, latent_F):
save_gen_image(output_dir, 'FS', f'{name}.png', im_F)
save_latents(output_dir, 'FS', f'{name}.npz', latent_S=lat_S, latent_F=lat_F)
return name_to_embed
import argparse
import clip
import torch
import torch.nn as nn
from torch.nn import Linear, LayerNorm, LeakyReLU, Sequential
from torchvision import transforms as T
from models.Net import FeatureEncoderMult, IBasicBlock, conv1x1
from models.stylegan2.model import PixelNorm
class ModulationModule(nn.Module):
def __init__(self, layernum, last=False, inp=512, middle=512):
super().__init__()
self.layernum = layernum
self.last = last
self.fc = Linear(512, 512)
self.norm = LayerNorm([self.layernum, 512], elementwise_affine=False)
self.gamma_function = Sequential(Linear(inp, middle), LayerNorm([middle]), LeakyReLU(), Linear(middle, 512))
self.beta_function = Sequential(Linear(inp, middle), LayerNorm([middle]), LeakyReLU(), Linear(middle, 512))
self.leakyrelu = LeakyReLU()
def forward(self, x, embedding):
x = self.fc(x)
x = self.norm(x)
gamma = self.gamma_function(embedding)
beta = self.beta_function(embedding)
out = x * (1 + gamma) + beta
if not self.last:
out = self.leakyrelu(out)
return out
class FeatureiResnet(nn.Module):
def __init__(self, blocks, inplanes=1024):
super().__init__()
self.res_blocks = {}
for n, block in enumerate(blocks, start=1):
planes, num_blocks = block
for k in range(1, num_blocks + 1):
downsample = None
if inplanes != planes:
downsample = nn.Sequential(conv1x1(inplanes, planes, 1), nn.BatchNorm2d(planes, eps=1e-05, ), )
self.res_blocks[f'res_block_{n}_{k}'] = IBasicBlock(inplanes, planes, 1, downsample, 1, 64, 1)
inplanes = planes
self.res_blocks = nn.ModuleDict(self.res_blocks)
def forward(self, x):
for module in self.res_blocks.values():
x = module(x)
return x
class RotateModel(nn.Module):
def __init__(self):
super().__init__()
self.pixelnorm = PixelNorm()
self.modulation_module_list = nn.ModuleList([ModulationModule(6, i == 4) for i in range(5)])
def forward(self, latent_from, latent_to):
dt_latent = self.pixelnorm(latent_from)
for modulation_module in self.modulation_module_list:
dt_latent = modulation_module(dt_latent, latent_to)
output = latent_from + 0.1 * dt_latent
return output
class ClipBlendingModel(nn.Module):
def __init__(self, clip_model="ViT-B/32"):
super().__init__()
self.pixelnorm = PixelNorm()
self.clip_model, _ = clip.load(clip_model, device="cuda")
self.transform = T.Compose(
[T.Normalize((0.48145466, 0.4578275, 0.40821073), (0.26862954, 0.26130258, 0.27577711))])
self.face_pool = torch.nn.AdaptiveAvgPool2d((224, 224))
self.modulation_module_list = nn.ModuleList(
[ModulationModule(12, i == 4, inp=512 * 3, middle=1024) for i in range(5)]
)
for param in self.clip_model.parameters():
param.requires_grad = False
def get_image_embed(self, image_tensor):
resized_tensor = self.face_pool(image_tensor)
renormed_tensor = self.transform(resized_tensor * 0.5 + 0.5)
return self.clip_model.encode_image(renormed_tensor)
def forward(self, latent_face, latent_color, target_face, hair_color):
embed_face = self.get_image_embed(target_face).unsqueeze(1).expand(-1, 12, -1)
embed_color = self.get_image_embed(hair_color).unsqueeze(1).expand(-1, 12, -1)
latent_in = torch.cat((latent_color, embed_face, embed_color), dim=-1)
dt_latent = self.pixelnorm(latent_face)
for modulation_module in self.modulation_module_list:
dt_latent = modulation_module(dt_latent, latent_in)
output = latent_face + 0.1 * dt_latent
return output
class PostProcessModel(nn.Module):
def __init__(self):
super().__init__()
self.encoder_face = FeatureEncoderMult(fs_layers=[9], opts=argparse.Namespace(
**{'arcface_model_path': "pretrained_models/ArcFace/backbone_ir50.pth"}))
self.latent_avg = torch.load('pretrained_models/PostProcess/latent_avg.pt', map_location=torch.device('cuda'))
self.to_feature = FeatureiResnet([[1024, 2], [768, 2], [512, 2]])
self.to_latent_1 = nn.ModuleList([ModulationModule(18, i == 4) for i in range(5)])
self.to_latent_2 = nn.ModuleList([ModulationModule(18, i == 4) for i in range(5)])
self.pixelnorm = PixelNorm()
def forward(self, source, target):
s_face, [f_face] = self.encoder_face(source)
s_hair, [f_hair] = self.encoder_face(target)
dt_latent_face = self.pixelnorm(s_face)
dt_latent_hair = self.pixelnorm(s_hair)
for mod_module in self.to_latent_1:
dt_latent_face = mod_module(dt_latent_face, s_hair)
for mod_module in self.to_latent_2:
dt_latent_hair = mod_module(dt_latent_hair, s_face)
finall_s = self.latent_avg + 0.1 * (dt_latent_face + dt_latent_hair)
cat_f = torch.cat((f_face, f_hair), dim=1)
finall_f = self.to_feature(cat_f)
return finall_s, finall_f
class ClipModel(nn.Module):
def __init__(self):
super().__init__()
self.clip_model, _ = clip.load("ViT-B/32", device="cuda")
self.transform = T.Compose(
[T.Normalize((0.48145466, 0.4578275, 0.40821073), (0.26862954, 0.26130258, 0.27577711))]
)
self.face_pool = torch.nn.AdaptiveAvgPool2d((224, 224))
for param in self.clip_model.parameters():
param.requires_grad = False
def forward(self, image_tensor):
if not image_tensor.is_cuda:
image_tensor = image_tensor.to("cuda")
if image_tensor.dtype == torch.uint8:
image_tensor = image_tensor / 255
resized_tensor = self.face_pool(image_tensor)
renormed_tensor = self.transform(resized_tensor)
return self.clip_model.encode_image(renormed_tensor)
do,to,cpu,on,every,part,needed,hereimport argparse
from pathlib import Path
from hair_swap import HairFast, get_parser
model_args = get_parser()
hair_fast = HairFast(model_args.parse_args([]))