Compare commits
47
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5e3c92d55c | ||
|
|
cc0501a2db | ||
|
|
3dc822fd2f | ||
|
|
697b9a78e6 | ||
|
|
2a7bd6116f | ||
|
|
a7d09f5d49 | ||
|
|
8e85d5b96d | ||
|
|
30989a9d37 | ||
|
|
eecf645603 | ||
|
|
1b080706df | ||
|
|
336f3f7c23 | ||
|
|
f27e1cca13 | ||
|
|
86e91a6e9d | ||
|
|
92529f7ca8 | ||
|
|
c0959056ae | ||
|
|
2e40fe3820 | ||
|
|
0a5e187637 | ||
|
|
4e19dbd6d1 | ||
|
|
9c190804a7 | ||
|
|
d9ca40e1d6 | ||
|
|
b68cf8788c | ||
|
|
e702b26895 | ||
|
|
c21705edb5 | ||
|
|
a284bb52b2 | ||
|
|
9884aac18a | ||
|
|
f8aada81db | ||
|
|
6735771664 | ||
|
|
857ddbc6d7 | ||
|
|
a6edcda97d | ||
|
|
ca01d706d0 | ||
|
|
0dc9a8a695 | ||
|
|
ba6b3f5f68 | ||
|
|
ee7d5b4241 | ||
|
|
03df9f35cd | ||
|
|
ec6b5c8c85 | ||
|
|
8509d9a551 | ||
|
|
e724da1161 | ||
|
|
68d0ddf72a | ||
|
|
6f9dba7777 | ||
|
|
811ca557fb | ||
|
|
6d790bdcc3 | ||
|
|
ef8b4263b4 | ||
|
|
eb5fddf4de | ||
|
|
9c7db3c59a | ||
|
|
bf3410cd0d | ||
|
|
24c65627db | ||
|
|
72bb6910e9 |
@@ -1,500 +0,0 @@
|
||||
{
|
||||
"last_node_id": 31,
|
||||
"last_link_id": 68,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 4,
|
||||
"type": "LoadImage",
|
||||
"pos": [
|
||||
138,
|
||||
323
|
||||
],
|
||||
"size": {
|
||||
"0": 272.85791015625,
|
||||
"1": 331.60894775390625
|
||||
},
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
59
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadImage"
|
||||
},
|
||||
"widgets_values": [
|
||||
"oldman.jpg",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 19,
|
||||
"type": "ImageResizeKJ",
|
||||
"pos": [
|
||||
507,
|
||||
675
|
||||
],
|
||||
"size": {
|
||||
"0": 315,
|
||||
"1": 242
|
||||
},
|
||||
"flags": {},
|
||||
"order": 4,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image",
|
||||
"type": "IMAGE",
|
||||
"link": 30
|
||||
},
|
||||
{
|
||||
"name": "get_image_size",
|
||||
"type": "IMAGE",
|
||||
"link": 68
|
||||
},
|
||||
{
|
||||
"name": "width_input",
|
||||
"type": "INT",
|
||||
"link": null,
|
||||
"widget": {
|
||||
"name": "width_input"
|
||||
}
|
||||
},
|
||||
{
|
||||
"name": "height_input",
|
||||
"type": "INT",
|
||||
"link": null,
|
||||
"widget": {
|
||||
"name": "height_input"
|
||||
}
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
32
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "width",
|
||||
"type": "INT",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
},
|
||||
{
|
||||
"name": "height",
|
||||
"type": "INT",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "ImageResizeKJ"
|
||||
},
|
||||
"widgets_values": [
|
||||
512,
|
||||
512,
|
||||
"nearest-exact",
|
||||
false,
|
||||
2,
|
||||
0,
|
||||
0
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 18,
|
||||
"type": "ImageConcatMulti",
|
||||
"pos": [
|
||||
860,
|
||||
679
|
||||
],
|
||||
"size": {
|
||||
"0": 210,
|
||||
"1": 150
|
||||
},
|
||||
"flags": {},
|
||||
"order": 5,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image_1",
|
||||
"type": "IMAGE",
|
||||
"link": 32
|
||||
},
|
||||
{
|
||||
"name": "image_2",
|
||||
"type": "IMAGE",
|
||||
"link": 67
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
64
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {},
|
||||
"widgets_values": [
|
||||
2,
|
||||
"right",
|
||||
false,
|
||||
null
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 23,
|
||||
"type": "VHS_VideoCombine",
|
||||
"pos": [
|
||||
1098,
|
||||
240
|
||||
],
|
||||
"size": [
|
||||
1253.234130859375,
|
||||
940.6170654296875
|
||||
],
|
||||
"flags": {},
|
||||
"order": 6,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 64
|
||||
},
|
||||
{
|
||||
"name": "audio",
|
||||
"type": "VHS_AUDIO",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "meta_batch",
|
||||
"type": "VHS_BatchManager",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "vae",
|
||||
"type": "VAE",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "Filenames",
|
||||
"type": "VHS_FILENAMES",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VHS_VideoCombine"
|
||||
},
|
||||
"widgets_values": {
|
||||
"frame_rate": 24,
|
||||
"loop_count": 0,
|
||||
"filename_prefix": "LivePortrait",
|
||||
"format": "video/h264-mp4",
|
||||
"pix_fmt": "yuv420p",
|
||||
"crf": 19,
|
||||
"save_metadata": true,
|
||||
"pingpong": false,
|
||||
"save_output": false,
|
||||
"videopreview": {
|
||||
"hidden": false,
|
||||
"paused": false,
|
||||
"params": {
|
||||
"filename": "LivePortrait_00001.mp4",
|
||||
"subfolder": "",
|
||||
"type": "temp",
|
||||
"format": "video/h264-mp4",
|
||||
"frame_rate": 24
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 1,
|
||||
"type": "DownloadAndLoadLivePortraitModels",
|
||||
"pos": [
|
||||
142,
|
||||
205
|
||||
],
|
||||
"size": {
|
||||
"0": 252,
|
||||
"1": 58
|
||||
},
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"outputs": [
|
||||
{
|
||||
"name": "live_portrait_pipe",
|
||||
"type": "LIVEPORTRAITPIPE",
|
||||
"links": [
|
||||
58
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "DownloadAndLoadLivePortraitModels"
|
||||
},
|
||||
"widgets_values": [
|
||||
"auto"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 8,
|
||||
"type": "VHS_LoadVideo",
|
||||
"pos": [
|
||||
161,
|
||||
714
|
||||
],
|
||||
"size": [
|
||||
235.1999969482422,
|
||||
491.1999969482422
|
||||
],
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "meta_batch",
|
||||
"type": "VHS_BatchManager",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "vae",
|
||||
"type": "VAE",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
30,
|
||||
60
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "frame_count",
|
||||
"type": "INT",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
},
|
||||
{
|
||||
"name": "audio",
|
||||
"type": "VHS_AUDIO",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
},
|
||||
{
|
||||
"name": "video_info",
|
||||
"type": "VHS_VIDEOINFO",
|
||||
"links": null,
|
||||
"shape": 3
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VHS_LoadVideo"
|
||||
},
|
||||
"widgets_values": {
|
||||
"video": "d3.mp4",
|
||||
"force_rate": 0,
|
||||
"force_size": "Disabled",
|
||||
"custom_width": 512,
|
||||
"custom_height": 512,
|
||||
"frame_load_cap": 0,
|
||||
"skip_first_frames": 0,
|
||||
"select_every_nth": 1,
|
||||
"choose video to upload": "image",
|
||||
"videopreview": {
|
||||
"hidden": false,
|
||||
"paused": false,
|
||||
"params": {
|
||||
"frame_load_cap": 0,
|
||||
"skip_first_frames": 0,
|
||||
"force_rate": 0,
|
||||
"filename": "d3.mp4",
|
||||
"type": "input",
|
||||
"format": "video/mp4",
|
||||
"select_every_nth": 1
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 30,
|
||||
"type": "LivePortraitProcess",
|
||||
"pos": [
|
||||
500,
|
||||
249
|
||||
],
|
||||
"size": {
|
||||
"0": 367.79998779296875,
|
||||
"1": 362
|
||||
},
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "pipeline",
|
||||
"type": "LIVEPORTRAITPIPE",
|
||||
"link": 58
|
||||
},
|
||||
{
|
||||
"name": "source_image",
|
||||
"type": "IMAGE",
|
||||
"link": 59
|
||||
},
|
||||
{
|
||||
"name": "driving_images",
|
||||
"type": "IMAGE",
|
||||
"link": 60
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "cropped_images",
|
||||
"type": "IMAGE",
|
||||
"links": [],
|
||||
"shape": 3,
|
||||
"slot_index": 0
|
||||
},
|
||||
{
|
||||
"name": "full_images",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
67,
|
||||
68
|
||||
],
|
||||
"shape": 3,
|
||||
"slot_index": 1
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LivePortraitProcess"
|
||||
},
|
||||
"widgets_values": [
|
||||
512,
|
||||
2.3,
|
||||
0,
|
||||
-0.11,
|
||||
true,
|
||||
false,
|
||||
1,
|
||||
false,
|
||||
1,
|
||||
true,
|
||||
true,
|
||||
"CPU"
|
||||
]
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
30,
|
||||
8,
|
||||
0,
|
||||
19,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
32,
|
||||
19,
|
||||
0,
|
||||
18,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
58,
|
||||
1,
|
||||
0,
|
||||
30,
|
||||
0,
|
||||
"LIVEPORTRAITPIPE"
|
||||
],
|
||||
[
|
||||
59,
|
||||
4,
|
||||
0,
|
||||
30,
|
||||
1,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
60,
|
||||
8,
|
||||
0,
|
||||
30,
|
||||
2,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
64,
|
||||
18,
|
||||
0,
|
||||
23,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
67,
|
||||
30,
|
||||
1,
|
||||
18,
|
||||
1,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
68,
|
||||
30,
|
||||
1,
|
||||
19,
|
||||
1,
|
||||
"IMAGE"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 0.8264462809917354,
|
||||
"offset": {
|
||||
"0": 173.40487670898438,
|
||||
"1": -0.9636010527610779
|
||||
}
|
||||
}
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,44 +0,0 @@
|
||||
# coding: utf-8
|
||||
|
||||
"""
|
||||
config for user
|
||||
"""
|
||||
|
||||
import os.path as osp
|
||||
from dataclasses import dataclass
|
||||
#import tyro
|
||||
from typing_extensions import Annotated
|
||||
from .base_config import PrintableConfig, make_abs_path
|
||||
|
||||
|
||||
@dataclass(repr=False) # use repr from PrintableConfig
|
||||
class ArgumentConfig(PrintableConfig):
|
||||
########## input arguments ##########
|
||||
#source_image: Annotated[str, tyro.conf.arg(aliases=["-s"])] = make_abs_path('../../assets/examples/source/s6.jpg') # path to the reference portrait
|
||||
#driving_info: Annotated[str, tyro.conf.arg(aliases=["-d"])] = make_abs_path('../../assets/examples/driving/d0.mp4') # path to driving video or template (.pkl format)
|
||||
#output_dir: Annotated[str, tyro.conf.arg(aliases=["-o"])] = 'animations/' # directory to save output video
|
||||
#####################################
|
||||
|
||||
########## inference arguments ##########
|
||||
device_id: int = 0
|
||||
flag_lip_zero : bool = True # whether let the lip to close state before animation, only take effect when flag_eye_retargeting and flag_lip_retargeting is False
|
||||
flag_eye_retargeting: bool = False
|
||||
flag_lip_retargeting: bool = False
|
||||
flag_stitching: bool = True # we recommend setting it to True!
|
||||
flag_relative: bool = True # whether to use relative pose
|
||||
flag_pasteback: bool = True # whether to paste-back/stitch the animated face cropping from the face-cropping space to the original image space
|
||||
flag_do_crop: bool = True # whether to crop the reference portrait to the face-cropping space
|
||||
flag_do_rot: bool = True # whether to conduct the rotation when flag_do_crop is True
|
||||
#########################################
|
||||
|
||||
########## crop arguments ##########
|
||||
dsize: int = 512
|
||||
scale: float = 2.3
|
||||
vx_ratio: float = 0 # vx ratio
|
||||
vy_ratio: float = -0.125 # vy ratio +up, -down
|
||||
####################################
|
||||
|
||||
########## gradio arguments ##########
|
||||
#server_port: Annotated[int, tyro.conf.arg(aliases=["-p"])] = 8890
|
||||
#share: bool = False
|
||||
#server_name: str = "0.0.0.0"
|
||||
@@ -1,18 +0,0 @@
|
||||
# coding: utf-8
|
||||
|
||||
"""
|
||||
parameters used for crop faces
|
||||
"""
|
||||
|
||||
import os.path as osp
|
||||
from dataclasses import dataclass
|
||||
from typing import Union, List
|
||||
from .base_config import PrintableConfig
|
||||
|
||||
|
||||
@dataclass(repr=False) # use repr from PrintableConfig
|
||||
class CropConfig(PrintableConfig):
|
||||
dsize: int = 512 # crop size
|
||||
scale: float = 2.3 # scale factor
|
||||
vx_ratio: float = 0 # vx ratio
|
||||
vy_ratio: float = -0.125 # vy ratio +up, -down
|
||||
@@ -29,20 +29,15 @@ class InferenceConfig(PrintableConfig):
|
||||
flag_stitching: bool = True # we recommend setting it to True!
|
||||
|
||||
flag_relative: bool = True # whether to use relative pose
|
||||
anchor_frame: int = 0 # set this value if find_best_frame is True
|
||||
|
||||
input_shape: Tuple[int, int] = (256, 256) # input shape
|
||||
output_format: Literal['mp4', 'gif'] = 'mp4' # output video format
|
||||
output_fps: int = 30 # fps for output video
|
||||
crf: int = 15 # crf for output video
|
||||
|
||||
flag_write_result: bool = True # whether to write output video
|
||||
flag_pasteback: bool = True # whether to paste-back/stitch the animated face cropping from the face-cropping space to the original image space
|
||||
mask_crop = None
|
||||
flag_write_gif: bool = False
|
||||
size_gif: int = 256
|
||||
ref_max_shape: int = 1280
|
||||
ref_shape_n: int = 2
|
||||
|
||||
device_id: int = 0
|
||||
flag_do_crop: bool = False # whether to crop the reference portrait to the face-cropping space
|
||||
|
||||
@@ -4,184 +4,310 @@
|
||||
Pipeline of LivePortrait
|
||||
"""
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
import os.path as osp
|
||||
from tqdm import tqdm
|
||||
|
||||
from .config.inference_config import InferenceConfig
|
||||
|
||||
#from .utils.cropper import Cropper
|
||||
from .utils.camera import get_rotation_matrix
|
||||
#from .utils.video import images2video, concat_frames
|
||||
from .utils.crop import _transform_img
|
||||
#from .utils.retargeting_utils import calc_lip_close_ratio
|
||||
#from .utils.io import load_image_rgb, load_driving_info
|
||||
#from .utils.helper import mkdir, basename, dct2cuda, is_video, is_template, resize_to_limit
|
||||
from .utils.helper import resize_to_limit
|
||||
#from .utils.rprint import rlog as log
|
||||
from .live_portrait_wrapper import LivePortraitWrapper
|
||||
|
||||
import comfy.utils
|
||||
from tqdm import tqdm
|
||||
import numpy as np
|
||||
from .config.inference_config import InferenceConfig
|
||||
import torch
|
||||
from .utils.camera import get_rotation_matrix
|
||||
from .utils.crop import _transform_img, _transform_img_kornia
|
||||
from .live_portrait_wrapper import LivePortraitWrapper
|
||||
from .utils.retargeting_utils import calc_eye_close_ratio, calc_lip_close_ratio
|
||||
from .utils.filter import smooth
|
||||
|
||||
def make_abs_path(fn):
|
||||
return osp.join(osp.dirname(osp.realpath(__file__)), fn)
|
||||
|
||||
import os
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
class LivePortraitPipeline(object):
|
||||
|
||||
def __init__(self, appearance_feature_extractor, motion_extractor, warping_module,
|
||||
spade_generator, stitching_retargeting_module, inference_cfg: InferenceConfig):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
appearance_feature_extractor,
|
||||
motion_extractor,
|
||||
warping_module,
|
||||
spade_generator,
|
||||
stitching_retargeting_module,
|
||||
inference_cfg: InferenceConfig,
|
||||
):
|
||||
self.live_portrait_wrapper: LivePortraitWrapper = LivePortraitWrapper(
|
||||
appearance_feature_extractor, motion_extractor, warping_module,
|
||||
spade_generator, stitching_retargeting_module, cfg=inference_cfg)
|
||||
appearance_feature_extractor,
|
||||
motion_extractor,
|
||||
warping_module,
|
||||
spade_generator,
|
||||
stitching_retargeting_module,
|
||||
cfg=inference_cfg,
|
||||
)
|
||||
|
||||
def execute(self, img_rgb, driving_images_np):
|
||||
inference_cfg = self.live_portrait_wrapper.cfg # for convenience
|
||||
######## process reference portrait ########
|
||||
#img_rgb = load_image_rgb(args.source_image)
|
||||
img_rgb = resize_to_limit(img_rgb, inference_cfg.ref_max_shape, inference_cfg.ref_shape_n)
|
||||
#log(f"Load source image from {args.source_image}")
|
||||
crop_info = self.cropper.crop_single_image(img_rgb)
|
||||
source_lmk = crop_info['lmk_crop']
|
||||
_, img_crop_256x256 = crop_info['img_crop'], crop_info['img_crop_256x256']
|
||||
if inference_cfg.flag_do_crop:
|
||||
I_s = self.live_portrait_wrapper.prepare_source(img_crop_256x256)
|
||||
else:
|
||||
I_s = self.live_portrait_wrapper.prepare_source(img_rgb)
|
||||
x_s_info = self.live_portrait_wrapper.get_kp_info(I_s)
|
||||
x_c_s = x_s_info['kp']
|
||||
R_s = get_rotation_matrix(x_s_info['pitch'], x_s_info['yaw'], x_s_info['roll'])
|
||||
f_s = self.live_portrait_wrapper.extract_feature_3d(I_s)
|
||||
x_s = self.live_portrait_wrapper.transform_keypoint(x_s_info)
|
||||
def _get_source_frame(self, source_np, idx, method):
|
||||
if source_np.shape[0] == 1:
|
||||
return source_np[0]
|
||||
|
||||
if inference_cfg.flag_lip_zero:
|
||||
# let lip-open scalar to be 0 at first
|
||||
c_d_lip_before_animation = [0.]
|
||||
combined_lip_ratio_tensor_before_animation = self.live_portrait_wrapper.calc_combined_lip_ratio(c_d_lip_before_animation, source_lmk)
|
||||
if combined_lip_ratio_tensor_before_animation[0][0] < inference_cfg.lip_zero_threshold:
|
||||
inference_cfg.flag_lip_zero = False
|
||||
else:
|
||||
lip_delta_before_animation = self.live_portrait_wrapper.retarget_lip(x_s, combined_lip_ratio_tensor_before_animation)
|
||||
############################################
|
||||
if method == "constant":
|
||||
return source_np[min(idx, source_np.shape[0] - 1)]
|
||||
elif method == "cycle":
|
||||
return source_np[idx % source_np.shape[0]]
|
||||
elif method == "mirror":
|
||||
cycle_length = 2 * source_np.shape[0] - 2
|
||||
mirror_idx = idx % cycle_length
|
||||
if mirror_idx >= source_np.shape[0]:
|
||||
mirror_idx = cycle_length - mirror_idx
|
||||
return source_np[mirror_idx]
|
||||
|
||||
######## process driving info ########
|
||||
#if is_video(args.driving_info):
|
||||
#log(f"Load from video file (mp4 mov avi etc...): {args.driving_info}")
|
||||
# TODO: 这里track一下驱动视频 -> 构建模板
|
||||
#driving_rgb_lst = load_driving_info(args.driving_info)
|
||||
def execute(
|
||||
self, source_np, driving_images, crop_info, driving_landmarks, delta_multiplier, relative_motion_mode, driving_smooth_observation_variance, mismatch_method="constant",
|
||||
):
|
||||
inference_cfg = self.live_portrait_wrapper.cfg
|
||||
device = inference_cfg.device_id
|
||||
|
||||
driving_rgb_lst = driving_images_np
|
||||
|
||||
driving_rgb_lst_256 = [cv2.resize(_, (256, 256)) for _ in driving_rgb_lst]
|
||||
I_d_lst = self.live_portrait_wrapper.prepare_driving_videos(driving_rgb_lst_256)
|
||||
n_frames = I_d_lst.shape[0]
|
||||
if inference_cfg.flag_eye_retargeting or inference_cfg.flag_lip_retargeting:
|
||||
driving_lmk_lst = self.cropper.get_retargeting_lmk_info(driving_rgb_lst)
|
||||
input_eye_ratio_lst, input_lip_ratio_lst = self.live_portrait_wrapper.calc_retargeting_ratio(source_lmk, driving_lmk_lst)
|
||||
|
||||
# elif is_template(args.driving_info):
|
||||
# log(f"Load from video templates {args.driving_info}")
|
||||
# with open(args.driving_info, 'rb') as f:
|
||||
# template_lst, driving_lmk_lst = pickle.load(f)
|
||||
# n_frames = template_lst[0]['n_frames']
|
||||
# input_eye_ratio_lst, input_lip_ratio_lst = self.live_portrait_wrapper.calc_retargeting_ratio(source_lmk, driving_lmk_lst)
|
||||
# else:
|
||||
# raise Exception("Unsupported driving types!")
|
||||
#########################################
|
||||
|
||||
######## prepare for pasteback ########
|
||||
if inference_cfg.flag_pasteback:
|
||||
if inference_cfg.mask_crop is None:
|
||||
inference_cfg.mask_crop = cv2.imread(make_abs_path('./utils/resources/mask_template.png'), cv2.IMREAD_COLOR)
|
||||
mask_ori = _transform_img(inference_cfg.mask_crop, crop_info['M_c2o'], dsize=(img_rgb.shape[1], img_rgb.shape[0]))
|
||||
mask_ori = mask_ori.astype(np.float32) / 255.
|
||||
I_p_paste_lst = []
|
||||
#########################################
|
||||
|
||||
I_p_lst = []
|
||||
cropped_image_list = []
|
||||
composited_image_list = []
|
||||
out_mask_list = []
|
||||
R_d_0, x_d_0_info = None, None
|
||||
pbar = comfy.utils.ProgressBar(n_frames)
|
||||
for i in tqdm(range(n_frames), desc='Animating...', total=n_frames):
|
||||
#if is_video(args.driving_info):
|
||||
# extract kp info by M
|
||||
I_d_i = I_d_lst[i]
|
||||
x_d_i_info = self.live_portrait_wrapper.get_kp_info(I_d_i)
|
||||
R_d_i = get_rotation_matrix(x_d_i_info['pitch'], x_d_i_info['yaw'], x_d_i_info['roll'])
|
||||
# else:
|
||||
# # from template
|
||||
# x_d_i_info = template_lst[i]
|
||||
# x_d_i_info = dct2cuda(x_d_i_info, inference_cfg.device_id)
|
||||
# R_d_i = x_d_i_info['R_d']
|
||||
|
||||
if mismatch_method == "cut":
|
||||
total_frames = source_np.shape[0]
|
||||
else:
|
||||
total_frames = driving_images.shape[0]
|
||||
|
||||
|
||||
source_info = []
|
||||
source_rot_list = []
|
||||
f_s_list = []
|
||||
for i in tqdm(range(source_np.shape[0]), desc='Processing source images...', total=source_np.shape[0]):
|
||||
#get source keypoints info
|
||||
img_crop_256x256 = crop_info["crop_info_list"][i]["img_crop_256x256"]
|
||||
I_s = self.live_portrait_wrapper.prepare_source(img_crop_256x256)
|
||||
x_s_info = self.live_portrait_wrapper.get_kp_info(I_s)
|
||||
f_s = self.live_portrait_wrapper.extract_feature_3d(I_s)
|
||||
f_s_list.append(f_s)
|
||||
source_info.append(x_s_info)
|
||||
|
||||
R_s = get_rotation_matrix(
|
||||
x_s_info["pitch"], x_s_info["yaw"], x_s_info["roll"]
|
||||
)
|
||||
source_rot_list.append(R_s)
|
||||
|
||||
driving_info = []
|
||||
driving_exp_list = []
|
||||
driving_rot_list = []
|
||||
|
||||
for i in tqdm(range(driving_images.shape[0]), desc='Processing driving images...', total=driving_images.shape[0]):
|
||||
#get driving keypoints info
|
||||
x_d_info = self.live_portrait_wrapper.get_kp_info(driving_images[i].unsqueeze(0).to(device))
|
||||
safe_index = min(i, len(crop_info["crop_info_list"]) - 1)
|
||||
if i == 0:
|
||||
first = x_d_info
|
||||
|
||||
driving_info.append(x_d_info)
|
||||
|
||||
driving_exp = source_info[safe_index]["exp"] + x_d_info["exp"] - first["exp"]
|
||||
driving_exp_list.append(driving_exp.cpu())
|
||||
|
||||
R_d = get_rotation_matrix(
|
||||
x_d_info["pitch"], x_d_info["yaw"], x_d_info["roll"]
|
||||
)
|
||||
driving_rot_list.append(R_d)
|
||||
|
||||
if relative_motion_mode == "source_video_smoothed":
|
||||
x_d_r_lst = []
|
||||
first_driving_rot = driving_rot_list[0].cpu().numpy().astype(np.float32).transpose(0, 2, 1)
|
||||
for i in tqdm(range(source_np.shape[0]), desc='Smoothing...', total=source_np.shape[0]):
|
||||
driving_rot = driving_rot_list[i].cpu().numpy().astype(np.float32)
|
||||
source_rot = source_rot_list[i].cpu().numpy().astype(np.float32)
|
||||
dot = np.dot(driving_rot, first_driving_rot) @ source_rot
|
||||
x_d_r_lst.append(dot)
|
||||
|
||||
driving_exp_list_smooth = smooth(driving_exp_list, source_info[0]["exp"].shape, device, observation_variance=driving_smooth_observation_variance)
|
||||
driving_rot_list_smooth = smooth(x_d_r_lst, source_rot_list[0].shape, device, observation_variance=driving_smooth_observation_variance)
|
||||
|
||||
pbar = comfy.utils.ProgressBar(total_frames)
|
||||
|
||||
for i in tqdm(range(total_frames), desc='Animating...', total=total_frames):
|
||||
|
||||
safe_index = min(i, len(crop_info["crop_info_list"]) - 1)
|
||||
|
||||
# skip and return empty frames if no crop due to no face detected
|
||||
if not crop_info["crop_info_list"][safe_index]:
|
||||
composited_image_list.append(source_np[safe_index])
|
||||
cropped_image_list.append(torch.zeros(1, 512, 512, 3, dtype=torch.float32, device = device))
|
||||
out_mask_list.append(np.zeros((source_np.shape[1], source_np.shape[2], 3), dtype=np.uint8))
|
||||
continue
|
||||
|
||||
source_lmk = crop_info["crop_info_list"][safe_index]["lmk_crop"]
|
||||
|
||||
x_d_info = driving_info[i]
|
||||
x_s_info = source_info[safe_index]
|
||||
|
||||
x_c_s = x_s_info["kp"]
|
||||
|
||||
R_s = source_rot_list[safe_index]
|
||||
f_s = f_s_list[safe_index]
|
||||
x_s = self.live_portrait_wrapper.transform_keypoint(x_s_info)
|
||||
|
||||
#lip zero
|
||||
if inference_cfg.flag_lip_zero:
|
||||
c_d_lip_before_animation = [0.0]
|
||||
combined_lip_ratio_tensor_before_animation = (self.live_portrait_wrapper.calc_combined_lip_ratio(c_d_lip_before_animation, source_lmk))
|
||||
|
||||
if (combined_lip_ratio_tensor_before_animation[0][0] < inference_cfg.lip_zero_threshold):
|
||||
inference_cfg.flag_lip_zero = False
|
||||
else:
|
||||
lip_delta_before_animation = (self.live_portrait_wrapper.retarget_lip(x_s, combined_lip_ratio_tensor_before_animation))
|
||||
|
||||
R_d = driving_rot_list[i]
|
||||
|
||||
if i == 0:
|
||||
R_d_0 = R_d_i
|
||||
x_d_0_info = x_d_i_info
|
||||
R_d_0 = R_d
|
||||
x_d_0_info = x_d_info
|
||||
|
||||
if inference_cfg.flag_relative:
|
||||
R_new = (R_d_i @ R_d_0.permute(0, 2, 1)) @ R_s
|
||||
delta_new = x_s_info['exp'] + (x_d_i_info['exp'] - x_d_0_info['exp'])
|
||||
scale_new = x_s_info['scale'] * (x_d_i_info['scale'] / x_d_0_info['scale'])
|
||||
t_new = x_s_info['t'] + (x_d_i_info['t'] - x_d_0_info['t'])
|
||||
if relative_motion_mode == "relative":
|
||||
R_new = (R_d @ R_d_0.permute(0, 2, 1)) @ R_s
|
||||
delta_new = x_s_info["exp"] + (x_d_info["exp"] - x_d_0_info["exp"])
|
||||
scale_new = x_s_info["scale"] * (x_d_info["scale"] / x_d_0_info["scale"])
|
||||
t_new = x_s_info["t"] + (x_d_info["t"] - x_d_0_info["t"])
|
||||
elif relative_motion_mode == "source_video_smoothed":
|
||||
R_new = driving_rot_list_smooth[i]
|
||||
delta_new = driving_exp_list_smooth[i]
|
||||
scale_new = x_s_info["scale"]
|
||||
t_new = x_d_info["t"]
|
||||
elif relative_motion_mode == "relative_rotation_only":
|
||||
R_new = R_s
|
||||
delta_new = x_s_info['exp']
|
||||
scale_new = x_s_info["scale"]
|
||||
t_new = x_d_info["t"]
|
||||
else:
|
||||
R_new = R_d_i
|
||||
delta_new = x_d_i_info['exp']
|
||||
scale_new = x_s_info['scale']
|
||||
t_new = x_d_i_info['t']
|
||||
R_new = R_d
|
||||
delta_new = x_s_info['exp']
|
||||
scale_new = x_s_info["scale"]
|
||||
t_new = x_d_info["t"]
|
||||
|
||||
t_new[..., 2].fill_(0) # zero tz
|
||||
t_new[..., 2].fill_(0) # zero tz
|
||||
|
||||
delta_new = delta_new * delta_multiplier
|
||||
|
||||
x_d_i_new = scale_new * (x_c_s @ R_new + delta_new) + t_new
|
||||
|
||||
# Algorithm 1:
|
||||
if not inference_cfg.flag_stitching and not inference_cfg.flag_eye_retargeting and not inference_cfg.flag_lip_retargeting:
|
||||
if (
|
||||
not inference_cfg.flag_stitching
|
||||
and not inference_cfg.flag_eye_retargeting
|
||||
and not inference_cfg.flag_lip_retargeting
|
||||
):
|
||||
# without stitching or retargeting
|
||||
if inference_cfg.flag_lip_zero:
|
||||
x_d_i_new += lip_delta_before_animation.reshape(-1, x_s.shape[1], 3)
|
||||
else:
|
||||
pass
|
||||
elif inference_cfg.flag_stitching and not inference_cfg.flag_eye_retargeting and not inference_cfg.flag_lip_retargeting:
|
||||
elif (
|
||||
inference_cfg.flag_stitching
|
||||
and not inference_cfg.flag_eye_retargeting
|
||||
and not inference_cfg.flag_lip_retargeting
|
||||
):
|
||||
# with stitching and without retargeting
|
||||
if inference_cfg.flag_lip_zero:
|
||||
x_d_i_new = self.live_portrait_wrapper.stitching(x_s, x_d_i_new) + lip_delta_before_animation.reshape(-1, x_s.shape[1], 3)
|
||||
x_d_i_new = self.live_portrait_wrapper.stitching(
|
||||
x_s, x_d_i_new
|
||||
) + lip_delta_before_animation.reshape(-1, x_s.shape[1], 3)
|
||||
else:
|
||||
x_d_i_new = self.live_portrait_wrapper.stitching(x_s, x_d_i_new)
|
||||
else:
|
||||
eyes_delta, lip_delta = None, None
|
||||
if inference_cfg.flag_eye_retargeting:
|
||||
c_d_eyes_i = input_eye_ratio_lst[i]
|
||||
combined_eye_ratio_tensor = self.live_portrait_wrapper.calc_combined_eye_ratio(c_d_eyes_i, source_lmk)
|
||||
combined_eye_ratio_tensor = combined_eye_ratio_tensor * inference_cfg.eyes_retargeting_multiplier
|
||||
c_d_eyes_i = calc_eye_close_ratio(driving_landmarks[i][None])
|
||||
combined_eye_ratio_tensor = (
|
||||
self.live_portrait_wrapper.calc_combined_eye_ratio(
|
||||
c_d_eyes_i, source_lmk
|
||||
)
|
||||
)
|
||||
combined_eye_ratio_tensor = (
|
||||
combined_eye_ratio_tensor
|
||||
* inference_cfg.eyes_retargeting_multiplier
|
||||
)
|
||||
# ∆_eyes,i = R_eyes(x_s; c_s,eyes, c_d,eyes,i)
|
||||
eyes_delta = self.live_portrait_wrapper.retarget_eye(x_s, combined_eye_ratio_tensor)
|
||||
eyes_delta = self.live_portrait_wrapper.retarget_eye(
|
||||
x_s, combined_eye_ratio_tensor
|
||||
)
|
||||
if inference_cfg.flag_lip_retargeting:
|
||||
c_d_lip_i = input_lip_ratio_lst[i]
|
||||
combined_lip_ratio_tensor = self.live_portrait_wrapper.calc_combined_lip_ratio(c_d_lip_i, source_lmk)
|
||||
combined_lip_ratio_tensor = combined_lip_ratio_tensor * inference_cfg.lip_retargeting_multiplier
|
||||
c_d_lip_i = calc_lip_close_ratio(driving_landmarks[i][None])
|
||||
combined_lip_ratio_tensor = (
|
||||
self.live_portrait_wrapper.calc_combined_lip_ratio(
|
||||
c_d_lip_i, source_lmk
|
||||
)
|
||||
)
|
||||
combined_lip_ratio_tensor = (
|
||||
combined_lip_ratio_tensor
|
||||
* inference_cfg.lip_retargeting_multiplier
|
||||
)
|
||||
# ∆_lip,i = R_lip(x_s; c_s,lip, c_d,lip,i)
|
||||
lip_delta = self.live_portrait_wrapper.retarget_lip(x_s, combined_lip_ratio_tensor)
|
||||
lip_delta = self.live_portrait_wrapper.retarget_lip(
|
||||
x_s, combined_lip_ratio_tensor
|
||||
)
|
||||
|
||||
if inference_cfg.flag_relative: # use x_s
|
||||
x_d_i_new = x_s + \
|
||||
(eyes_delta.reshape(-1, x_s.shape[1], 3) if eyes_delta is not None else 0) + \
|
||||
(lip_delta.reshape(-1, x_s.shape[1], 3) if lip_delta is not None else 0)
|
||||
x_d_i_new = (
|
||||
x_s
|
||||
+ (
|
||||
eyes_delta.reshape(-1, x_s.shape[1], 3)
|
||||
if eyes_delta is not None
|
||||
else 0
|
||||
)
|
||||
+ (
|
||||
lip_delta.reshape(-1, x_s.shape[1], 3)
|
||||
if lip_delta is not None
|
||||
else 0
|
||||
)
|
||||
)
|
||||
else: # use x_d,i
|
||||
x_d_i_new = x_d_i_new + \
|
||||
(eyes_delta.reshape(-1, x_s.shape[1], 3) if eyes_delta is not None else 0) + \
|
||||
(lip_delta.reshape(-1, x_s.shape[1], 3) if lip_delta is not None else 0)
|
||||
x_d_i_new = (
|
||||
x_d_i_new
|
||||
+ (
|
||||
eyes_delta.reshape(-1, x_s.shape[1], 3)
|
||||
if eyes_delta is not None
|
||||
else 0
|
||||
)
|
||||
+ (
|
||||
lip_delta.reshape(-1, x_s.shape[1], 3)
|
||||
if lip_delta is not None
|
||||
else 0
|
||||
)
|
||||
)
|
||||
|
||||
if inference_cfg.flag_stitching:
|
||||
x_d_i_new = self.live_portrait_wrapper.stitching(x_s, x_d_i_new)
|
||||
|
||||
if inference_cfg.flag_stitching:
|
||||
x_d_i_new = self.live_portrait_wrapper.stitching(x_s, x_d_i_new)
|
||||
|
||||
out = self.live_portrait_wrapper.warp_decode(f_s, x_s, x_d_i_new)
|
||||
I_p_i = self.live_portrait_wrapper.parse_output(out['out'])[0]
|
||||
I_p_lst.append(I_p_i)
|
||||
|
||||
cropped_image = torch.clamp(out["out"], 0, 1).permute(0, 2, 3, 1)
|
||||
|
||||
cropped_image_list.append(cropped_image)
|
||||
|
||||
if mismatch_method == "cut" or inference_cfg.flag_eye_retargeting or inference_cfg.flag_lip_retargeting:
|
||||
source_frame_rgb = source_np[safe_index]
|
||||
else:
|
||||
source_frame_rgb = self._get_source_frame(source_np, i, mismatch_method)
|
||||
|
||||
# Transform and blend
|
||||
if inference_cfg.flag_pasteback:
|
||||
cropped_image_to_original = _transform_img_kornia(
|
||||
cropped_image,
|
||||
crop_info["crop_info_list"][safe_index]["M_c2o"],
|
||||
dsize=(source_frame_rgb.shape[1], source_frame_rgb.shape[0]),
|
||||
)
|
||||
|
||||
mask_ori = _transform_img_kornia(
|
||||
inference_cfg.mask_crop,
|
||||
crop_info["crop_info_list"][safe_index]["M_c2o"],
|
||||
dsize=(source_frame_rgb.shape[1], source_frame_rgb.shape[0]),
|
||||
)
|
||||
|
||||
source_frame_torch = torch.from_numpy(source_frame_rgb).unsqueeze(0).permute(0, 3, 1, 2).to(mask_ori.device) / 255
|
||||
|
||||
cropped_image_to_original_blend = torch.clip(
|
||||
mask_ori * cropped_image_to_original + (1 - mask_ori) * source_frame_torch, 0, 1
|
||||
)
|
||||
|
||||
composited_image_list.append(cropped_image_to_original_blend)
|
||||
out_mask_list.append(mask_ori)
|
||||
pbar.update(1)
|
||||
|
||||
#if inference_cfg.flag_pasteback:
|
||||
I_p_i_to_ori = _transform_img(I_p_i, crop_info['M_c2o'], dsize=(img_rgb.shape[1], img_rgb.shape[0]))
|
||||
I_p_i_to_ori_blend = np.clip(mask_ori * I_p_i_to_ori + (1 - mask_ori) * img_rgb, 0, 255).astype(np.uint8)
|
||||
out = np.hstack([I_p_i_to_ori, I_p_i_to_ori_blend])
|
||||
I_p_paste_lst.append(I_p_i_to_ori_blend)
|
||||
|
||||
return I_p_lst, I_p_paste_lst
|
||||
return cropped_image_list, composited_image_list, out_mask_list
|
||||
|
||||
@@ -32,11 +32,6 @@ class LivePortraitWrapper(object):
|
||||
self.device_id = cfg.device_id
|
||||
self.timer = Timer()
|
||||
|
||||
def update_config(self, user_args):
|
||||
for k, v in user_args.items():
|
||||
if hasattr(self.cfg, k):
|
||||
setattr(self.cfg, k, v)
|
||||
|
||||
def prepare_source(self, img: np.ndarray) -> torch.Tensor:
|
||||
""" construct the input as standard
|
||||
img: HxWx3, uint8, 256x256
|
||||
@@ -58,24 +53,6 @@ class LivePortraitWrapper(object):
|
||||
x = x.to(self.device_id)
|
||||
return x
|
||||
|
||||
def prepare_driving_videos(self, imgs) -> torch.Tensor:
|
||||
""" construct the input as standard
|
||||
imgs: NxBxHxWx3, uint8
|
||||
"""
|
||||
if isinstance(imgs, list):
|
||||
_imgs = np.array(imgs)[..., np.newaxis] # TxHxWx3x1
|
||||
elif isinstance(imgs, np.ndarray):
|
||||
_imgs = imgs
|
||||
else:
|
||||
raise ValueError(f'imgs type error: {type(imgs)}')
|
||||
|
||||
y = _imgs.astype(np.float32) / 255.
|
||||
y = np.clip(y, 0, 1) # clip to 0~1
|
||||
y = torch.from_numpy(y).permute(0, 4, 3, 1, 2) # TxHxWx3x1 -> Tx1x3xHxW
|
||||
y = y.to(self.device_id)
|
||||
|
||||
return y
|
||||
|
||||
def extract_feature_3d(self, x: torch.Tensor) -> torch.Tensor:
|
||||
""" get the appearance feature of the image by F
|
||||
x: Bx3xHxW, normalized to 0~1
|
||||
@@ -264,7 +241,7 @@ class LivePortraitWrapper(object):
|
||||
kp_source: BxNx3
|
||||
kp_driving: BxNx3
|
||||
"""
|
||||
# The line 18 in Algorithm 1: D(W(f_s; x_s, x′_d,i))
|
||||
# The line 18 in Algorithm 1: D(W(f_s; x_s, x′_d,i)
|
||||
with torch.autocast(get_autocast_device(self.device_id), dtype=torch.float16) if self.cfg.flag_use_half_precision else nullcontext():
|
||||
# get decoder input
|
||||
ret_dct = self.warping_module(feature_3d, kp_source=kp_source, kp_driving=kp_driving)
|
||||
@@ -301,17 +278,19 @@ class LivePortraitWrapper(object):
|
||||
|
||||
def calc_combined_eye_ratio(self, input_eye_ratio, source_lmk):
|
||||
eye_close_ratio = calc_eye_close_ratio(source_lmk[None])
|
||||
eye_close_ratio_tensor = torch.from_numpy(eye_close_ratio).float().to(self.device_id)
|
||||
input_eye_ratio_tensor = torch.Tensor([input_eye_ratio[0][0]]).reshape(1, 1).to(self.device_id)
|
||||
eye_close_ratios_tensor = torch.from_numpy(eye_close_ratio).float().to(self.device_id)
|
||||
input_eye_ratio_array = np.array(input_eye_ratio[0][0]).reshape(1, 1)
|
||||
input_eye_ratio_tensor = torch.from_numpy(input_eye_ratio_array).float().to(self.device_id)
|
||||
# [c_s,eyes, c_d,eyes,i]
|
||||
combined_eye_ratio_tensor = torch.cat([eye_close_ratio_tensor, input_eye_ratio_tensor], dim=1)
|
||||
return combined_eye_ratio_tensor
|
||||
combined_eye_ratios_tensor = torch.cat([eye_close_ratios_tensor, input_eye_ratio_tensor], dim=1)
|
||||
return combined_eye_ratios_tensor
|
||||
|
||||
def calc_combined_lip_ratio(self, input_lip_ratio, source_lmk):
|
||||
lip_close_ratio = calc_lip_close_ratio(source_lmk[None])
|
||||
lip_close_ratio_tensor = torch.from_numpy(lip_close_ratio).float().to(self.device_id)
|
||||
# [c_s,lip, c_d,lip,i]
|
||||
input_lip_ratio_tensor = torch.Tensor([input_lip_ratio[0]]).to(self.device_id)
|
||||
input_lip_ratio_array = np.array([input_lip_ratio[0]])
|
||||
input_lip_ratio_tensor = torch.from_numpy(input_lip_ratio_array).float().to(self.device_id)
|
||||
if input_lip_ratio_tensor.shape != [1, 1]:
|
||||
input_lip_ratio_tensor = input_lip_ratio_tensor.reshape(1, 1)
|
||||
combined_lip_ratio_tensor = torch.cat([lip_close_ratio_tensor, input_lip_ratio_tensor], dim=1)
|
||||
|
||||
@@ -1,65 +0,0 @@
|
||||
# coding: utf-8
|
||||
|
||||
"""
|
||||
Make video template
|
||||
"""
|
||||
|
||||
import os
|
||||
import cv2
|
||||
import numpy as np
|
||||
import pickle
|
||||
from tqdm import tqdm
|
||||
from .utils.cropper import Cropper
|
||||
|
||||
from .utils.io import load_driving_info
|
||||
from .utils.camera import get_rotation_matrix
|
||||
from .utils.helper import mkdir, basename
|
||||
from .utils.rprint import rlog as log
|
||||
from .config.crop_config import CropConfig
|
||||
from .config.inference_config import InferenceConfig
|
||||
from .live_portrait_wrapper import LivePortraitWrapper
|
||||
|
||||
class TemplateMaker:
|
||||
|
||||
def __init__(self, inference_cfg: InferenceConfig, crop_cfg: CropConfig):
|
||||
self.live_portrait_wrapper: LivePortraitWrapper = LivePortraitWrapper(cfg=inference_cfg)
|
||||
self.cropper = Cropper(crop_cfg=crop_cfg)
|
||||
|
||||
def make_motion_template(self, video_fp: str, output_path: str, **kwargs):
|
||||
""" make video template (.pkl format)
|
||||
video_fp: driving video file path
|
||||
output_path: where to save the pickle file
|
||||
"""
|
||||
|
||||
driving_rgb_lst = load_driving_info(video_fp)
|
||||
driving_rgb_lst = [cv2.resize(_, (256, 256)) for _ in driving_rgb_lst]
|
||||
driving_lmk_lst = self.cropper.get_retargeting_lmk_info(driving_rgb_lst)
|
||||
I_d_lst = self.live_portrait_wrapper.prepare_driving_videos(driving_rgb_lst)
|
||||
|
||||
n_frames = I_d_lst.shape[0]
|
||||
|
||||
templates = []
|
||||
|
||||
|
||||
for i in tqdm(range(n_frames), desc='Making templates...', total=n_frames):
|
||||
I_d_i = I_d_lst[i]
|
||||
x_d_i_info = self.live_portrait_wrapper.get_kp_info(I_d_i)
|
||||
R_d_i = get_rotation_matrix(x_d_i_info['pitch'], x_d_i_info['yaw'], x_d_i_info['roll'])
|
||||
# collect s_d, R_d, δ_d and t_d for inference
|
||||
template_dct = {
|
||||
'n_frames': n_frames,
|
||||
'frames_index': i,
|
||||
}
|
||||
template_dct['scale'] = x_d_i_info['scale'].cpu().numpy().astype(np.float32)
|
||||
template_dct['R_d'] = R_d_i.cpu().numpy().astype(np.float32)
|
||||
template_dct['exp'] = x_d_i_info['exp'].cpu().numpy().astype(np.float32)
|
||||
template_dct['t'] = x_d_i_info['t'].cpu().numpy().astype(np.float32)
|
||||
|
||||
templates.append(template_dct)
|
||||
|
||||
mkdir(output_path)
|
||||
# Save the dictionary as a pickle file
|
||||
pickle_fp = os.path.join(output_path, f'{basename(video_fp)}.pkl')
|
||||
with open(pickle_fp, 'wb') as f:
|
||||
pickle.dump([templates, driving_lmk_lst], f)
|
||||
log(f"Template saved at {pickle_fp}")
|
||||
@@ -4,14 +4,12 @@
|
||||
cropping function and the related preprocess functions for cropping
|
||||
"""
|
||||
|
||||
import cv2; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False) # NOTE: enforce single thread
|
||||
import cv2#; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False) # NOTE: enforce single thread
|
||||
import numpy as np
|
||||
from .rprint import rprint as print
|
||||
from math import sin, cos, acos, degrees
|
||||
|
||||
DTYPE = np.float32
|
||||
CV2_INTERP = cv2.INTER_LINEAR
|
||||
|
||||
import comfy.model_management as mm
|
||||
|
||||
def _transform_img(img, M, dsize, flags=CV2_INTERP, borderMode=None):
|
||||
""" conduct similarity or affine transformation to the image, do not do border operation!
|
||||
@@ -29,6 +27,43 @@ def _transform_img(img, M, dsize, flags=CV2_INTERP, borderMode=None):
|
||||
else:
|
||||
return cv2.warpAffine(img, M[:2, :], dsize=_dsize, flags=flags)
|
||||
|
||||
import torch
|
||||
import kornia.geometry.transform as KGT
|
||||
|
||||
def _transform_img_kornia(img, M, dsize, flags='bilinear', borderMode='zeros'):
|
||||
"""Conduct similarity or affine transformation to the image using Kornia.
|
||||
|
||||
img: Input image as a PyTorch tensor of shape (C, H, W).
|
||||
M: 2x3 transformation matrix as a PyTorch tensor.
|
||||
dsize: Target shape (width, height).
|
||||
"""
|
||||
device = mm.get_torch_device()
|
||||
# Convert dsize to tensor shape (H, W)
|
||||
_dsize = torch.tensor([dsize[1], dsize[0]]) # Kornia expects (H, W)
|
||||
|
||||
# Convert M from numpy.ndarray to PyTorch tensor
|
||||
M = torch.from_numpy(M).float().to(device)
|
||||
if M.shape == (3, 3):
|
||||
M = M[:2, :].unsqueeze(0) # Adjust M to the expected shape Bx2x3
|
||||
elif M.shape == (2, 3):
|
||||
M = M.unsqueeze(0) # Add batch dimension if not present
|
||||
|
||||
# Reshape M for Kornia (1, 2, 3) and upscale to 3D affine matrix if not already
|
||||
if M.shape == (2, 3):
|
||||
M = M.unsqueeze(0) # Add batch dimension
|
||||
|
||||
# Convert image to floating point tensor if not already
|
||||
if img.dtype != torch.float32:
|
||||
img = img.float()
|
||||
img = img.to(device)
|
||||
|
||||
# Reshape img for Kornia (B, C, H, W)
|
||||
img = img.permute(0, 3, 1, 2)
|
||||
|
||||
# Apply the affine transformation
|
||||
img_warped = KGT.warp_affine(img, M, _dsize, mode=flags, padding_mode=borderMode)
|
||||
|
||||
return img_warped
|
||||
|
||||
def _transform_pts(pts, M):
|
||||
""" conduct similarity or affine transformation to the pts
|
||||
@@ -350,13 +385,15 @@ def crop_image(img, pts: np.ndarray, **kwargs):
|
||||
dsize = kwargs.get('dsize', 224)
|
||||
scale = kwargs.get('scale', 1.5) # 1.5 | 1.6
|
||||
vy_ratio = kwargs.get('vy_ratio', -0.1) # -0.0625 | -0.1
|
||||
vx_ratio = kwargs.get('vx_ratio', 0)
|
||||
|
||||
M_INV, _ = _estimate_similar_transform_from_pts(
|
||||
pts,
|
||||
dsize=dsize,
|
||||
scale=scale,
|
||||
vy_ratio=vy_ratio,
|
||||
flag_do_rot=kwargs.get('flag_do_rot', True),
|
||||
vx_ratio=vx_ratio,
|
||||
flag_do_rot=kwargs.get('rotate', True),
|
||||
)
|
||||
|
||||
if img is None:
|
||||
|
||||
@@ -1,32 +1,22 @@
|
||||
# coding: utf-8
|
||||
|
||||
import numpy as np
|
||||
import os.path as osp
|
||||
from typing import List, Union, Tuple
|
||||
from dataclasses import dataclass, field
|
||||
import cv2; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False)
|
||||
import cv2#; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False)
|
||||
|
||||
from .landmark_runner import LandmarkRunner
|
||||
from .face_analysis_diy import FaceAnalysisDIY
|
||||
#from .helper import prefix
|
||||
from .crop import crop_image, crop_image_by_bbox, parse_bbox_from_landmark, average_bbox_lst
|
||||
#from .timer import Timer
|
||||
from .rprint import rlog as log
|
||||
from .io import load_image_rgb
|
||||
#from .video import VideoWriter, get_fps, change_video_fps
|
||||
from .crop import crop_image
|
||||
|
||||
import folder_paths
|
||||
import os
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
def make_abs_path(fn):
|
||||
return osp.join(osp.dirname(osp.realpath(__file__)), fn)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Trajectory:
|
||||
start: int = -1 # 起始帧 闭区间
|
||||
end: int = -1 # 结束帧 闭区间
|
||||
start: int = -1
|
||||
end: int = -1
|
||||
lmk_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # lmk list
|
||||
bbox_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # bbox list
|
||||
frame_rgb_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # frame list
|
||||
@@ -34,10 +24,10 @@ class Trajectory:
|
||||
|
||||
|
||||
class Cropper(object):
|
||||
def __init__(self, provider, **kwargs) -> None:
|
||||
def __init__(self, **kwargs) -> None:
|
||||
device_id = kwargs.get('device_id', 0)
|
||||
provider = kwargs.get('onnx_device', 'CPU')
|
||||
self.landmark_runner = LandmarkRunner(
|
||||
#ckpt_path=make_abs_path('../../pretrained_weights/liveportrait/landmark.onnx'),
|
||||
ckpt_path=os.path.join(folder_paths.models_dir, 'liveportrait', 'landmark.onnx'),
|
||||
onnx_provider=provider,
|
||||
device_id=device_id
|
||||
@@ -52,21 +42,8 @@ class Cropper(object):
|
||||
self.face_analysis_wrapper.prepare(ctx_id=device_id, det_size=(512, 512))
|
||||
self.face_analysis_wrapper.warmup()
|
||||
|
||||
self.crop_cfg = kwargs.get('crop_cfg', None)
|
||||
|
||||
def update_config(self, user_args):
|
||||
for k, v in user_args.items():
|
||||
if hasattr(self.crop_cfg, k):
|
||||
setattr(self.crop_cfg, k, v)
|
||||
|
||||
def crop_single_image(self, obj, **kwargs):
|
||||
direction = kwargs.get('direction', 'large-small')
|
||||
|
||||
# crop and align a single image
|
||||
if isinstance(obj, str):
|
||||
img_rgb = load_image_rgb(obj)
|
||||
elif isinstance(obj, np.ndarray):
|
||||
img_rgb = obj
|
||||
def crop_single_image(self, img_rgb, dsize, scale, vy_ratio, vx_ratio, face_index, face_index_order, rotate):
|
||||
direction = face_index_order
|
||||
|
||||
src_face = self.face_analysis_wrapper.get(
|
||||
img_rgb,
|
||||
@@ -75,73 +52,34 @@ class Cropper(object):
|
||||
)
|
||||
|
||||
if len(src_face) == 0:
|
||||
log('No face detected in the source image.')
|
||||
raise Exception("No face detected in the source image!")
|
||||
elif len(src_face) > 1:
|
||||
log(f'More than one face detected in the image, only pick one face by rule {direction}.')
|
||||
ret_dct = {}
|
||||
return ret_dct
|
||||
#raise Exception("No face detected in the source image!")
|
||||
#elif len(src_face) > 1:
|
||||
# print(f'More than one face detected in the image, only pick one face by rule {direction}.')
|
||||
|
||||
src_face = src_face[0]
|
||||
src_face = src_face[face_index] # choose the index if multiple faces detected
|
||||
pts = src_face.landmark_2d_106
|
||||
|
||||
|
||||
# crop the face
|
||||
ret_dct = crop_image(
|
||||
img_rgb, # ndarray
|
||||
pts, # 106x2 or Nx2
|
||||
dsize=kwargs.get('dsize', 512),
|
||||
scale=kwargs.get('scale', 2.3),
|
||||
vy_ratio=kwargs.get('vy_ratio', -0.15),
|
||||
dsize=dsize,
|
||||
scale=scale,
|
||||
vy_ratio=vy_ratio,
|
||||
vx_ratio=vx_ratio,
|
||||
rotate=rotate
|
||||
)
|
||||
# update a 256x256 version for network input or else
|
||||
ret_dct['img_crop_256x256'] = cv2.resize(ret_dct['img_crop'], (256, 256), interpolation=cv2.INTER_AREA)
|
||||
ret_dct['pt_crop_256x256'] = ret_dct['pt_crop'] * 256 / kwargs.get('dsize', 512)
|
||||
ret_dct['pt_crop_256x256'] = ret_dct['pt_crop'] * 256 / dsize
|
||||
|
||||
input_image_size = img_rgb.shape[:2]
|
||||
ret_dct['input_image_size'] = input_image_size
|
||||
|
||||
recon_ret = self.landmark_runner.run(img_rgb, pts)
|
||||
lmk = recon_ret['pts']
|
||||
ret_dct['lmk_crop'] = lmk
|
||||
|
||||
return ret_dct
|
||||
|
||||
def get_retargeting_lmk_info(self, driving_rgb_lst):
|
||||
# TODO: implement a tracking-based version
|
||||
driving_lmk_lst = []
|
||||
for driving_image in driving_rgb_lst:
|
||||
ret_dct = self.crop_single_image(driving_image)
|
||||
driving_lmk_lst.append(ret_dct['lmk_crop'])
|
||||
return driving_lmk_lst
|
||||
|
||||
def make_video_clip(self, driving_rgb_lst, output_path, output_fps=30, **kwargs):
|
||||
trajectory = Trajectory()
|
||||
direction = kwargs.get('direction', 'large-small')
|
||||
for idx, driving_image in enumerate(driving_rgb_lst):
|
||||
if idx == 0 or trajectory.start == -1:
|
||||
src_face = self.face_analysis_wrapper.get(
|
||||
driving_image,
|
||||
flag_do_landmark_2d_106=True,
|
||||
direction=direction
|
||||
)
|
||||
if len(src_face) == 0:
|
||||
# No face detected in the driving_image
|
||||
continue
|
||||
elif len(src_face) > 1:
|
||||
log(f'More than one face detected in the driving frame_{idx}, only pick one face by rule {direction}.')
|
||||
src_face = src_face[0]
|
||||
pts = src_face.landmark_2d_106
|
||||
lmk_203 = self.landmark_runner(driving_image, pts)['pts']
|
||||
trajectory.start, trajectory.end = idx, idx
|
||||
else:
|
||||
lmk_203 = self.face_recon_wrapper(driving_image, trajectory.lmk_lst[-1])['pts']
|
||||
trajectory.end = idx
|
||||
|
||||
trajectory.lmk_lst.append(lmk_203)
|
||||
ret_bbox = parse_bbox_from_landmark(lmk_203, scale=self.crop_cfg.globalscale, vy_ratio=elf.crop_cfg.vy_ratio)['bbox']
|
||||
bbox = [ret_bbox[0, 0], ret_bbox[0, 1], ret_bbox[2, 0], ret_bbox[2, 1]] # 4,
|
||||
trajectory.bbox_lst.append(bbox) # bbox
|
||||
trajectory.frame_rgb_lst.append(driving_image)
|
||||
|
||||
global_bbox = average_bbox_lst(trajectory.bbox_lst)
|
||||
for idx, (frame_rgb, lmk) in enumerate(zip(trajectory.frame_rgb_lst, trajectory.lmk_lst)):
|
||||
ret_dct = crop_image_by_bbox(
|
||||
frame_rgb, global_bbox, lmk=lmk,
|
||||
dsize=self.video_crop_cfg.dsize, flag_rot=self.video_crop_cfg.flag_rot, borderValue=self.video_crop_cfg.borderValue
|
||||
)
|
||||
frame_rgb_crop = ret_dct['img_crop']
|
||||
return ret_dct
|
||||
@@ -1,16 +1,30 @@
|
||||
# coding: utf-8
|
||||
|
||||
"""
|
||||
face detectoin and alignment using InsightFace
|
||||
face detection and alignment using InsightFace
|
||||
"""
|
||||
from insightface.utils import transform
|
||||
|
||||
#patch Insightface function to get rid of the annoying warnings
|
||||
def patched_estimate_affine_matrix_3d23d(X, Y):
|
||||
''' Using least-squares solution
|
||||
Args:
|
||||
X: [n, 3]. 3d points(fixed)
|
||||
Y: [n, 3]. corresponding 3d points(moving). Y = PX
|
||||
Returns:
|
||||
P_Affine: (3, 4). Affine camera matrix (the third row is [0, 0, 0, 1]).
|
||||
'''
|
||||
X_homo = np.hstack((X, np.ones([X.shape[0],1]))) # n x 4
|
||||
P = np.linalg.lstsq(X_homo, Y, rcond=None)[0].T # Affine matrix. 3 x 4
|
||||
return P
|
||||
|
||||
transform.estimate_affine_matrix_3d23d = patched_estimate_affine_matrix_3d23d
|
||||
|
||||
import numpy as np
|
||||
from .rprint import rlog as log
|
||||
from insightface.app import FaceAnalysis
|
||||
from insightface.app.common import Face
|
||||
from .timer import Timer
|
||||
|
||||
|
||||
def sort_by_direction(faces, direction: str = 'large-small', face_center=None):
|
||||
if len(faces) <= 0:
|
||||
return faces
|
||||
@@ -76,4 +90,4 @@ class FaceAnalysisDIY(FaceAnalysis):
|
||||
self.get(img_bgr)
|
||||
|
||||
elapse = self.timer.toc()
|
||||
log(f'FaceAnalysisDIY warmup time: {elapse:.3f}s')
|
||||
print(f'FaceAnalysisDIY warmup time: {elapse:.3f}s')
|
||||
|
||||
@@ -0,0 +1,17 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
from pykalman import KalmanFilter
|
||||
|
||||
|
||||
def smooth(x_d_lst, shape, device, observation_variance=3e-6, process_variance=1e-5):
|
||||
x_d_lst_reshape = [x.reshape(-1) for x in x_d_lst]
|
||||
x_d_stacked = np.vstack(x_d_lst_reshape)
|
||||
kf = KalmanFilter(
|
||||
initial_state_mean=x_d_stacked[0],
|
||||
n_dim_obs=x_d_stacked.shape[1],
|
||||
transition_covariance=process_variance * np.eye(x_d_stacked.shape[1]),
|
||||
observation_covariance=observation_variance * np.eye(x_d_stacked.shape[1])
|
||||
)
|
||||
smoothed_state_means, _ = kf.smooth(x_d_stacked)
|
||||
x_d_lst_smooth = [torch.tensor(state_mean.reshape(shape[-2:]), dtype=torch.float32, device=device) for state_mean in smoothed_state_means]
|
||||
return x_d_lst_smooth
|
||||
@@ -4,58 +4,15 @@
|
||||
utility functions and classes to handle feature extraction and model loading
|
||||
"""
|
||||
|
||||
import os
|
||||
import os.path as osp
|
||||
import cv2
|
||||
import torch
|
||||
from collections import OrderedDict
|
||||
|
||||
def suffix(filename):
|
||||
"""a.jpg -> jpg"""
|
||||
pos = filename.rfind(".")
|
||||
if pos == -1:
|
||||
return ""
|
||||
return filename[pos + 1:]
|
||||
|
||||
|
||||
def prefix(filename):
|
||||
"""a.jpg -> a"""
|
||||
pos = filename.rfind(".")
|
||||
if pos == -1:
|
||||
return filename
|
||||
return filename[:pos]
|
||||
|
||||
|
||||
def basename(filename):
|
||||
"""a/b/c.jpg -> c"""
|
||||
return prefix(osp.basename(filename))
|
||||
|
||||
|
||||
def is_video(file_path):
|
||||
if file_path.lower().endswith((".mp4", ".mov", ".avi", ".webm")) or osp.isdir(file_path):
|
||||
return True
|
||||
return False
|
||||
|
||||
def is_template(file_path):
|
||||
if file_path.endswith(".pkl"):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def mkdir(d, log=False):
|
||||
# return self-assined `d`, for one line code
|
||||
if not osp.exists(d):
|
||||
os.makedirs(d, exist_ok=True)
|
||||
if log:
|
||||
print(f"Make dir: {d}")
|
||||
return d
|
||||
|
||||
|
||||
def squeeze_tensor_to_numpy(tensor):
|
||||
out = tensor.data.squeeze(0).cpu().numpy()
|
||||
return out
|
||||
|
||||
|
||||
def dct2cuda(dct: dict, device_id: int):
|
||||
for key in dct:
|
||||
dct[key] = torch.tensor(dct[key]).to(device_id)
|
||||
@@ -95,11 +52,6 @@ def calculate_transformation(config, s_kp_info, t_0_kp_info, t_i_kp_info, R_s, R
|
||||
new_scale = s_kp_info['scale'] * (t_i_kp_info['scale'] / t_0_kp_info['scale'])
|
||||
return new_rotation, new_expression, new_translation, new_scale
|
||||
|
||||
def load_description(fp):
|
||||
with open(fp, 'r', encoding='utf-8') as f:
|
||||
content = f.read()
|
||||
return content
|
||||
|
||||
|
||||
def resize_to_limit(img, max_dim=1280, n=2):
|
||||
h, w = img.shape[:2]
|
||||
|
||||
@@ -1,97 +0,0 @@
|
||||
# coding: utf-8
|
||||
|
||||
import os
|
||||
from glob import glob
|
||||
import os.path as osp
|
||||
import imageio
|
||||
import numpy as np
|
||||
import cv2; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False)
|
||||
|
||||
|
||||
def load_image_rgb(image_path: str):
|
||||
if not osp.exists(image_path):
|
||||
raise FileNotFoundError(f"Image not found: {image_path}")
|
||||
img = cv2.imread(image_path, cv2.IMREAD_COLOR)
|
||||
return cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
|
||||
|
||||
|
||||
def load_driving_info(driving_info):
|
||||
driving_video_ori = []
|
||||
|
||||
def load_images_from_directory(directory):
|
||||
image_paths = sorted(glob(osp.join(directory, '*.png')) + glob(osp.join(directory, '*.jpg')))
|
||||
return [load_image_rgb(im_path) for im_path in image_paths]
|
||||
|
||||
def load_images_from_video(file_path):
|
||||
reader = imageio.get_reader(file_path)
|
||||
return [image for idx, image in enumerate(reader)]
|
||||
|
||||
if osp.isdir(driving_info):
|
||||
driving_video_ori = load_images_from_directory(driving_info)
|
||||
elif osp.isfile(driving_info):
|
||||
driving_video_ori = load_images_from_video(driving_info)
|
||||
|
||||
return driving_video_ori
|
||||
|
||||
|
||||
def contiguous(obj):
|
||||
if not obj.flags.c_contiguous:
|
||||
obj = obj.copy(order="C")
|
||||
return obj
|
||||
|
||||
|
||||
def _resize_to_limit(img: np.ndarray, max_dim=1920, n=2):
|
||||
"""
|
||||
ajust the size of the image so that the maximum dimension does not exceed max_dim, and the width and the height of the image are multiples of n.
|
||||
:param img: the image to be processed.
|
||||
:param max_dim: the maximum dimension constraint.
|
||||
:param n: the number that needs to be multiples of.
|
||||
:return: the adjusted image.
|
||||
"""
|
||||
h, w = img.shape[:2]
|
||||
|
||||
# ajust the size of the image according to the maximum dimension
|
||||
if max_dim > 0 and max(h, w) > max_dim:
|
||||
if h > w:
|
||||
new_h = max_dim
|
||||
new_w = int(w * (max_dim / h))
|
||||
else:
|
||||
new_w = max_dim
|
||||
new_h = int(h * (max_dim / w))
|
||||
img = cv2.resize(img, (new_w, new_h))
|
||||
|
||||
# ensure that the image dimensions are multiples of n
|
||||
n = max(n, 1)
|
||||
new_h = img.shape[0] - (img.shape[0] % n)
|
||||
new_w = img.shape[1] - (img.shape[1] % n)
|
||||
|
||||
if new_h == 0 or new_w == 0:
|
||||
# when the width or height is less than n, no need to process
|
||||
return img
|
||||
|
||||
if new_h != img.shape[0] or new_w != img.shape[1]:
|
||||
img = img[:new_h, :new_w]
|
||||
|
||||
return img
|
||||
|
||||
|
||||
def load_img_online(obj, mode="bgr", **kwargs):
|
||||
max_dim = kwargs.get("max_dim", 1920)
|
||||
n = kwargs.get("n", 2)
|
||||
if isinstance(obj, str):
|
||||
if mode.lower() == "gray":
|
||||
img = cv2.imread(obj, cv2.IMREAD_GRAYSCALE)
|
||||
else:
|
||||
img = cv2.imread(obj, cv2.IMREAD_COLOR)
|
||||
else:
|
||||
img = obj
|
||||
|
||||
# Resize image to satisfy constraints
|
||||
img = _resize_to_limit(img, max_dim=max_dim, n=n)
|
||||
|
||||
if mode.lower() == "bgr":
|
||||
return contiguous(img)
|
||||
elif mode.lower() == "rgb":
|
||||
return contiguous(img[..., ::-1])
|
||||
else:
|
||||
raise Exception(f"Unknown mode {mode}")
|
||||
@@ -1,19 +1,12 @@
|
||||
# coding: utf-8
|
||||
|
||||
import os.path as osp
|
||||
import cv2; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False)
|
||||
import cv2#; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False)
|
||||
import torch
|
||||
import numpy as np
|
||||
import onnxruntime
|
||||
from .timer import Timer
|
||||
from .rprint import rlog
|
||||
from .crop import crop_image, _transform_pts
|
||||
|
||||
|
||||
def make_abs_path(fn):
|
||||
return osp.join(osp.dirname(osp.realpath(__file__)), fn)
|
||||
|
||||
|
||||
def to_ndarray(obj):
|
||||
if isinstance(obj, torch.Tensor):
|
||||
return obj.cpu().numpy()
|
||||
@@ -22,12 +15,11 @@ def to_ndarray(obj):
|
||||
else:
|
||||
return np.array(obj)
|
||||
|
||||
|
||||
class LandmarkRunner(object):
|
||||
"""landmark runner"""
|
||||
def __init__(self, **kwargs):
|
||||
ckpt_path = kwargs.get('ckpt_path')
|
||||
onnx_provider = kwargs.get('onnx_provider', 'cuda') # 默认用cuda
|
||||
onnx_provider = kwargs.get('onnx_provider', 'cuda')
|
||||
device_id = kwargs.get('device_id', 0)
|
||||
self.dsize = kwargs.get('dsize', 224)
|
||||
self.timer = Timer()
|
||||
@@ -40,7 +32,7 @@ class LandmarkRunner(object):
|
||||
)
|
||||
else:
|
||||
opts = onnxruntime.SessionOptions()
|
||||
opts.intra_op_num_threads = 4 # 默认线程数为 4
|
||||
opts.intra_op_num_threads = 4
|
||||
self.session = onnxruntime.InferenceSession(
|
||||
ckpt_path, providers=['CPUExecutionProvider'],
|
||||
sess_options=opts
|
||||
@@ -78,7 +70,6 @@ class LandmarkRunner(object):
|
||||
}
|
||||
|
||||
def warmup(self):
|
||||
# 构造dummy image进行warmup
|
||||
self.timer.tic()
|
||||
|
||||
dummy_image = np.zeros((1, 3, self.dsize, self.dsize), dtype=np.float32)
|
||||
@@ -86,4 +77,4 @@ class LandmarkRunner(object):
|
||||
_ = self._run(dummy_image)
|
||||
|
||||
elapse = self.timer.toc()
|
||||
rlog(f'LandmarkRunner warmup time: {elapse:.3f}s')
|
||||
print(f'LandmarkRunner warmup time: {elapse:.3f}s')
|
||||
|
||||
@@ -1,16 +0,0 @@
|
||||
# coding: utf-8
|
||||
|
||||
"""
|
||||
custom print and log functions
|
||||
"""
|
||||
|
||||
__all__ = ['rprint', 'rlog']
|
||||
|
||||
try:
|
||||
from rich.console import Console
|
||||
console = Console()
|
||||
rprint = console.print
|
||||
rlog = console.log
|
||||
except:
|
||||
rprint = print
|
||||
rlog = print
|
||||
@@ -1,139 +0,0 @@
|
||||
# coding: utf-8
|
||||
|
||||
"""
|
||||
functions for processing video
|
||||
"""
|
||||
|
||||
import os.path as osp
|
||||
import numpy as np
|
||||
import subprocess
|
||||
import imageio
|
||||
import cv2
|
||||
|
||||
from tqdm import tqdm
|
||||
from .helper import prefix
|
||||
from .rprint import rprint as print
|
||||
|
||||
|
||||
def exec_cmd(cmd):
|
||||
subprocess.run(cmd, shell=True, check=True, stdout=subprocess.PIPE, stderr=subprocess.STDOUT)
|
||||
|
||||
|
||||
def images2video(images, wfp, **kwargs):
|
||||
fps = kwargs.get('fps', 30)
|
||||
video_format = kwargs.get('format', 'mp4') # default is mp4 format
|
||||
codec = kwargs.get('codec', 'libx264') # default is libx264 encoding
|
||||
quality = kwargs.get('quality') # video quality
|
||||
pixelformat = kwargs.get('pixelformat', 'yuv420p') # video pixel format
|
||||
image_mode = kwargs.get('image_mode', 'rgb')
|
||||
macro_block_size = kwargs.get('macro_block_size', 2)
|
||||
ffmpeg_params = ['-crf', str(kwargs.get('crf', 18))]
|
||||
|
||||
writer = imageio.get_writer(
|
||||
wfp, fps=fps, format=video_format,
|
||||
codec=codec, quality=quality, ffmpeg_params=ffmpeg_params, pixelformat=pixelformat, macro_block_size=macro_block_size
|
||||
)
|
||||
|
||||
n = len(images)
|
||||
for i in tqdm(range(n), desc='writing', transient=True):
|
||||
if image_mode.lower() == 'bgr':
|
||||
writer.append_data(images[i][..., ::-1])
|
||||
else:
|
||||
writer.append_data(images[i])
|
||||
|
||||
writer.close()
|
||||
|
||||
# print(f':smiley: Dump to {wfp}\n', style="bold green")
|
||||
print(f'Dump to {wfp}\n')
|
||||
|
||||
|
||||
def video2gif(video_fp, fps=30, size=256):
|
||||
if osp.exists(video_fp):
|
||||
d = osp.split(video_fp)[0]
|
||||
fn = prefix(osp.basename(video_fp))
|
||||
palette_wfp = osp.join(d, 'palette.png')
|
||||
gif_wfp = osp.join(d, f'{fn}.gif')
|
||||
# generate the palette
|
||||
cmd = f'ffmpeg -i {video_fp} -vf "fps={fps},scale={size}:-1:flags=lanczos,palettegen" {palette_wfp} -y'
|
||||
exec_cmd(cmd)
|
||||
# use the palette to generate the gif
|
||||
cmd = f'ffmpeg -i {video_fp} -i {palette_wfp} -filter_complex "fps={fps},scale={size}:-1:flags=lanczos[x];[x][1:v]paletteuse" {gif_wfp} -y'
|
||||
exec_cmd(cmd)
|
||||
else:
|
||||
print(f'video_fp: {video_fp} not exists!')
|
||||
|
||||
|
||||
def merge_audio_video(video_fp, audio_fp, wfp):
|
||||
if osp.exists(video_fp) and osp.exists(audio_fp):
|
||||
cmd = f'ffmpeg -i {video_fp} -i {audio_fp} -c:v copy -c:a aac {wfp} -y'
|
||||
exec_cmd(cmd)
|
||||
print(f'merge {video_fp} and {audio_fp} to {wfp}')
|
||||
else:
|
||||
print(f'video_fp: {video_fp} or audio_fp: {audio_fp} not exists!')
|
||||
|
||||
|
||||
def blend(img: np.ndarray, mask: np.ndarray, background_color=(255, 255, 255)):
|
||||
mask_float = mask.astype(np.float32) / 255.
|
||||
background_color = np.array(background_color).reshape([1, 1, 3])
|
||||
bg = np.ones_like(img) * background_color
|
||||
img = np.clip(mask_float * img + (1 - mask_float) * bg, 0, 255).astype(np.uint8)
|
||||
return img
|
||||
|
||||
|
||||
def concat_frames(I_p_lst, driving_rgb_lst, img_rgb):
|
||||
# TODO: add more concat style, e.g., left-down corner driving
|
||||
out_lst = []
|
||||
for idx, _ in tqdm(enumerate(I_p_lst), total=len(I_p_lst), desc='Concatenating result...'):
|
||||
source_image_drived = I_p_lst[idx]
|
||||
image_drive = driving_rgb_lst[idx]
|
||||
|
||||
# resize images to match source_image_drived shape
|
||||
h, w, _ = source_image_drived.shape
|
||||
image_drive_resized = cv2.resize(image_drive, (w, h))
|
||||
img_rgb_resized = cv2.resize(img_rgb, (w, h))
|
||||
|
||||
# concatenate images horizontally
|
||||
frame = np.concatenate((image_drive_resized, img_rgb_resized, source_image_drived), axis=1)
|
||||
out_lst.append(frame)
|
||||
return out_lst
|
||||
|
||||
|
||||
class VideoWriter:
|
||||
def __init__(self, **kwargs):
|
||||
self.fps = kwargs.get('fps', 30)
|
||||
self.wfp = kwargs.get('wfp', 'video.mp4')
|
||||
self.video_format = kwargs.get('format', 'mp4')
|
||||
self.codec = kwargs.get('codec', 'libx264')
|
||||
self.quality = kwargs.get('quality')
|
||||
self.pixelformat = kwargs.get('pixelformat', 'yuv420p')
|
||||
self.image_mode = kwargs.get('image_mode', 'rgb')
|
||||
self.ffmpeg_params = kwargs.get('ffmpeg_params')
|
||||
|
||||
self.writer = imageio.get_writer(
|
||||
self.wfp, fps=self.fps, format=self.video_format,
|
||||
codec=self.codec, quality=self.quality,
|
||||
ffmpeg_params=self.ffmpeg_params, pixelformat=self.pixelformat
|
||||
)
|
||||
|
||||
def write(self, image):
|
||||
if self.image_mode.lower() == 'bgr':
|
||||
self.writer.append_data(image[..., ::-1])
|
||||
else:
|
||||
self.writer.append_data(image)
|
||||
|
||||
def close(self):
|
||||
if self.writer is not None:
|
||||
self.writer.close()
|
||||
|
||||
|
||||
def change_video_fps(input_file, output_file, fps=20, codec='libx264', crf=5):
|
||||
cmd = f"ffmpeg -i {input_file} -c:v {codec} -crf {crf} -r {fps} {output_file} -y"
|
||||
exec_cmd(cmd)
|
||||
|
||||
|
||||
def get_fps(filepath):
|
||||
import ffmpeg
|
||||
probe = ffmpeg.probe(filepath)
|
||||
video_stream = next((stream for stream in probe['streams'] if stream['codec_type'] == 'video'), None)
|
||||
fps = eval(video_stream['avg_frame_rate'])
|
||||
return fps
|
||||
@@ -4,6 +4,9 @@ import yaml
|
||||
import folder_paths
|
||||
import comfy.model_management as mm
|
||||
import comfy.utils
|
||||
import numpy as np
|
||||
import cv2
|
||||
from tqdm import tqdm
|
||||
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
@@ -12,28 +15,35 @@ from .liveportrait.utils.cropper import Cropper
|
||||
from .liveportrait.modules.spade_generator import SPADEDecoder
|
||||
from .liveportrait.modules.warping_network import WarpingNetwork
|
||||
from .liveportrait.modules.motion_extractor import MotionExtractor
|
||||
from .liveportrait.modules.appearance_feature_extractor import AppearanceFeatureExtractor
|
||||
from .liveportrait.modules.stitching_retargeting_network import StitchingRetargetingNetwork
|
||||
from .liveportrait.modules.appearance_feature_extractor import (
|
||||
AppearanceFeatureExtractor,
|
||||
)
|
||||
from .liveportrait.modules.stitching_retargeting_network import (
|
||||
StitchingRetargetingNetwork,
|
||||
)
|
||||
|
||||
import logging
|
||||
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
class InferenceConfig:
|
||||
def __init__(self,
|
||||
mask_crop = None,
|
||||
flag_use_half_precision=True,
|
||||
flag_lip_zero=True,
|
||||
lip_zero_threshold=0.03,
|
||||
flag_eye_retargeting=False,
|
||||
flag_lip_retargeting=False,
|
||||
flag_stitching=True,
|
||||
flag_relative=True,
|
||||
anchor_frame=0,
|
||||
input_shape=(256, 256),
|
||||
flag_write_result=True,
|
||||
flag_pasteback=True,
|
||||
ref_max_shape=1280,
|
||||
ref_shape_n=2,
|
||||
device_id=0,
|
||||
flag_do_crop=True,
|
||||
flag_do_rot=True):
|
||||
def __init__(
|
||||
self,
|
||||
mask_crop=None,
|
||||
flag_use_half_precision=True,
|
||||
flag_lip_zero=True,
|
||||
lip_zero_threshold=0.03,
|
||||
flag_eye_retargeting=False,
|
||||
flag_lip_retargeting=False,
|
||||
flag_stitching=True,
|
||||
flag_relative=True,
|
||||
flag_relative_rotation_only=False,
|
||||
input_shape=(256, 256),
|
||||
flag_pasteback=True,
|
||||
device_id=0,
|
||||
flag_do_crop=True,
|
||||
flag_do_rot=True,
|
||||
):
|
||||
self.flag_use_half_precision = flag_use_half_precision
|
||||
self.flag_lip_zero = flag_lip_zero
|
||||
self.lip_zero_threshold = lip_zero_threshold
|
||||
@@ -41,69 +51,29 @@ class InferenceConfig:
|
||||
self.flag_lip_retargeting = flag_lip_retargeting
|
||||
self.flag_stitching = flag_stitching
|
||||
self.flag_relative = flag_relative
|
||||
self.anchor_frame = anchor_frame
|
||||
self.flag_relative_rotation_only = flag_relative_rotation_only
|
||||
self.input_shape = input_shape
|
||||
self.flag_write_result = flag_write_result
|
||||
self.flag_pasteback = flag_pasteback
|
||||
self.ref_max_shape = ref_max_shape
|
||||
self.ref_shape_n = ref_shape_n
|
||||
self.device_id = device_id
|
||||
self.flag_do_crop = flag_do_crop
|
||||
self.flag_do_rot = flag_do_rot
|
||||
self.mask_crop=mask_crop
|
||||
|
||||
class CropConfig:
|
||||
def __init__(self, dsize=512, scale=2.3, vx_ratio=0, vy_ratio=-0.125):
|
||||
self.dsize = dsize
|
||||
self.scale = scale
|
||||
self.vx_ratio = vx_ratio
|
||||
self.vy_ratio = vy_ratio
|
||||
|
||||
class ArgumentConfig:
|
||||
def __init__(self,
|
||||
device_id=0,
|
||||
flag_lip_zero=True,
|
||||
flag_eye_retargeting=False,
|
||||
flag_lip_retargeting=False,
|
||||
flag_stitching=True,
|
||||
flag_relative=True,
|
||||
flag_pasteback=True,
|
||||
flag_do_crop=True,
|
||||
flag_do_rot=True,
|
||||
dsize=512,
|
||||
scale=2.3,
|
||||
vx_ratio=0,
|
||||
vy_ratio=-0.125,
|
||||
):
|
||||
self.device_id = device_id
|
||||
self.flag_lip_zero = flag_lip_zero
|
||||
self.flag_eye_retargeting = flag_eye_retargeting
|
||||
self.flag_lip_retargeting = flag_lip_retargeting
|
||||
self.flag_stitching = flag_stitching
|
||||
self.flag_relative = flag_relative
|
||||
self.flag_pasteback = flag_pasteback
|
||||
self.flag_do_crop = flag_do_crop
|
||||
self.flag_do_rot = flag_do_rot
|
||||
self.dsize = dsize
|
||||
self.scale = scale
|
||||
self.vx_ratio = vx_ratio
|
||||
self.vy_ratio = vy_ratio
|
||||
self.mask_crop = mask_crop
|
||||
|
||||
class DownloadAndLoadLivePortraitModels:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
},
|
||||
return {
|
||||
"required": {},
|
||||
"optional": {
|
||||
"precision": (
|
||||
"precision": (
|
||||
[
|
||||
'auto',
|
||||
'fp16',
|
||||
'fp32',
|
||||
], {
|
||||
"default": 'auto'
|
||||
}),
|
||||
}
|
||||
"fp16",
|
||||
"fp32",
|
||||
"auto",
|
||||
],
|
||||
{"default": "auto"},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LIVEPORTRAITPIPE",)
|
||||
@@ -111,7 +81,7 @@ class DownloadAndLoadLivePortraitModels:
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "LivePortrait"
|
||||
|
||||
def loadmodel(self, precision='auto'):
|
||||
def loadmodel(self, precision="fp16"):
|
||||
device = mm.get_torch_device()
|
||||
mm.soft_empty_cache()
|
||||
|
||||
@@ -138,86 +108,111 @@ class DownloadAndLoadLivePortraitModels:
|
||||
model_path = os.path.join(download_path)
|
||||
|
||||
if not os.path.exists(model_path):
|
||||
print(f"Downloading model to: {model_path}")
|
||||
log.info(f"Downloading model to: {model_path}")
|
||||
from huggingface_hub import snapshot_download
|
||||
snapshot_download(repo_id="Kijai/LivePortrait_safetensors",
|
||||
local_dir=download_path,
|
||||
local_dir_use_symlinks=False)
|
||||
|
||||
model_config_path = os.path.join(script_directory, 'liveportrait', 'config', 'models.yaml')
|
||||
with open(model_config_path, 'r') as file:
|
||||
snapshot_download(
|
||||
repo_id="Kijai/LivePortrait_safetensors",
|
||||
local_dir=download_path,
|
||||
local_dir_use_symlinks=False,
|
||||
)
|
||||
|
||||
model_config_path = os.path.join(
|
||||
script_directory, "liveportrait", "config", "models.yaml"
|
||||
)
|
||||
with open(model_config_path, "r") as file:
|
||||
model_config = yaml.safe_load(file)
|
||||
|
||||
feature_extractor_path = os.path.join(model_path, 'appearance_feature_extractor.safetensors')
|
||||
motion_extractor_path = os.path.join(model_path, 'motion_extractor.safetensors')
|
||||
warping_module_path = os.path.join(model_path, 'warping_module.safetensors')
|
||||
spade_generator_path = os.path.join(model_path, 'spade_generator.safetensors')
|
||||
stitching_retargeting_path = os.path.join(model_path, 'stitching_retargeting_module.safetensors')
|
||||
|
||||
feature_extractor_path = os.path.join(
|
||||
model_path, "appearance_feature_extractor.safetensors"
|
||||
)
|
||||
motion_extractor_path = os.path.join(model_path, "motion_extractor.safetensors")
|
||||
warping_module_path = os.path.join(model_path, "warping_module.safetensors")
|
||||
spade_generator_path = os.path.join(model_path, "spade_generator.safetensors")
|
||||
stitching_retargeting_path = os.path.join(
|
||||
model_path, "stitching_retargeting_module.safetensors"
|
||||
)
|
||||
|
||||
# init F
|
||||
model_params = model_config['model_params']['appearance_feature_extractor_params']
|
||||
self.appearance_feature_extractor = AppearanceFeatureExtractor(**model_params).to(device)
|
||||
self.appearance_feature_extractor.load_state_dict(comfy.utils.load_torch_file(feature_extractor_path))
|
||||
model_params = model_config["model_params"][
|
||||
"appearance_feature_extractor_params"
|
||||
]
|
||||
self.appearance_feature_extractor = AppearanceFeatureExtractor(
|
||||
**model_params
|
||||
).to(device)
|
||||
self.appearance_feature_extractor.load_state_dict(
|
||||
comfy.utils.load_torch_file(feature_extractor_path)
|
||||
)
|
||||
self.appearance_feature_extractor.eval()
|
||||
print('Load appearance_feature_extractor done.')
|
||||
log.info("Load appearance_feature_extractor done.")
|
||||
pbar.update(1)
|
||||
# init M
|
||||
model_params = model_config['model_params']['motion_extractor_params']
|
||||
model_params = model_config["model_params"]["motion_extractor_params"]
|
||||
self.motion_extractor = MotionExtractor(**model_params).to(device)
|
||||
self.motion_extractor.load_state_dict(comfy.utils.load_torch_file(motion_extractor_path))
|
||||
self.motion_extractor.load_state_dict(
|
||||
comfy.utils.load_torch_file(motion_extractor_path)
|
||||
)
|
||||
self.motion_extractor.eval()
|
||||
print('Load motion_extractor done.')
|
||||
log.info("Load motion_extractor done.")
|
||||
pbar.update(1)
|
||||
# init W
|
||||
model_params = model_config['model_params']['warping_module_params']
|
||||
model_params = model_config["model_params"]["warping_module_params"]
|
||||
self.warping_module = WarpingNetwork(**model_params).to(device)
|
||||
self.warping_module.load_state_dict(comfy.utils.load_torch_file(warping_module_path))
|
||||
self.warping_module.load_state_dict(
|
||||
comfy.utils.load_torch_file(warping_module_path)
|
||||
)
|
||||
self.warping_module.eval()
|
||||
print('Load warping_module done.')
|
||||
log.info("Load warping_module done.")
|
||||
pbar.update(1)
|
||||
# init G
|
||||
model_params = model_config['model_params']['spade_generator_params']
|
||||
model_params = model_config["model_params"]["spade_generator_params"]
|
||||
self.spade_generator = SPADEDecoder(**model_params).to(device)
|
||||
self.spade_generator.load_state_dict(comfy.utils.load_torch_file(spade_generator_path))
|
||||
self.spade_generator.load_state_dict(
|
||||
comfy.utils.load_torch_file(spade_generator_path)
|
||||
)
|
||||
self.spade_generator.eval()
|
||||
print('Load spade_generator done.')
|
||||
log.info("Load spade_generator done.")
|
||||
pbar.update(1)
|
||||
|
||||
def filter_checkpoint_for_model(checkpoint, prefix):
|
||||
"""Filter and adjust the checkpoint dictionary for a specific model based on the prefix."""
|
||||
# Create a new dictionary where keys are adjusted by removing the prefix and the model name
|
||||
filtered_checkpoint = {key.replace(prefix + "_module.", ""): value for key, value in checkpoint.items() if key.startswith(prefix)}
|
||||
filtered_checkpoint = {
|
||||
key.replace(prefix + "_module.", ""): value
|
||||
for key, value in checkpoint.items()
|
||||
if key.startswith(prefix)
|
||||
}
|
||||
return filtered_checkpoint
|
||||
|
||||
config = model_config['model_params']['stitching_retargeting_module_params']
|
||||
config = model_config["model_params"]["stitching_retargeting_module_params"]
|
||||
checkpoint = comfy.utils.load_torch_file(stitching_retargeting_path)
|
||||
|
||||
stitcher_prefix = 'retarget_shoulder'
|
||||
stitcher_prefix = "retarget_shoulder"
|
||||
stitcher_checkpoint = filter_checkpoint_for_model(checkpoint, stitcher_prefix)
|
||||
stitcher = StitchingRetargetingNetwork(**config.get('stitching'))
|
||||
stitcher = StitchingRetargetingNetwork(**config.get("stitching"))
|
||||
stitcher.load_state_dict(stitcher_checkpoint)
|
||||
stitcher = stitcher.to(device)
|
||||
stitcher.eval()
|
||||
|
||||
lip_prefix = 'retarget_mouth'
|
||||
lip_prefix = "retarget_mouth"
|
||||
lip_checkpoint = filter_checkpoint_for_model(checkpoint, lip_prefix)
|
||||
retargetor_lip = StitchingRetargetingNetwork(**config.get('lip'))
|
||||
retargetor_lip = StitchingRetargetingNetwork(**config.get("lip"))
|
||||
retargetor_lip.load_state_dict(lip_checkpoint)
|
||||
retargetor_lip = retargetor_lip.to(device)
|
||||
retargetor_lip.eval()
|
||||
|
||||
eye_prefix = 'retarget_eye'
|
||||
eye_prefix = "retarget_eye"
|
||||
eye_checkpoint = filter_checkpoint_for_model(checkpoint, eye_prefix)
|
||||
retargetor_eye = StitchingRetargetingNetwork(**config.get('eye'))
|
||||
retargetor_eye = StitchingRetargetingNetwork(**config.get("eye"))
|
||||
retargetor_eye.load_state_dict(eye_checkpoint)
|
||||
retargetor_eye = retargetor_eye.to(device)
|
||||
retargetor_eye.eval()
|
||||
print('Load stitching_retargeting_module done.')
|
||||
log.info("Load stitching_retargeting_module done.")
|
||||
|
||||
self.stich_retargeting_module = {
|
||||
'stitching': stitcher,
|
||||
'lip': retargetor_lip,
|
||||
'eye': retargetor_eye
|
||||
"stitching": stitcher,
|
||||
"lip": retargetor_lip,
|
||||
"eye": retargetor_eye,
|
||||
}
|
||||
|
||||
pipeline = LivePortraitPipeline(
|
||||
@@ -227,98 +222,379 @@ class DownloadAndLoadLivePortraitModels:
|
||||
self.spade_generator,
|
||||
self.stich_retargeting_module,
|
||||
InferenceConfig(
|
||||
device_id=device,
|
||||
flag_use_half_precision = True if dtype == 'fp16' else False
|
||||
)
|
||||
device_id=device,
|
||||
flag_use_half_precision=True if precision == "fp16" else False,
|
||||
),
|
||||
)
|
||||
|
||||
return (pipeline,)
|
||||
|
||||
|
||||
class LivePortraitProcess:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
|
||||
"pipeline": ("LIVEPORTRAITPIPE",),
|
||||
"crop_info": ("CROPINFO", {"default": {}}),
|
||||
"source_image": ("IMAGE",),
|
||||
"driving_images": ("IMAGE",),
|
||||
"lip_zero": ("BOOLEAN", {"default": False}),
|
||||
"lip_zero_threshold": ("FLOAT", {"default": 0.03, "min": 0.001, "max": 4.0, "step": 0.001}),
|
||||
"stitching": ("BOOLEAN", {"default": True}),
|
||||
"delta_multiplier": ("FLOAT", {"default": 1.0, "min": -10.0, "max": 10.0, "step": 0.001}),
|
||||
"mismatch_method": (
|
||||
[
|
||||
"constant",
|
||||
"cycle",
|
||||
"mirror",
|
||||
"cut"
|
||||
],
|
||||
{"default": "constant"},
|
||||
),
|
||||
|
||||
"relative_motion_mode": (
|
||||
[
|
||||
"relative",
|
||||
"source_video_smoothed",
|
||||
"relative_rotation_only",
|
||||
"off"
|
||||
],
|
||||
),
|
||||
"driving_smooth_observation_variance": ("FLOAT", {"default": 3e-6, "min": 1e-11, "max": 1e-2, "step": 1e-11}),
|
||||
},
|
||||
|
||||
"optional": {
|
||||
"mask": ("MASK", {"default": None}),
|
||||
"opt_retargeting_info": ("RETARGETINGINFO", {"default": None}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = (
|
||||
"IMAGE",
|
||||
"IMAGE",
|
||||
"MASK",
|
||||
)
|
||||
RETURN_NAMES = (
|
||||
"cropped_images",
|
||||
"full_images",
|
||||
"mask",
|
||||
)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "LivePortrait"
|
||||
|
||||
def process(
|
||||
self,
|
||||
source_image: torch.Tensor,
|
||||
driving_images: torch.Tensor,
|
||||
crop_info: dict,
|
||||
pipeline: LivePortraitPipeline,
|
||||
lip_zero: bool,
|
||||
lip_zero_threshold: float,
|
||||
stitching: bool,
|
||||
relative_motion_mode: str,
|
||||
driving_smooth_observation_variance: float,
|
||||
delta_multiplier: float = 1.0,
|
||||
mismatch_method: str = "constant",
|
||||
mask: torch.Tensor = None,
|
||||
opt_retargeting_info: dict = None,
|
||||
):
|
||||
if driving_images.shape[0] < source_image.shape[0]:
|
||||
raise ValueError("The number of driving images should be larger than the number of source images.")
|
||||
source_np = (source_image * 255).byte().numpy()
|
||||
|
||||
if opt_retargeting_info is not None:
|
||||
pipeline.live_portrait_wrapper.cfg.flag_eye_retargeting = opt_retargeting_info["eye_retargeting"]
|
||||
pipeline.live_portrait_wrapper.cfg.eyes_retargeting_multiplier = (opt_retargeting_info["eyes_retargeting_multiplier"])
|
||||
pipeline.live_portrait_wrapper.cfg.flag_lip_retargeting = opt_retargeting_info["lip_retargeting"]
|
||||
pipeline.live_portrait_wrapper.cfg.lip_retargeting_multiplier = (opt_retargeting_info["lip_retargeting_multiplier"])
|
||||
driving_landmarks = opt_retargeting_info["driving_landmarks"]
|
||||
else:
|
||||
pipeline.live_portrait_wrapper.cfg.flag_eye_retargeting = False
|
||||
pipeline.live_portrait_wrapper.cfg.eyes_retargeting_multiplier = 1.0
|
||||
pipeline.live_portrait_wrapper.cfg.flag_lip_retargeting = False
|
||||
pipeline.live_portrait_wrapper.cfg.lip_retargeting_multiplier = 1.0
|
||||
driving_landmarks = None
|
||||
|
||||
pipeline.live_portrait_wrapper.cfg.flag_stitching = stitching
|
||||
pipeline.live_portrait_wrapper.cfg.flag_lip_zero = lip_zero
|
||||
pipeline.live_portrait_wrapper.cfg.lip_zero_threshold = lip_zero_threshold
|
||||
|
||||
if relative_motion_mode != "off":
|
||||
pipeline.live_portrait_wrapper.cfg.flag_relative = True
|
||||
else:
|
||||
pipeline.live_portrait_wrapper.cfg.flag_relative = False
|
||||
|
||||
if lip_zero and opt_retargeting_info is not None:
|
||||
log.warning("Warning: lip_zero only has an effect with lip or eye retargeting")
|
||||
|
||||
if mask is not None:
|
||||
crop_mask = mask.unsqueeze(-1).expand(-1, -1, -1, 3)
|
||||
pipeline.live_portrait_wrapper.cfg.mask_crop = crop_mask
|
||||
else:
|
||||
log.info("Using default mask template")
|
||||
pipeline.live_portrait_wrapper.cfg.mask_crop = cv2.imread(os.path.join(script_directory, "liveportrait", "utils", "resources", "mask_template.png"), cv2.IMREAD_COLOR)
|
||||
|
||||
driving_images_256 = comfy.utils.common_upscale(driving_images.permute(0, 3, 1, 2), 256, 256, "lanczos", "disabled")
|
||||
if pipeline.live_portrait_wrapper.cfg.flag_use_half_precision:
|
||||
driving_images_256 = driving_images_256.to(torch.float16)
|
||||
|
||||
cropped_out_list = []
|
||||
full_out_list = []
|
||||
|
||||
cropped_out_list, full_out_list, out_mask_list = pipeline.execute(
|
||||
source_np,
|
||||
driving_images_256,
|
||||
crop_info,
|
||||
driving_landmarks,
|
||||
delta_multiplier,
|
||||
relative_motion_mode,
|
||||
driving_smooth_observation_variance,
|
||||
mismatch_method
|
||||
)
|
||||
|
||||
cropped_out_tensors = torch.cat(cropped_out_list, dim=0)
|
||||
|
||||
full_tensors_out = torch.cat(full_out_list, dim=0)
|
||||
full_tensors_out = full_tensors_out.permute(0, 2, 3, 1)
|
||||
|
||||
mask_tensors_out = torch.cat(out_mask_list, dim=0)
|
||||
mask_tensors_out = mask_tensors_out[:, :, :, 0]
|
||||
|
||||
return (
|
||||
cropped_out_tensors.cpu().float(),
|
||||
full_tensors_out.cpu().float(),
|
||||
mask_tensors_out.cpu().float()
|
||||
)
|
||||
|
||||
class LivePortraitLoadCropper:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
|
||||
"onnx_device": (
|
||||
['CPU', 'CUDA', 'ROCM', 'CoreML'], {
|
||||
"default": 'CPU'
|
||||
}),
|
||||
"keep_model_loaded": ("BOOLEAN", {"default": True})
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LPCROPPER",)
|
||||
RETURN_NAMES = ("cropper",)
|
||||
FUNCTION = "crop"
|
||||
CATEGORY = "LivePortrait"
|
||||
|
||||
def crop(self, onnx_device, keep_model_loaded):
|
||||
cropper_init_config = {
|
||||
'keep_model_loaded': keep_model_loaded,
|
||||
'onnx_device': onnx_device
|
||||
}
|
||||
|
||||
if not hasattr(self, 'cropper') or self.cropper is None or self.current_config != cropper_init_config:
|
||||
self.current_config = cropper_init_config
|
||||
self.cropper = Cropper(**cropper_init_config)
|
||||
|
||||
return (self.cropper,)
|
||||
|
||||
class LivePortraitCropper:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"cropper": ("LPCROPPER",),
|
||||
"source_image": ("IMAGE",),
|
||||
"dsize": ("INT", {"default": 512, "min": 64, "max": 2048}),
|
||||
"scale": ("FLOAT", {"default": 2.3, "min": 1.0, "max": 4.0, "step": 0.01}),
|
||||
"vx_ratio": ("FLOAT", {"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.01}),
|
||||
"vy_ratio": ("FLOAT", {"default": -0.125, "min": -1.0, "max": 1.0, "step": 0.01}),
|
||||
"lip_zero": ("BOOLEAN", {"default": True}),
|
||||
"vx_ratio": ("FLOAT", {"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.001}),
|
||||
"vy_ratio": ("FLOAT", {"default": -0.125, "min": -1.0, "max": 1.0, "step": 0.001}),
|
||||
"face_index": ("INT", {"default": 0, "min": 0, "max": 100}),
|
||||
"face_index_order": (
|
||||
[
|
||||
'large-small',
|
||||
'left-right',
|
||||
'right-left',
|
||||
'top-bottom',
|
||||
'bottom-top',
|
||||
'small-large',
|
||||
'distance-from-retarget-face'
|
||||
],
|
||||
),
|
||||
"rotate": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "CROPINFO",)
|
||||
RETURN_NAMES = ("cropped_image", "crop_info",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "LivePortrait"
|
||||
|
||||
def process(self, cropper, source_image, dsize, scale, vx_ratio, vy_ratio, face_index, face_index_order, rotate):
|
||||
source_image_np = (source_image * 255).byte().numpy()
|
||||
|
||||
crop_info_list = []
|
||||
cropped_images_list = []
|
||||
|
||||
pbar = comfy.utils.ProgressBar(len(source_image_np))
|
||||
for i in tqdm(range(len(source_image_np)), desc='Detecting and cropping..', total=len(source_image_np)):
|
||||
crop_info = cropper.crop_single_image(source_image_np[i], dsize, scale, vy_ratio, vx_ratio, face_index, face_index_order, rotate)
|
||||
crop_info_list.append(crop_info)
|
||||
if crop_info:
|
||||
cropped_image = crop_info['img_crop_256x256']
|
||||
else:
|
||||
cropped_image = np.zeros((256, 256, 3), dtype=np.uint8)
|
||||
cropped_images_list.append(cropped_image)
|
||||
|
||||
pbar.update(1)
|
||||
|
||||
cropped_tensors_out = (
|
||||
torch.stack([torch.from_numpy(np_array) for np_array in cropped_images_list])
|
||||
/ 255
|
||||
)
|
||||
|
||||
crop_info_dict = {
|
||||
'crop_info_list': crop_info_list
|
||||
}
|
||||
|
||||
return (cropped_tensors_out, crop_info_dict)
|
||||
|
||||
class LivePortraitRetargeting:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"driving_crop_info": ("CROPINFO", {"default": []}),
|
||||
"eye_retargeting": ("BOOLEAN", {"default": False}),
|
||||
"eyes_retargeting_multiplier": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001}),
|
||||
"lip_retargeting": ("BOOLEAN", {"default": False}),
|
||||
"lip_retargeting_multiplier": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001}),
|
||||
"stitching": ("BOOLEAN", {"default": True}),
|
||||
"relative": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
"optional": {
|
||||
"onnx_device": (
|
||||
[
|
||||
'CPU',
|
||||
'CUDA',
|
||||
], {
|
||||
"default": 'CPU'
|
||||
}),
|
||||
}
|
||||
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "IMAGE",)
|
||||
RETURN_NAMES = ("cropped_images", "full_images",)
|
||||
RETURN_TYPES = ("RETARGETINGINFO",)
|
||||
RETURN_NAMES = ("retargeting_info",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "LivePortrait"
|
||||
|
||||
def process(self, source_image, driving_images, dsize, scale, vx_ratio, vy_ratio, pipeline,
|
||||
lip_zero, eye_retargeting, lip_retargeting, stitching, relative, eyes_retargeting_multiplier, lip_retargeting_multiplier, onnx_device='CUDA'):
|
||||
source_image_np = (source_image * 255).byte().numpy()
|
||||
driving_images_np = (driving_images * 255).byte().numpy()
|
||||
def process(self, driving_crop_info, eye_retargeting, eyes_retargeting_multiplier, lip_retargeting, lip_retargeting_multiplier):
|
||||
|
||||
crop_cfg = CropConfig(
|
||||
dsize = dsize,
|
||||
scale = scale,
|
||||
vx_ratio = vx_ratio,
|
||||
vy_ratio = vy_ratio,
|
||||
)
|
||||
driving_landmarks = []
|
||||
for crop in driving_crop_info["crop_info_list"]:
|
||||
driving_landmarks.append(crop['lmk_crop'])
|
||||
|
||||
retargeting_info = {
|
||||
'eye_retargeting': eye_retargeting,
|
||||
'eyes_retargeting_multiplier': eyes_retargeting_multiplier,
|
||||
'lip_retargeting': lip_retargeting,
|
||||
'lip_retargeting_multiplier': lip_retargeting_multiplier,
|
||||
'driving_landmarks': driving_landmarks
|
||||
}
|
||||
|
||||
return (retargeting_info,)
|
||||
|
||||
|
||||
class KeypointsToImage:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"crop_info": ("CROPINFO", {"default": []}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("keypoints_image",)
|
||||
FUNCTION = "drawkeypoints"
|
||||
CATEGORY = "LivePortrait"
|
||||
|
||||
def drawkeypoints(self, crop_info):
|
||||
height, width = crop_info["crop_info_list"][0]['input_image_size']
|
||||
keypoints_img_list = []
|
||||
pbar = comfy.utils.ProgressBar(len(crop_info))
|
||||
for crop in crop_info["crop_info_list"]:
|
||||
if crop:
|
||||
keypoints = crop['lmk_crop'].copy()
|
||||
# Draw each landmark as a circle
|
||||
blank_image = np.zeros((height, width, 3), dtype=np.uint8) * 255
|
||||
for (x, y) in keypoints:
|
||||
# Ensure the coordinates are within the dimensions of the blank image
|
||||
if 0 <= x < width and 0 <= y < height:
|
||||
cv2.circle(blank_image, (int(x), int(y)), radius=2, color=(0, 0, 255))
|
||||
|
||||
keypoints_image = cv2.cvtColor(blank_image, cv2.COLOR_BGR2RGB)
|
||||
else:
|
||||
keypoints_image = np.zeros((height, width, 3), dtype=np.uint8) * 255
|
||||
keypoints_img_list.append(keypoints_image)
|
||||
pbar.update(1)
|
||||
|
||||
keypoints_img_tensor = (
|
||||
torch.stack([torch.from_numpy(np_array) for np_array in keypoints_img_list]) / 255).float()
|
||||
|
||||
|
||||
return (keypoints_img_tensor,)
|
||||
|
||||
class KeypointScaler:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"crop_info": ("CROPINFO", {"default": {}}),
|
||||
"scale": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001}),
|
||||
"offset_x": ("INT", {"default": 0, "min": -1024, "max": 1024, "step": 1}),
|
||||
"offset_y": ("INT", {"default": 0, "min": -1024, "max": 1024, "step": 1}),
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CROPINFO", "IMAGE",)
|
||||
RETURN_NAMES = ("crop_info", "keypoints_image",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "LivePortrait"
|
||||
|
||||
def process(self, crop_info, offset_x, offset_y, scale):
|
||||
|
||||
keypoints = crop_info['crop_info']['lmk_crop'].copy()
|
||||
|
||||
# Create an offset array
|
||||
# Calculate the centroid of the keypoints
|
||||
centroid = keypoints.mean(axis=0)
|
||||
|
||||
# Translate keypoints to origin by subtracting the centroid
|
||||
translated_keypoints = keypoints - centroid
|
||||
|
||||
# Scale the translated keypoints
|
||||
scaled_keypoints = translated_keypoints * scale
|
||||
|
||||
# Translate scaled keypoints back to original position and then apply the offset
|
||||
final_keypoints = scaled_keypoints + centroid + np.array([offset_x, offset_y])
|
||||
|
||||
crop_info['crop_info']['lmk_crop'] = final_keypoints #fix this
|
||||
|
||||
# Draw each landmark as a circle
|
||||
width, height = 512, 512
|
||||
blank_image = np.zeros((height, width, 3), dtype=np.uint8) * 255
|
||||
for (x, y) in final_keypoints:
|
||||
# Ensure the coordinates are within the dimensions of the blank image
|
||||
if 0 <= x < width and 0 <= y < height:
|
||||
cv2.circle(blank_image, (int(x), int(y)), radius=2, color=(0, 0, 255))
|
||||
|
||||
keypoints_image = cv2.cvtColor(blank_image, cv2.COLOR_BGR2RGB)
|
||||
keypoints_image_tensor = torch.from_numpy(keypoints_image) / 255
|
||||
keypoints_image_tensor = keypoints_image_tensor.unsqueeze(0).cpu().float()
|
||||
|
||||
cropper = Cropper(crop_cfg=crop_cfg, provider=onnx_device)
|
||||
pipeline.cropper = cropper
|
||||
pipeline.live_portrait_wrapper.cfg.flag_eye_retargeting = eye_retargeting
|
||||
pipeline.live_portrait_wrapper.cfg.eyes_retargeting_multiplier = eyes_retargeting_multiplier
|
||||
pipeline.live_portrait_wrapper.cfg.flag_lip_retargeting = lip_retargeting
|
||||
pipeline.live_portrait_wrapper.cfg.lip_retargeting_multiplier = lip_retargeting_multiplier
|
||||
pipeline.live_portrait_wrapper.cfg.flag_stitching = stitching
|
||||
pipeline.live_portrait_wrapper.cfg.flag_relative = relative
|
||||
pipeline.live_portrait_wrapper.cfg.flag_lip_zero = lip_zero
|
||||
|
||||
cropped_out_list = []
|
||||
full_out_list = []
|
||||
for img in source_image_np:
|
||||
cropped_frames, full_frame = pipeline.execute(img, driving_images_np)
|
||||
cropped_tensors = [torch.from_numpy(np_array) for np_array in cropped_frames]
|
||||
cropped_tensors_out = torch.stack(cropped_tensors) / 255
|
||||
cropped_tensors_out = cropped_tensors_out.cpu().float()
|
||||
|
||||
full_tensors = [torch.from_numpy(np_array) for np_array in full_frame]
|
||||
full_tensors_out = torch.stack(full_tensors) / 255
|
||||
full_tensors_out = full_tensors_out.cpu().float()
|
||||
|
||||
cropped_out_list.append(cropped_tensors_out)
|
||||
full_out_list.append(full_tensors_out)
|
||||
|
||||
cropped_tensors_out = torch.cat(cropped_out_list, dim=0)
|
||||
full_tensors_out = torch.cat(full_out_list, dim=0)
|
||||
|
||||
return (cropped_tensors_out, full_tensors_out)
|
||||
return (crop_info, keypoints_image_tensor,)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"DownloadAndLoadLivePortraitModels": DownloadAndLoadLivePortraitModels,
|
||||
"LivePortraitProcess": LivePortraitProcess,
|
||||
"LivePortraitCropper": LivePortraitCropper,
|
||||
"LivePortraitRetargeting": LivePortraitRetargeting,
|
||||
#"KeypointScaler": KeypointScaler,
|
||||
"KeypointsToImage": KeypointsToImage,
|
||||
"LivePortraitLoadCropper": LivePortraitLoadCropper
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"DownloadAndLoadLivePortraitModels": "(Down)Load LivePortraitModels",
|
||||
"LivePortraitProcess": "LivePortraitProcess",
|
||||
"LivePortraitCropper": "LivePortraitCropper",
|
||||
"LivePortraitRetargeting": "LivePortraitRetargeting",
|
||||
#"KeypointScaler": "KeypointScaler",
|
||||
"KeypointsToImage": "LivePortrait KeypointsToImage",
|
||||
"LivePortraitLoadCropper": "LivePortrait LoadCropper"
|
||||
}
|
||||
@@ -1,5 +1,4 @@
|
||||
pyyaml
|
||||
numpy
|
||||
opencv-python
|
||||
rich
|
||||
onnxruntime-gpu
|
||||
Reference in New Issue
Block a user