Compare commits
47
Commits
main
...
kornia_test
| 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": [
|
|
||||||
"fp16"
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"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_stitching: bool = True # we recommend setting it to True!
|
||||||
|
|
||||||
flag_relative: bool = True # whether to use relative pose
|
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
|
input_shape: Tuple[int, int] = (256, 256) # input shape
|
||||||
output_format: Literal['mp4', 'gif'] = 'mp4' # output video format
|
output_format: Literal['mp4', 'gif'] = 'mp4' # output video format
|
||||||
output_fps: int = 30 # fps for output video
|
output_fps: int = 30 # fps for output video
|
||||||
crf: int = 15 # crf 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
|
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
|
mask_crop = None
|
||||||
flag_write_gif: bool = False
|
flag_write_gif: bool = False
|
||||||
size_gif: int = 256
|
|
||||||
ref_max_shape: int = 1280
|
|
||||||
ref_shape_n: int = 2
|
|
||||||
|
|
||||||
device_id: int = 0
|
device_id: int = 0
|
||||||
flag_do_crop: bool = False # whether to crop the reference portrait to the face-cropping space
|
flag_do_crop: bool = False # whether to crop the reference portrait to the face-cropping space
|
||||||
|
|||||||
@@ -4,184 +4,310 @@
|
|||||||
Pipeline of LivePortrait
|
Pipeline of LivePortrait
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import cv2
|
|
||||||
import numpy as np
|
|
||||||
import os.path as osp
|
|
||||||
from rich.progress import track
|
|
||||||
|
|
||||||
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
|
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):
|
import os
|
||||||
return osp.join(osp.dirname(osp.realpath(__file__)), fn)
|
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
|
||||||
|
|
||||||
class LivePortraitPipeline(object):
|
class LivePortraitPipeline(object):
|
||||||
|
def __init__(
|
||||||
def __init__(self, appearance_feature_extractor, motion_extractor, warping_module,
|
self,
|
||||||
spade_generator, stitching_retargeting_module, inference_cfg: InferenceConfig):
|
appearance_feature_extractor,
|
||||||
|
motion_extractor,
|
||||||
|
warping_module,
|
||||||
|
spade_generator,
|
||||||
|
stitching_retargeting_module,
|
||||||
|
inference_cfg: InferenceConfig,
|
||||||
|
):
|
||||||
self.live_portrait_wrapper: LivePortraitWrapper = LivePortraitWrapper(
|
self.live_portrait_wrapper: LivePortraitWrapper = LivePortraitWrapper(
|
||||||
appearance_feature_extractor, motion_extractor, warping_module,
|
appearance_feature_extractor,
|
||||||
spade_generator, stitching_retargeting_module, cfg=inference_cfg)
|
motion_extractor,
|
||||||
|
warping_module,
|
||||||
|
spade_generator,
|
||||||
|
stitching_retargeting_module,
|
||||||
|
cfg=inference_cfg,
|
||||||
|
)
|
||||||
|
|
||||||
def execute(self, img_rgb, driving_images_np):
|
def _get_source_frame(self, source_np, idx, method):
|
||||||
inference_cfg = self.live_portrait_wrapper.cfg # for convenience
|
if source_np.shape[0] == 1:
|
||||||
######## process reference portrait ########
|
return source_np[0]
|
||||||
#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)
|
|
||||||
|
|
||||||
if inference_cfg.flag_lip_zero:
|
if method == "constant":
|
||||||
# let lip-open scalar to be 0 at first
|
return source_np[min(idx, source_np.shape[0] - 1)]
|
||||||
c_d_lip_before_animation = [0.]
|
elif method == "cycle":
|
||||||
combined_lip_ratio_tensor_before_animation = self.live_portrait_wrapper.calc_combined_lip_ratio(c_d_lip_before_animation, source_lmk)
|
return source_np[idx % source_np.shape[0]]
|
||||||
if combined_lip_ratio_tensor_before_animation[0][0] < inference_cfg.lip_zero_threshold:
|
elif method == "mirror":
|
||||||
inference_cfg.flag_lip_zero = False
|
cycle_length = 2 * source_np.shape[0] - 2
|
||||||
else:
|
mirror_idx = idx % cycle_length
|
||||||
lip_delta_before_animation = self.live_portrait_wrapper.retarget_lip(x_s, combined_lip_ratio_tensor_before_animation)
|
if mirror_idx >= source_np.shape[0]:
|
||||||
############################################
|
mirror_idx = cycle_length - mirror_idx
|
||||||
|
return source_np[mirror_idx]
|
||||||
|
|
||||||
######## process driving info ########
|
def execute(
|
||||||
#if is_video(args.driving_info):
|
self, source_np, driving_images, crop_info, driving_landmarks, delta_multiplier, relative_motion_mode, driving_smooth_observation_variance, mismatch_method="constant",
|
||||||
#log(f"Load from video file (mp4 mov avi etc...): {args.driving_info}")
|
):
|
||||||
# TODO: 这里track一下驱动视频 -> 构建模板
|
inference_cfg = self.live_portrait_wrapper.cfg
|
||||||
#driving_rgb_lst = load_driving_info(args.driving_info)
|
device = inference_cfg.device_id
|
||||||
|
|
||||||
driving_rgb_lst = driving_images_np
|
cropped_image_list = []
|
||||||
|
composited_image_list = []
|
||||||
driving_rgb_lst_256 = [cv2.resize(_, (256, 256)) for _ in driving_rgb_lst]
|
out_mask_list = []
|
||||||
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 = []
|
|
||||||
R_d_0, x_d_0_info = None, None
|
R_d_0, x_d_0_info = None, None
|
||||||
pbar = comfy.utils.ProgressBar(n_frames)
|
|
||||||
for i in track(range(n_frames), description='Animating...', total=n_frames):
|
if mismatch_method == "cut":
|
||||||
#if is_video(args.driving_info):
|
total_frames = source_np.shape[0]
|
||||||
# extract kp info by M
|
else:
|
||||||
I_d_i = I_d_lst[i]
|
total_frames = driving_images.shape[0]
|
||||||
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:
|
source_info = []
|
||||||
# # from template
|
source_rot_list = []
|
||||||
# x_d_i_info = template_lst[i]
|
f_s_list = []
|
||||||
# x_d_i_info = dct2cuda(x_d_i_info, inference_cfg.device_id)
|
for i in tqdm(range(source_np.shape[0]), desc='Processing source images...', total=source_np.shape[0]):
|
||||||
# R_d_i = x_d_i_info['R_d']
|
#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:
|
if i == 0:
|
||||||
R_d_0 = R_d_i
|
R_d_0 = R_d
|
||||||
x_d_0_info = x_d_i_info
|
x_d_0_info = x_d_info
|
||||||
|
|
||||||
if inference_cfg.flag_relative:
|
if relative_motion_mode == "relative":
|
||||||
R_new = (R_d_i @ R_d_0.permute(0, 2, 1)) @ R_s
|
R_new = (R_d @ 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'])
|
delta_new = x_s_info["exp"] + (x_d_info["exp"] - x_d_0_info["exp"])
|
||||||
scale_new = x_s_info['scale'] * (x_d_i_info['scale'] / x_d_0_info['scale'])
|
scale_new = x_s_info["scale"] * (x_d_info["scale"] / x_d_0_info["scale"])
|
||||||
t_new = x_s_info['t'] + (x_d_i_info['t'] - x_d_0_info['t'])
|
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:
|
else:
|
||||||
R_new = R_d_i
|
R_new = R_d
|
||||||
delta_new = x_d_i_info['exp']
|
delta_new = x_s_info['exp']
|
||||||
scale_new = x_s_info['scale']
|
scale_new = x_s_info["scale"]
|
||||||
t_new = x_d_i_info['t']
|
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
|
x_d_i_new = scale_new * (x_c_s @ R_new + delta_new) + t_new
|
||||||
|
if (
|
||||||
# Algorithm 1:
|
not inference_cfg.flag_stitching
|
||||||
if not inference_cfg.flag_stitching and not inference_cfg.flag_eye_retargeting and not inference_cfg.flag_lip_retargeting:
|
and not inference_cfg.flag_eye_retargeting
|
||||||
|
and not inference_cfg.flag_lip_retargeting
|
||||||
|
):
|
||||||
# without stitching or retargeting
|
# without stitching or retargeting
|
||||||
if inference_cfg.flag_lip_zero:
|
if inference_cfg.flag_lip_zero:
|
||||||
x_d_i_new += lip_delta_before_animation.reshape(-1, x_s.shape[1], 3)
|
x_d_i_new += lip_delta_before_animation.reshape(-1, x_s.shape[1], 3)
|
||||||
else:
|
else:
|
||||||
pass
|
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
|
# with stitching and without retargeting
|
||||||
if inference_cfg.flag_lip_zero:
|
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:
|
else:
|
||||||
x_d_i_new = self.live_portrait_wrapper.stitching(x_s, x_d_i_new)
|
x_d_i_new = self.live_portrait_wrapper.stitching(x_s, x_d_i_new)
|
||||||
else:
|
else:
|
||||||
eyes_delta, lip_delta = None, None
|
eyes_delta, lip_delta = None, None
|
||||||
if inference_cfg.flag_eye_retargeting:
|
if inference_cfg.flag_eye_retargeting:
|
||||||
c_d_eyes_i = input_eye_ratio_lst[i]
|
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 = combined_eye_ratio_tensor * inference_cfg.eyes_retargeting_multiplier
|
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,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:
|
if inference_cfg.flag_lip_retargeting:
|
||||||
c_d_lip_i = input_lip_ratio_lst[i]
|
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 = combined_lip_ratio_tensor * inference_cfg.lip_retargeting_multiplier
|
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,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
|
if inference_cfg.flag_relative: # use x_s
|
||||||
x_d_i_new = x_s + \
|
x_d_i_new = (
|
||||||
(eyes_delta.reshape(-1, x_s.shape[1], 3) if eyes_delta is not None else 0) + \
|
x_s
|
||||||
(lip_delta.reshape(-1, x_s.shape[1], 3) if lip_delta is not None else 0)
|
+ (
|
||||||
|
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
|
else: # use x_d,i
|
||||||
x_d_i_new = 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) + \
|
x_d_i_new
|
||||||
(lip_delta.reshape(-1, x_s.shape[1], 3) if lip_delta is not None else 0)
|
+ (
|
||||||
|
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:
|
if inference_cfg.flag_stitching:
|
||||||
x_d_i_new = self.live_portrait_wrapper.stitching(x_s, x_d_i_new)
|
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)
|
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)
|
pbar.update(1)
|
||||||
|
|
||||||
#if inference_cfg.flag_pasteback:
|
return cropped_image_list, composited_image_list, out_mask_list
|
||||||
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
|
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ from .utils.retargeting_utils import compute_eye_delta, compute_lip_delta
|
|||||||
from .utils.camera import headpose_pred_to_degree, get_rotation_matrix
|
from .utils.camera import headpose_pred_to_degree, get_rotation_matrix
|
||||||
from .utils.retargeting_utils import calc_eye_close_ratio, calc_lip_close_ratio
|
from .utils.retargeting_utils import calc_eye_close_ratio, calc_lip_close_ratio
|
||||||
from .config.inference_config import InferenceConfig
|
from .config.inference_config import InferenceConfig
|
||||||
|
from contextlib import nullcontext
|
||||||
|
|
||||||
from comfy.model_management import get_autocast_device
|
from comfy.model_management import get_autocast_device
|
||||||
|
|
||||||
@@ -31,11 +32,6 @@ class LivePortraitWrapper(object):
|
|||||||
self.device_id = cfg.device_id
|
self.device_id = cfg.device_id
|
||||||
self.timer = Timer()
|
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:
|
def prepare_source(self, img: np.ndarray) -> torch.Tensor:
|
||||||
""" construct the input as standard
|
""" construct the input as standard
|
||||||
img: HxWx3, uint8, 256x256
|
img: HxWx3, uint8, 256x256
|
||||||
@@ -57,31 +53,12 @@ class LivePortraitWrapper(object):
|
|||||||
x = x.to(self.device_id)
|
x = x.to(self.device_id)
|
||||||
return x
|
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:
|
def extract_feature_3d(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
""" get the appearance feature of the image by F
|
""" get the appearance feature of the image by F
|
||||||
x: Bx3xHxW, normalized to 0~1
|
x: Bx3xHxW, normalized to 0~1
|
||||||
"""
|
"""
|
||||||
with torch.no_grad():
|
with torch.autocast(get_autocast_device(self.device_id), dtype=torch.float16) if self.cfg.flag_use_half_precision else nullcontext():
|
||||||
with torch.autocast(device_type=get_autocast_device(self.device_id), dtype=torch.float16, enabled=self.cfg.flag_use_half_precision):
|
feature_3d = self.appearance_feature_extractor(x)
|
||||||
feature_3d = self.appearance_feature_extractor(x)
|
|
||||||
|
|
||||||
return feature_3d.float()
|
return feature_3d.float()
|
||||||
|
|
||||||
@@ -91,15 +68,14 @@ class LivePortraitWrapper(object):
|
|||||||
flag_refine_info: whether to trandform the pose to degrees and the dimention of the reshape
|
flag_refine_info: whether to trandform the pose to degrees and the dimention of the reshape
|
||||||
return: A dict contains keys: 'pitch', 'yaw', 'roll', 't', 'exp', 'scale', 'kp'
|
return: A dict contains keys: 'pitch', 'yaw', 'roll', 't', 'exp', 'scale', 'kp'
|
||||||
"""
|
"""
|
||||||
with torch.no_grad():
|
with torch.autocast(get_autocast_device(self.device_id), dtype=torch.float16) if self.cfg.flag_use_half_precision else nullcontext():
|
||||||
with torch.autocast(device_type=get_autocast_device(self.device_id), dtype=torch.float16, enabled=self.cfg.flag_use_half_precision):
|
kp_info = self.motion_extractor(x)
|
||||||
kp_info = self.motion_extractor(x)
|
|
||||||
|
|
||||||
if self.cfg.flag_use_half_precision:
|
if self.cfg.flag_use_half_precision:
|
||||||
# float the dict
|
# float the dict
|
||||||
for k, v in kp_info.items():
|
for k, v in kp_info.items():
|
||||||
if isinstance(v, torch.Tensor):
|
if isinstance(v, torch.Tensor):
|
||||||
kp_info[k] = v.float()
|
kp_info[k] = v.float()
|
||||||
|
|
||||||
flag_refine_info: bool = kwargs.get('flag_refine_info', True)
|
flag_refine_info: bool = kwargs.get('flag_refine_info', True)
|
||||||
if flag_refine_info:
|
if flag_refine_info:
|
||||||
@@ -265,19 +241,18 @@ class LivePortraitWrapper(object):
|
|||||||
kp_source: BxNx3
|
kp_source: BxNx3
|
||||||
kp_driving: 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.no_grad():
|
with torch.autocast(get_autocast_device(self.device_id), dtype=torch.float16) if self.cfg.flag_use_half_precision else nullcontext():
|
||||||
with torch.autocast(device_type=get_autocast_device(self.device_id), dtype=torch.float16, enabled=self.cfg.flag_use_half_precision):
|
# get decoder input
|
||||||
# get decoder input
|
ret_dct = self.warping_module(feature_3d, kp_source=kp_source, kp_driving=kp_driving)
|
||||||
ret_dct = self.warping_module(feature_3d, kp_source=kp_source, kp_driving=kp_driving)
|
# decode
|
||||||
# decode
|
ret_dct['out'] = self.spade_generator(feature=ret_dct['out'])
|
||||||
ret_dct['out'] = self.spade_generator(feature=ret_dct['out'])
|
|
||||||
|
|
||||||
# float the dict
|
# float the dict
|
||||||
if self.cfg.flag_use_half_precision:
|
if self.cfg.flag_use_half_precision:
|
||||||
for k, v in ret_dct.items():
|
for k, v in ret_dct.items():
|
||||||
if isinstance(v, torch.Tensor):
|
if isinstance(v, torch.Tensor):
|
||||||
ret_dct[k] = v.float()
|
ret_dct[k] = v.float()
|
||||||
|
|
||||||
return ret_dct
|
return ret_dct
|
||||||
|
|
||||||
@@ -303,17 +278,19 @@ class LivePortraitWrapper(object):
|
|||||||
|
|
||||||
def calc_combined_eye_ratio(self, input_eye_ratio, source_lmk):
|
def calc_combined_eye_ratio(self, input_eye_ratio, source_lmk):
|
||||||
eye_close_ratio = calc_eye_close_ratio(source_lmk[None])
|
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)
|
eye_close_ratios_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)
|
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]
|
# [c_s,eyes, c_d,eyes,i]
|
||||||
combined_eye_ratio_tensor = torch.cat([eye_close_ratio_tensor, input_eye_ratio_tensor], dim=1)
|
combined_eye_ratios_tensor = torch.cat([eye_close_ratios_tensor, input_eye_ratio_tensor], dim=1)
|
||||||
return combined_eye_ratio_tensor
|
return combined_eye_ratios_tensor
|
||||||
|
|
||||||
def calc_combined_lip_ratio(self, input_lip_ratio, source_lmk):
|
def calc_combined_lip_ratio(self, input_lip_ratio, source_lmk):
|
||||||
lip_close_ratio = calc_lip_close_ratio(source_lmk[None])
|
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)
|
lip_close_ratio_tensor = torch.from_numpy(lip_close_ratio).float().to(self.device_id)
|
||||||
# [c_s,lip, c_d,lip,i]
|
# [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]:
|
if input_lip_ratio_tensor.shape != [1, 1]:
|
||||||
input_lip_ratio_tensor = input_lip_ratio_tensor.reshape(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)
|
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 rich.progress import track
|
|
||||||
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 track(range(n_frames), description='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
|
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
|
import numpy as np
|
||||||
from .rprint import rprint as print
|
|
||||||
from math import sin, cos, acos, degrees
|
from math import sin, cos, acos, degrees
|
||||||
|
|
||||||
DTYPE = np.float32
|
DTYPE = np.float32
|
||||||
CV2_INTERP = cv2.INTER_LINEAR
|
CV2_INTERP = cv2.INTER_LINEAR
|
||||||
|
import comfy.model_management as mm
|
||||||
|
|
||||||
def _transform_img(img, M, dsize, flags=CV2_INTERP, borderMode=None):
|
def _transform_img(img, M, dsize, flags=CV2_INTERP, borderMode=None):
|
||||||
""" conduct similarity or affine transformation to the image, do not do border operation!
|
""" 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:
|
else:
|
||||||
return cv2.warpAffine(img, M[:2, :], dsize=_dsize, flags=flags)
|
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):
|
def _transform_pts(pts, M):
|
||||||
""" conduct similarity or affine transformation to the pts
|
""" 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)
|
dsize = kwargs.get('dsize', 224)
|
||||||
scale = kwargs.get('scale', 1.5) # 1.5 | 1.6
|
scale = kwargs.get('scale', 1.5) # 1.5 | 1.6
|
||||||
vy_ratio = kwargs.get('vy_ratio', -0.1) # -0.0625 | -0.1
|
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(
|
M_INV, _ = _estimate_similar_transform_from_pts(
|
||||||
pts,
|
pts,
|
||||||
dsize=dsize,
|
dsize=dsize,
|
||||||
scale=scale,
|
scale=scale,
|
||||||
vy_ratio=vy_ratio,
|
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:
|
if img is None:
|
||||||
|
|||||||
@@ -1,32 +1,22 @@
|
|||||||
# coding: utf-8
|
# coding: utf-8
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import os.path as osp
|
|
||||||
from typing import List, Union, Tuple
|
from typing import List, Union, Tuple
|
||||||
from dataclasses import dataclass, field
|
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 .landmark_runner import LandmarkRunner
|
||||||
from .face_analysis_diy import FaceAnalysisDIY
|
from .face_analysis_diy import FaceAnalysisDIY
|
||||||
#from .helper import prefix
|
from .crop import crop_image
|
||||||
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
|
|
||||||
|
|
||||||
import folder_paths
|
import folder_paths
|
||||||
import os
|
import os
|
||||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
|
||||||
def make_abs_path(fn):
|
|
||||||
return osp.join(osp.dirname(osp.realpath(__file__)), fn)
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class Trajectory:
|
class Trajectory:
|
||||||
start: int = -1 # 起始帧 闭区间
|
start: int = -1
|
||||||
end: int = -1 # 结束帧 闭区间
|
end: int = -1
|
||||||
lmk_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # lmk list
|
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
|
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
|
frame_rgb_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # frame list
|
||||||
@@ -34,10 +24,10 @@ class Trajectory:
|
|||||||
|
|
||||||
|
|
||||||
class Cropper(object):
|
class Cropper(object):
|
||||||
def __init__(self, provider, **kwargs) -> None:
|
def __init__(self, **kwargs) -> None:
|
||||||
device_id = kwargs.get('device_id', 0)
|
device_id = kwargs.get('device_id', 0)
|
||||||
|
provider = kwargs.get('onnx_device', 'CPU')
|
||||||
self.landmark_runner = LandmarkRunner(
|
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'),
|
ckpt_path=os.path.join(folder_paths.models_dir, 'liveportrait', 'landmark.onnx'),
|
||||||
onnx_provider=provider,
|
onnx_provider=provider,
|
||||||
device_id=device_id
|
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.prepare(ctx_id=device_id, det_size=(512, 512))
|
||||||
self.face_analysis_wrapper.warmup()
|
self.face_analysis_wrapper.warmup()
|
||||||
|
|
||||||
self.crop_cfg = kwargs.get('crop_cfg', None)
|
def crop_single_image(self, img_rgb, dsize, scale, vy_ratio, vx_ratio, face_index, face_index_order, rotate):
|
||||||
|
direction = face_index_order
|
||||||
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
|
|
||||||
|
|
||||||
src_face = self.face_analysis_wrapper.get(
|
src_face = self.face_analysis_wrapper.get(
|
||||||
img_rgb,
|
img_rgb,
|
||||||
@@ -75,73 +52,34 @@ class Cropper(object):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if len(src_face) == 0:
|
if len(src_face) == 0:
|
||||||
log('No face detected in the source image.')
|
ret_dct = {}
|
||||||
raise Exception("No face detected in the source image!")
|
return ret_dct
|
||||||
elif len(src_face) > 1:
|
#raise Exception("No face detected in the source image!")
|
||||||
log(f'More than one face detected in the image, only pick one face by rule {direction}.')
|
#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
|
pts = src_face.landmark_2d_106
|
||||||
|
|
||||||
# crop the face
|
# crop the face
|
||||||
ret_dct = crop_image(
|
ret_dct = crop_image(
|
||||||
img_rgb, # ndarray
|
img_rgb, # ndarray
|
||||||
pts, # 106x2 or Nx2
|
pts, # 106x2 or Nx2
|
||||||
dsize=kwargs.get('dsize', 512),
|
dsize=dsize,
|
||||||
scale=kwargs.get('scale', 2.3),
|
scale=scale,
|
||||||
vy_ratio=kwargs.get('vy_ratio', -0.15),
|
vy_ratio=vy_ratio,
|
||||||
|
vx_ratio=vx_ratio,
|
||||||
|
rotate=rotate
|
||||||
)
|
)
|
||||||
# update a 256x256 version for network input or else
|
# 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['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)
|
recon_ret = self.landmark_runner.run(img_rgb, pts)
|
||||||
lmk = recon_ret['pts']
|
lmk = recon_ret['pts']
|
||||||
ret_dct['lmk_crop'] = lmk
|
ret_dct['lmk_crop'] = lmk
|
||||||
|
|
||||||
return ret_dct
|
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']
|
|
||||||
@@ -1,16 +1,30 @@
|
|||||||
# coding: utf-8
|
# 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
|
import numpy as np
|
||||||
from .rprint import rlog as log
|
|
||||||
from insightface.app import FaceAnalysis
|
from insightface.app import FaceAnalysis
|
||||||
from insightface.app.common import Face
|
from insightface.app.common import Face
|
||||||
from .timer import Timer
|
from .timer import Timer
|
||||||
|
|
||||||
|
|
||||||
def sort_by_direction(faces, direction: str = 'large-small', face_center=None):
|
def sort_by_direction(faces, direction: str = 'large-small', face_center=None):
|
||||||
if len(faces) <= 0:
|
if len(faces) <= 0:
|
||||||
return faces
|
return faces
|
||||||
@@ -76,4 +90,4 @@ class FaceAnalysisDIY(FaceAnalysis):
|
|||||||
self.get(img_bgr)
|
self.get(img_bgr)
|
||||||
|
|
||||||
elapse = self.timer.toc()
|
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
|
utility functions and classes to handle feature extraction and model loading
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import os
|
|
||||||
import os.path as osp
|
import os.path as osp
|
||||||
import cv2
|
import cv2
|
||||||
import torch
|
import torch
|
||||||
from collections import OrderedDict
|
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):
|
def squeeze_tensor_to_numpy(tensor):
|
||||||
out = tensor.data.squeeze(0).cpu().numpy()
|
out = tensor.data.squeeze(0).cpu().numpy()
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
def dct2cuda(dct: dict, device_id: int):
|
def dct2cuda(dct: dict, device_id: int):
|
||||||
for key in dct:
|
for key in dct:
|
||||||
dct[key] = torch.tensor(dct[key]).to(device_id)
|
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'])
|
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
|
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):
|
def resize_to_limit(img, max_dim=1280, n=2):
|
||||||
h, w = img.shape[: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
|
# 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 torch
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import onnxruntime
|
import onnxruntime
|
||||||
from .timer import Timer
|
from .timer import Timer
|
||||||
from .rprint import rlog
|
|
||||||
from .crop import crop_image, _transform_pts
|
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):
|
def to_ndarray(obj):
|
||||||
if isinstance(obj, torch.Tensor):
|
if isinstance(obj, torch.Tensor):
|
||||||
return obj.cpu().numpy()
|
return obj.cpu().numpy()
|
||||||
@@ -22,12 +15,11 @@ def to_ndarray(obj):
|
|||||||
else:
|
else:
|
||||||
return np.array(obj)
|
return np.array(obj)
|
||||||
|
|
||||||
|
|
||||||
class LandmarkRunner(object):
|
class LandmarkRunner(object):
|
||||||
"""landmark runner"""
|
"""landmark runner"""
|
||||||
def __init__(self, **kwargs):
|
def __init__(self, **kwargs):
|
||||||
ckpt_path = kwargs.get('ckpt_path')
|
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)
|
device_id = kwargs.get('device_id', 0)
|
||||||
self.dsize = kwargs.get('dsize', 224)
|
self.dsize = kwargs.get('dsize', 224)
|
||||||
self.timer = Timer()
|
self.timer = Timer()
|
||||||
@@ -40,7 +32,7 @@ class LandmarkRunner(object):
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
opts = onnxruntime.SessionOptions()
|
opts = onnxruntime.SessionOptions()
|
||||||
opts.intra_op_num_threads = 4 # 默认线程数为 4
|
opts.intra_op_num_threads = 4
|
||||||
self.session = onnxruntime.InferenceSession(
|
self.session = onnxruntime.InferenceSession(
|
||||||
ckpt_path, providers=['CPUExecutionProvider'],
|
ckpt_path, providers=['CPUExecutionProvider'],
|
||||||
sess_options=opts
|
sess_options=opts
|
||||||
@@ -78,7 +70,6 @@ class LandmarkRunner(object):
|
|||||||
}
|
}
|
||||||
|
|
||||||
def warmup(self):
|
def warmup(self):
|
||||||
# 构造dummy image进行warmup
|
|
||||||
self.timer.tic()
|
self.timer.tic()
|
||||||
|
|
||||||
dummy_image = np.zeros((1, 3, self.dsize, self.dsize), dtype=np.float32)
|
dummy_image = np.zeros((1, 3, self.dsize, self.dsize), dtype=np.float32)
|
||||||
@@ -86,4 +77,4 @@ class LandmarkRunner(object):
|
|||||||
_ = self._run(dummy_image)
|
_ = self._run(dummy_image)
|
||||||
|
|
||||||
elapse = self.timer.toc()
|
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 rich.progress import track
|
|
||||||
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 track(range(n), description='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 track(enumerate(I_p_lst), total=len(I_p_lst), description='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,37 +4,46 @@ import yaml
|
|||||||
import folder_paths
|
import folder_paths
|
||||||
import comfy.model_management as mm
|
import comfy.model_management as mm
|
||||||
import comfy.utils
|
import comfy.utils
|
||||||
|
import numpy as np
|
||||||
|
import cv2
|
||||||
|
from tqdm import tqdm
|
||||||
|
|
||||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
|
||||||
from .liveportrait.config.argument_config import ArgumentConfig
|
|
||||||
from .liveportrait.live_portrait_pipeline import LivePortraitPipeline
|
from .liveportrait.live_portrait_pipeline import LivePortraitPipeline
|
||||||
from .liveportrait.utils.cropper import Cropper
|
from .liveportrait.utils.cropper import Cropper
|
||||||
from .liveportrait.modules.spade_generator import SPADEDecoder
|
from .liveportrait.modules.spade_generator import SPADEDecoder
|
||||||
from .liveportrait.modules.warping_network import WarpingNetwork
|
from .liveportrait.modules.warping_network import WarpingNetwork
|
||||||
from .liveportrait.modules.motion_extractor import MotionExtractor
|
from .liveportrait.modules.motion_extractor import MotionExtractor
|
||||||
from .liveportrait.modules.appearance_feature_extractor import AppearanceFeatureExtractor
|
from .liveportrait.modules.appearance_feature_extractor import (
|
||||||
from .liveportrait.modules.stitching_retargeting_network import StitchingRetargetingNetwork
|
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:
|
class InferenceConfig:
|
||||||
def __init__(self,
|
def __init__(
|
||||||
mask_crop = None,
|
self,
|
||||||
flag_use_half_precision=True,
|
mask_crop=None,
|
||||||
flag_lip_zero=True,
|
flag_use_half_precision=True,
|
||||||
lip_zero_threshold=0.03,
|
flag_lip_zero=True,
|
||||||
flag_eye_retargeting=False,
|
lip_zero_threshold=0.03,
|
||||||
flag_lip_retargeting=False,
|
flag_eye_retargeting=False,
|
||||||
flag_stitching=True,
|
flag_lip_retargeting=False,
|
||||||
flag_relative=True,
|
flag_stitching=True,
|
||||||
anchor_frame=0,
|
flag_relative=True,
|
||||||
input_shape=(256, 256),
|
flag_relative_rotation_only=False,
|
||||||
flag_write_result=True,
|
input_shape=(256, 256),
|
||||||
flag_pasteback=True,
|
flag_pasteback=True,
|
||||||
ref_max_shape=1280,
|
device_id=0,
|
||||||
ref_shape_n=2,
|
flag_do_crop=True,
|
||||||
device_id=0,
|
flag_do_rot=True,
|
||||||
flag_do_crop=True,
|
):
|
||||||
flag_do_rot=True):
|
|
||||||
self.flag_use_half_precision = flag_use_half_precision
|
self.flag_use_half_precision = flag_use_half_precision
|
||||||
self.flag_lip_zero = flag_lip_zero
|
self.flag_lip_zero = flag_lip_zero
|
||||||
self.lip_zero_threshold = lip_zero_threshold
|
self.lip_zero_threshold = lip_zero_threshold
|
||||||
@@ -42,68 +51,29 @@ class InferenceConfig:
|
|||||||
self.flag_lip_retargeting = flag_lip_retargeting
|
self.flag_lip_retargeting = flag_lip_retargeting
|
||||||
self.flag_stitching = flag_stitching
|
self.flag_stitching = flag_stitching
|
||||||
self.flag_relative = flag_relative
|
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.input_shape = input_shape
|
||||||
self.flag_write_result = flag_write_result
|
|
||||||
self.flag_pasteback = flag_pasteback
|
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.device_id = device_id
|
||||||
self.flag_do_crop = flag_do_crop
|
self.flag_do_crop = flag_do_crop
|
||||||
self.flag_do_rot = flag_do_rot
|
self.flag_do_rot = flag_do_rot
|
||||||
self.mask_crop=mask_crop
|
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
|
|
||||||
|
|
||||||
class DownloadAndLoadLivePortraitModels:
|
class DownloadAndLoadLivePortraitModels:
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
return {"required": {
|
return {
|
||||||
},
|
"required": {},
|
||||||
"optional": {
|
"optional": {
|
||||||
"precision": (
|
"precision": (
|
||||||
[
|
[
|
||||||
'fp16',
|
"fp16",
|
||||||
'fp32',
|
"fp32",
|
||||||
], {
|
"auto",
|
||||||
"default": 'fp16'
|
],
|
||||||
}),
|
{"default": "auto"},
|
||||||
}
|
),
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("LIVEPORTRAITPIPE",)
|
RETURN_TYPES = ("LIVEPORTRAITPIPE",)
|
||||||
@@ -111,96 +81,138 @@ class DownloadAndLoadLivePortraitModels:
|
|||||||
FUNCTION = "loadmodel"
|
FUNCTION = "loadmodel"
|
||||||
CATEGORY = "LivePortrait"
|
CATEGORY = "LivePortrait"
|
||||||
|
|
||||||
def loadmodel(self, precision='fp16'):
|
def loadmodel(self, precision="fp16"):
|
||||||
device = mm.get_torch_device()
|
device = mm.get_torch_device()
|
||||||
mm.soft_empty_cache()
|
mm.soft_empty_cache()
|
||||||
|
|
||||||
|
if precision == 'auto':
|
||||||
|
try:
|
||||||
|
if mm.is_device_mps(device):
|
||||||
|
print("LivePortrait using fp32 for MPS")
|
||||||
|
dtype = 'fp32'
|
||||||
|
elif mm.should_use_fp16():
|
||||||
|
print("LivePortrait using fp16")
|
||||||
|
dtype = 'fp16'
|
||||||
|
else:
|
||||||
|
print("LivePortrait using fp32")
|
||||||
|
dtype = 'fp32'
|
||||||
|
except:
|
||||||
|
raise AttributeError("ComfyUI version too old, can't autodetect properly. Set your dtypes manually.")
|
||||||
|
else:
|
||||||
|
dtype = precision
|
||||||
|
print(f"LivePortrait using {dtype}")
|
||||||
|
|
||||||
pbar = comfy.utils.ProgressBar(3)
|
pbar = comfy.utils.ProgressBar(3)
|
||||||
|
|
||||||
download_path = os.path.join(folder_paths.models_dir, "liveportrait")
|
download_path = os.path.join(folder_paths.models_dir, "liveportrait")
|
||||||
model_path = os.path.join(download_path)
|
model_path = os.path.join(download_path)
|
||||||
|
|
||||||
if not os.path.exists(model_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
|
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')
|
snapshot_download(
|
||||||
with open(model_config_path, 'r') as file:
|
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)
|
model_config = yaml.safe_load(file)
|
||||||
|
|
||||||
feature_extractor_path = os.path.join(model_path, 'appearance_feature_extractor.safetensors')
|
feature_extractor_path = os.path.join(
|
||||||
motion_extractor_path = os.path.join(model_path, 'motion_extractor.safetensors')
|
model_path, "appearance_feature_extractor.safetensors"
|
||||||
warping_module_path = os.path.join(model_path, 'warping_module.safetensors')
|
)
|
||||||
spade_generator_path = os.path.join(model_path, 'spade_generator.safetensors')
|
motion_extractor_path = os.path.join(model_path, "motion_extractor.safetensors")
|
||||||
stitching_retargeting_path = os.path.join(model_path, 'stitching_retargeting_module.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
|
# init F
|
||||||
model_params = model_config['model_params']['appearance_feature_extractor_params']
|
model_params = model_config["model_params"][
|
||||||
self.appearance_feature_extractor = AppearanceFeatureExtractor(**model_params).to(device)
|
"appearance_feature_extractor_params"
|
||||||
self.appearance_feature_extractor.load_state_dict(comfy.utils.load_torch_file(feature_extractor_path))
|
]
|
||||||
|
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()
|
self.appearance_feature_extractor.eval()
|
||||||
print('Load appearance_feature_extractor done.')
|
log.info("Load appearance_feature_extractor done.")
|
||||||
pbar.update(1)
|
pbar.update(1)
|
||||||
# init M
|
# 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 = 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()
|
self.motion_extractor.eval()
|
||||||
print('Load motion_extractor done.')
|
log.info("Load motion_extractor done.")
|
||||||
pbar.update(1)
|
pbar.update(1)
|
||||||
# init W
|
# 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 = 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()
|
self.warping_module.eval()
|
||||||
print('Load warping_module done.')
|
log.info("Load warping_module done.")
|
||||||
pbar.update(1)
|
pbar.update(1)
|
||||||
# init G
|
# 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 = 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()
|
self.spade_generator.eval()
|
||||||
print('Load spade_generator done.')
|
log.info("Load spade_generator done.")
|
||||||
pbar.update(1)
|
pbar.update(1)
|
||||||
|
|
||||||
def filter_checkpoint_for_model(checkpoint, prefix):
|
def filter_checkpoint_for_model(checkpoint, prefix):
|
||||||
"""Filter and adjust the checkpoint dictionary for a specific model based on the 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
|
# 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
|
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)
|
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_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.load_state_dict(stitcher_checkpoint)
|
||||||
stitcher = stitcher.to(device)
|
stitcher = stitcher.to(device)
|
||||||
stitcher.eval()
|
stitcher.eval()
|
||||||
|
|
||||||
lip_prefix = 'retarget_mouth'
|
lip_prefix = "retarget_mouth"
|
||||||
lip_checkpoint = filter_checkpoint_for_model(checkpoint, lip_prefix)
|
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.load_state_dict(lip_checkpoint)
|
||||||
retargetor_lip = retargetor_lip.to(device)
|
retargetor_lip = retargetor_lip.to(device)
|
||||||
retargetor_lip.eval()
|
retargetor_lip.eval()
|
||||||
|
|
||||||
eye_prefix = 'retarget_eye'
|
eye_prefix = "retarget_eye"
|
||||||
eye_checkpoint = filter_checkpoint_for_model(checkpoint, eye_prefix)
|
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.load_state_dict(eye_checkpoint)
|
||||||
retargetor_eye = retargetor_eye.to(device)
|
retargetor_eye = retargetor_eye.to(device)
|
||||||
retargetor_eye.eval()
|
retargetor_eye.eval()
|
||||||
print('Load stitching_retargeting_module done.')
|
log.info("Load stitching_retargeting_module done.")
|
||||||
|
|
||||||
self.stich_retargeting_module = {
|
self.stich_retargeting_module = {
|
||||||
'stitching': stitcher,
|
"stitching": stitcher,
|
||||||
'lip': retargetor_lip,
|
"lip": retargetor_lip,
|
||||||
'eye': retargetor_eye
|
"eye": retargetor_eye,
|
||||||
}
|
}
|
||||||
|
|
||||||
pipeline = LivePortraitPipeline(
|
pipeline = LivePortraitPipeline(
|
||||||
@@ -210,98 +222,379 @@ class DownloadAndLoadLivePortraitModels:
|
|||||||
self.spade_generator,
|
self.spade_generator,
|
||||||
self.stich_retargeting_module,
|
self.stich_retargeting_module,
|
||||||
InferenceConfig(
|
InferenceConfig(
|
||||||
device_id=device,
|
device_id=device,
|
||||||
flag_use_half_precision = True if precision == 'fp16' else False
|
flag_use_half_precision=True if precision == "fp16" else False,
|
||||||
)
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
return (pipeline,)
|
return (pipeline,)
|
||||||
|
|
||||||
|
|
||||||
class LivePortraitProcess:
|
class LivePortraitProcess:
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
return {"required": {
|
return {"required": {
|
||||||
|
|
||||||
"pipeline": ("LIVEPORTRAITPIPE",),
|
"pipeline": ("LIVEPORTRAITPIPE",),
|
||||||
|
"crop_info": ("CROPINFO", {"default": {}}),
|
||||||
"source_image": ("IMAGE",),
|
"source_image": ("IMAGE",),
|
||||||
"driving_images": ("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}),
|
"dsize": ("INT", {"default": 512, "min": 64, "max": 2048}),
|
||||||
"scale": ("FLOAT", {"default": 2.3, "min": 1.0, "max": 4.0, "step": 0.01}),
|
"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}),
|
"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.01}),
|
"vy_ratio": ("FLOAT", {"default": -0.125, "min": -1.0, "max": 1.0, "step": 0.001}),
|
||||||
"lip_zero": ("BOOLEAN", {"default": True}),
|
"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}),
|
"eye_retargeting": ("BOOLEAN", {"default": False}),
|
||||||
"eyes_retargeting_multiplier": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001}),
|
"eyes_retargeting_multiplier": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001}),
|
||||||
"lip_retargeting": ("BOOLEAN", {"default": False}),
|
"lip_retargeting": ("BOOLEAN", {"default": False}),
|
||||||
"lip_retargeting_multiplier": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001}),
|
"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_TYPES = ("RETARGETINGINFO",)
|
||||||
RETURN_NAMES = ("cropped_images", "full_images",)
|
RETURN_NAMES = ("retargeting_info",)
|
||||||
FUNCTION = "process"
|
FUNCTION = "process"
|
||||||
CATEGORY = "LivePortrait"
|
CATEGORY = "LivePortrait"
|
||||||
|
|
||||||
def process(self, source_image, driving_images, dsize, scale, vx_ratio, vy_ratio, pipeline,
|
def process(self, driving_crop_info, eye_retargeting, eyes_retargeting_multiplier, lip_retargeting, lip_retargeting_multiplier):
|
||||||
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()
|
|
||||||
|
|
||||||
crop_cfg = CropConfig(
|
driving_landmarks = []
|
||||||
dsize = dsize,
|
for crop in driving_crop_info["crop_info_list"]:
|
||||||
scale = scale,
|
driving_landmarks.append(crop['lmk_crop'])
|
||||||
vx_ratio = vx_ratio,
|
|
||||||
vy_ratio = vy_ratio,
|
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)
|
return (crop_info, keypoints_image_tensor,)
|
||||||
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)
|
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
"DownloadAndLoadLivePortraitModels": DownloadAndLoadLivePortraitModels,
|
"DownloadAndLoadLivePortraitModels": DownloadAndLoadLivePortraitModels,
|
||||||
"LivePortraitProcess": LivePortraitProcess,
|
"LivePortraitProcess": LivePortraitProcess,
|
||||||
|
"LivePortraitCropper": LivePortraitCropper,
|
||||||
|
"LivePortraitRetargeting": LivePortraitRetargeting,
|
||||||
|
#"KeypointScaler": KeypointScaler,
|
||||||
|
"KeypointsToImage": KeypointsToImage,
|
||||||
|
"LivePortraitLoadCropper": LivePortraitLoadCropper
|
||||||
}
|
}
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
"DownloadAndLoadLivePortraitModels": "(Down)Load LivePortraitModels",
|
"DownloadAndLoadLivePortraitModels": "(Down)Load LivePortraitModels",
|
||||||
"LivePortraitProcess": "LivePortraitProcess",
|
"LivePortraitProcess": "LivePortraitProcess",
|
||||||
|
"LivePortraitCropper": "LivePortraitCropper",
|
||||||
|
"LivePortraitRetargeting": "LivePortraitRetargeting",
|
||||||
|
#"KeypointScaler": "KeypointScaler",
|
||||||
|
"KeypointsToImage": "LivePortrait KeypointsToImage",
|
||||||
|
"LivePortraitLoadCropper": "LivePortrait LoadCropper"
|
||||||
}
|
}
|
||||||
@@ -1,5 +1,4 @@
|
|||||||
pyyaml
|
pyyaml
|
||||||
numpy
|
numpy
|
||||||
opencv-python
|
opencv-python
|
||||||
rich
|
|
||||||
onnxruntime-gpu
|
onnxruntime-gpu
|
||||||
Reference in New Issue
Block a user