JoJoGAN 实践

JoJoGAN: One Shot Face Stylization. 只用一张人脸图片,就能学习其风格,然后迁移到其他图片。训练时长只用 1~2 min 即可。

效果:

主流程:

本文分享了个人在本地环境(非 colab)实践 JoJoGAN 的整个过程。你也可以依照本文上手训练自己喜欢的风格。

准备环境

安装:

1conda create -n torch python=3.9 -y 2conda activate torch 3 4conda install pytorch torchvision torchaudio cudatoolkit=11.3 -c pytorch -y

检查:

1$ python - <<EOF 2import torch, torchvision 3print(torch.__version__, torch.cuda.is_available()) 4EOF 51.10.1 True

准备代码

1git clone https://github.com/mchong6/JoJoGAN.git 2cd JoJoGAN 3 4pip install tqdm gdown matplotlib scipy opencv-python dlib lpips wandb 5 6# Ninja is required to load C++ extensions 7wget https://github.com/ninja-build/ninja/releases/download/v1.10.2/ninja-linux.zip 8sudo unzip ninja-linux.zip -d /usr/local/bin/ 9sudo update-alternatives --install /usr/bin/ninja ninja /usr/local/bin/ninja 1 --force

然后,将本文提供的几个 *.py 放进 JoJoGAN 目录,从这里获取: https://github.com/ikuokuo/start-deep-learning/tree/master/practice/JoJoGAN

  • download_models.py: 获取模型
  • generate_faces.py: 生成人脸
  • stylize.py: 风格化
  • train.py: 训练

之后,于训练流程一节,会结合代码,讲述下 JoJoGAN 的工作流程。其他些 *.py 只提下用法,实现就不多说了。

获取模型

python download_models.py 获取模型,如下:

1models/ 2├── arcane_caitlyn_preserve_color.pt 3├── arcane_caitlyn.pt 4├── arcane_jinx_preserve_color.pt 5├── arcane_jinx.pt 6├── arcane_multi_preserve_color.pt 7├── arcane_multi.pt 8├── art.pt 9├── disney_preserve_color.pt 10├── disney.pt 11├── dlibshape_predictor_68_face_landmarks.dat 12├── e4e_ffhq_encode.pt 13├── jojo_preserve_color.pt 14├── jojo.pt 15├── jojo_yasuho_preserve_color.pt 16├── jojo_yasuho.pt 17├── restyle_psp_ffhq_encode.pt 18├── stylegan2-ffhq-config-f.pt 19├── supergirl_preserve_color.pt 20└── supergirl.pt

生成人脸

用 StyleGAN2 预训练模型随机生成人脸,用于测试:

1python generate_faces.py -n 5 -s 2000 -o input

使用预训练风格

JoJoGAN 给了 8 个预训练模型,可以一并体验,与文首的效果图一样:

1# 预览 JoJoGAN 所有预训练模型 风格化某图片(test_input/iu.jpeg)的效果 2python stylize.py -i test_input/iu.jpeg -s all --save-all --show-all 3 4# 使用 JoJoGAN 所有预训练模型 风格化所有生成的测试人脸(input/*) 5find ./input -type f -print0 | xargs -0 -i python stylize.py -i {} -s all --save-all

训练自己的风格

首先,准备一张风格图:

之后,开始训练:

1python train.py -n yinshi -i style_images/yinshi.jpeg --alpha 1.0 --num_iter 500 --latent_dim 512 --use_wandb --log_interval 50

--use_wandb 时,可查看训练日志:

最后,测试效果:

1python stylize.py -i input/girl.jpeg --save-all --show-all --test_style yinshi --test_ckpt output/yinshi.pt --test_ref output/yinshi/style_images_aligned/yinshi.png

训练工作流程

准备风格图片,转为训练数据

将风格图片里的人脸裁减对齐:

1# dlib 预测人脸特征点,再裁减对齐 2from util import align_face 3style_aligned = align_face(img_path)

将风格图片 GAN Inversion 逆映射回预训练模型的隐向量空间(Latent Space):

1name, _ = os.path.splitext(os.path.basename(img_path)) 2style_code_path = os.path.join(latent_dir, f'{name}.pt') 3 4# e4e FFHQ encoder (pSp) > GAN inversion,得到 latent 5from e4e_projection import projection 6latent = projection(style_aligned, style_code_path, device)

载入 StyleGAN2 模型,训练微调

载入预训练模型:

1latent_dim = 512 2 3# 加载预训练模型 4original_generator = Generator(1024, latent_dim, 8, 2).to(device) 5ckpt = torch.load("models/stylegan2-ffhq-config-f.pt", map_location=lambda storage, loc: storage) 6original_generator.load_state_dict(ckpt["g_ema"], strict=False) 7 8# 准备微调的模型 9generator = deepcopy(original_generator)

训练可调参数:

1# 控制风格强度 [0, 1] 2alpha = 1.0 3alpha = 1-alpha 4 5# 是否保留原图像色彩 6preserve_color = True 7 8# 训练迭代次数(最好 500,Adam 学习率是基于 500 次迭代调优的) 9num_iter = 500 10 11# 风格图片 targets 及 latents 12targets = .. 13latents = ..

进行训练,拟合隐空间。最后保存:

1# 准备 LPIPS 计算 loss 2lpips_fn = lpips.LPIPS(net='vgg').to(device) 3 4# 准备优化器 5g_optim = torch.optim.Adam(generator.parameters(), lr=2e-3, betas=(0, 0.99)) 6 7# 哪些层用于交换,用于生成风格化图片 8if preserve_color: 9 id_swap = [7,9,11,15,16,17] 10else: 11 id_swap = list(range(7, generator.n_latent)) 12 13# 训练迭代 14for idx in tqdm(range(num_iter)): 15 # 交换层混合风格,并加噪声 16 mean_w = generator.get_latent(torch.randn([latents.size(0), latent_dim]) 17 .to(device)).unsqueeze(1).repeat(1, generator.n_latent, 1) 18 in_latent = latents.clone() 19 in_latent[:, id_swap] = alpha*latents[:, id_swap] + (1-alpha)*mean_w[:, id_swap] 20 21 # 以 latent 风格化图片,与目标风格对比 22 img = generator(in_latent, input_is_latent=True) 23 loss = lpips_fn(F.interpolate(img, size=(256,256), mode='area'), 24 F.interpolate(targets, size=(256,256), mode='area')).mean() 25 26 # 优化 27 g_optim.zero_grad() 28 loss.backward() 29 g_optim.step() 30 31# 保存权重,完成 32torch.save({"g": generator.state_dict()}, save_path)

结语

JoJoGAN 实践下来效果不错。使用本文给到的代码,更容易上手训练自己喜欢的风格,值得试试。

GoCoding 个人实践的经验分享,可关注公众号!

点赞
收藏

评论区

加载中...

相关推荐

MySQL:[Err] 1292 - Incorrect datetime value: ‘0000-00-00 00:00:00‘ for column ‘CREATE_TIME‘ at row 1

文章目录问题用navicat导入数据时,报错:原因这是因为当前的MySQL不支持datetime为0的情况。解决修改sql\mode:sql\mode:SQLMode定义了MySQL应支持的SQL语法、数据校验等,这样可以更容易地在不同的环境中使用MySQL。全局s

Oracle 分组与拼接字符串同时使用

SELECTT.,ROWNUMIDFROM(SELECTT.EMPLID,T.NAME,T.BU,T.REALDEPART,T.FORMATDATE,SUM(T.S0)S0,MAX(UPDATETIME)CREATETIME,LISTAGG(TOCHAR(

MySQL部分从库上面因为大量的临时表tmp_table造成慢查询

背景描述Time:20190124T00:08:14.70572408:00User@Host:@Id:Schema:sentrymetaLast_errno:0Killed:0Query_time:0.315758Lock_

2020年前端实用代码段,为你的工作保驾护航

有空的时候,自己总结了几个代码段,在开发中也经常使用,谢谢。1、使用解构获取json数据let jsonData  id: 1,status: "OK",data: 'a', 'b';let  id, status, data: number   jsonData;console.log(id, status, number )

roop 视频换脸

roop:oneclickfaceswap.只用一张人脸图片,就能完成视频换脸。本文是本地部署的实践记录。

Python3:sqlalchemy对mysql数据库操作,非sql语句

Python3:sqlalchemy对mysql数据库操作,非sql语句python3authorlizmdatetime2018020110:00:00coding:utf8'''