Author SHA1 Message Date
kijai a5e1130088 testing 2024-07-23 21:13:47 +03:00
kijai 46675b2016 update workflows, cleanup 2024-07-23 21:13:15 +03:00
kijai 806263dd25 cleanup, fixes 2024-07-22 20:39:43 +03:00
kijai ac89dc1e2f fix no face frame skip 2024-07-22 19:19:46 +03:00
kijai 27d745b53e add other examples 2024-07-22 17:52:34 +03:00
kijai 052762578c Update readme.md 2024-07-22 17:45:49 +03:00
kijai e825c51c87 separate composition to it's own node 2024-07-22 16:21:35 +03:00
kijai 177b324fcd Update live_portrait_pipeline.py 2024-07-22 01:02:56 +03:00
kijai 5c03bd8439 MPS fallbacks 2024-07-22 00:57:44 +03:00
kijai ef5ff7075f Update requirements.txt 2024-07-21 20:35:33 +03:00
kijai 92fad03ee5 restructure a bit for more caching 2024-07-21 20:20:46 +03:00
kijai 4cefac79b8 Add single_frame mode for webcam 2024-07-21 19:42:22 +03:00
24 changed files with 3978 additions and 833 deletions
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,535 @@
{
"last_node_id": 203,
"last_link_id": 477,
"nodes": [
{
"id": 129,
"type": "LivePortraitLoadCropper",
"pos": [
-1050,
-740
],
"size": {
"0": 315,
"1": 82
},
"flags": {},
"order": 0,
"mode": 0,
"outputs": [
{
"name": "cropper",
"type": "LPCROPPER",
"links": [
444
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "LivePortraitLoadCropper"
},
"widgets_values": [
"CPU",
true
]
},
{
"id": 1,
"type": "DownloadAndLoadLivePortraitModels",
"pos": [
-1040,
-850
],
"size": {
"0": 302.43463134765625,
"1": 58
},
"flags": {},
"order": 1,
"mode": 0,
"outputs": [
{
"name": "live_portrait_pipe",
"type": "LIVEPORTRAITPIPE",
"links": [
446,
448
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "DownloadAndLoadLivePortraitModels"
},
"widgets_values": [
"fp16"
]
},
{
"id": 165,
"type": "ImageResizeKJ",
"pos": [
-670,
-560
],
"size": {
"0": 315,
"1": 242
},
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 466
},
{
"name": "get_image_size",
"type": "IMAGE",
"link": null
},
{
"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": [
434
],
"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,
"lanczos",
true,
2,
0,
0
]
},
{
"id": 196,
"type": "LoadImage",
"pos": [
-1050,
-550
],
"size": {
"0": 315,
"1": 314
},
"flags": {},
"order": 2,
"mode": 0,
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
466
],
"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": 78,
"type": "GetImageSizeAndCount",
"pos": [
-310,
-550
],
"size": {
"0": 210,
"1": 86
},
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 434
}
],
"outputs": [
{
"name": "image",
"type": "IMAGE",
"links": [
445,
475
],
"shape": 3,
"slot_index": 0
},
{
"name": "512 width",
"type": "INT",
"links": null,
"shape": 3
},
{
"name": "512 height",
"type": "INT",
"links": null,
"shape": 3
},
{
"name": "1 count",
"type": "INT",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "GetImageSizeAndCount"
}
},
{
"id": 190,
"type": "LivePortraitProcess",
"pos": [
563,
-418
],
"size": {
"0": 430.8000183105469,
"1": 282
},
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "pipeline",
"type": "LIVEPORTRAITPIPE",
"link": 448
},
{
"name": "crop_info",
"type": "CROPINFO",
"link": 449
},
{
"name": "source_image",
"type": "IMAGE",
"link": 475
},
{
"name": "driving_images",
"type": "IMAGE",
"link": 477
},
{
"name": "opt_retargeting_info",
"type": "RETARGETINGINFO",
"link": null
}
],
"outputs": [
{
"name": "cropped_image",
"type": "IMAGE",
"links": [
470
],
"shape": 3,
"slot_index": 0
},
{
"name": "output",
"type": "LP_OUT",
"links": [],
"shape": 3,
"slot_index": 1
}
],
"properties": {
"Node name for S&R": "LivePortraitProcess"
},
"widgets_values": [
false,
0.03,
false,
1,
"constant",
"single_frame",
0.000003
]
},
{
"id": 198,
"type": "PreviewImage",
"pos": [
1027,
-409
],
"size": {
"0": 521.2196044921875,
"1": 566.1187133789062
},
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 470
}
],
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 203,
"type": "Screencap_mss",
"pos": [
5,
-277
],
"size": {
"0": 315,
"1": 178
},
"flags": {},
"order": 3,
"mode": 0,
"outputs": [
{
"name": "image",
"type": "IMAGE",
"links": [
477
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "Screencap_mss"
},
"widgets_values": [
0,
0,
512,
512,
1,
0.1
]
},
{
"id": 189,
"type": "LivePortraitCropper",
"pos": [
-48,
-851
],
"size": {
"0": 330,
"1": 242
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "pipeline",
"type": "LIVEPORTRAITPIPE",
"link": 446,
"slot_index": 0
},
{
"name": "cropper",
"type": "LPCROPPER",
"link": 444
},
{
"name": "source_image",
"type": "IMAGE",
"link": 445
}
],
"outputs": [
{
"name": "cropped_image",
"type": "IMAGE",
"links": null,
"shape": 3,
"slot_index": 0
},
{
"name": "crop_info",
"type": "CROPINFO",
"links": [
449
],
"shape": 3,
"slot_index": 1
}
],
"properties": {
"Node name for S&R": "LivePortraitCropper"
},
"widgets_values": [
512,
2.3,
0,
-0.125,
0,
"large-small",
false
]
}
],
"links": [
[
434,
165,
0,
78,
0,
"IMAGE"
],
[
444,
129,
0,
189,
1,
"LPCROPPER"
],
[
445,
78,
0,
189,
2,
"IMAGE"
],
[
446,
1,
0,
189,
0,
"LIVEPORTRAITPIPE"
],
[
448,
1,
0,
190,
0,
"LIVEPORTRAITPIPE"
],
[
449,
189,
1,
190,
1,
"CROPINFO"
],
[
466,
196,
0,
165,
0,
"IMAGE"
],
[
470,
190,
0,
198,
0,
"IMAGE"
],
[
475,
78,
0,
190,
2,
"IMAGE"
],
[
477,
203,
0,
190,
3,
"IMAGE"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.7513148009015781,
"offset": {
"0": 1170.2642381365986,
"1": 992.3601372540302
}
}
},
"version": 0.4
}
File diff suppressed because it is too large Load Diff
-2
View File
@@ -35,8 +35,6 @@ class InferenceConfig(PrintableConfig):
output_fps: int = 30 # fps for output video
crf: int = 15 # crf for output video
flag_pasteback: bool = True # whether to paste-back/stitch the animated face cropping from the face-cropping space to the original image space
mask_crop = None
flag_write_gif: bool = False
device_id: int = 0
+3
View File
@@ -0,0 +1,3 @@
from .predictor import EfficientLivePortraitPredictor
from .config.config import save_config_to_yaml
from .utils import *
@@ -0,0 +1 @@
from .config import Config
@@ -0,0 +1,29 @@
# coding: utf-8
"""
pretty printing class
"""
from __future__ import annotations
import os.path as osp
from typing import Tuple
def make_abs_path(fn):
return osp.join(osp.dirname(osp.realpath(__file__)), fn)
class PrintableConfig: # pylint: disable=too-few-public-methods
"""Printable Config defining str function"""
def __repr__(self):
lines = [self.__class__.__name__ + ":"]
for key, val in vars(self).items():
if isinstance(val, Tuple):
flattened_val = "["
for item in val:
flattened_val += str(item) + "\n"
flattened_val = flattened_val.rstrip("\n")
val = flattened_val + "]"
lines += f"{key}: {str(val)}".split("\n")
return "\n ".join(lines)
+155
View File
@@ -0,0 +1,155 @@
import os
import requests
from dataclasses import dataclass, asdict
from typing import Literal, Tuple
from tqdm import tqdm
import torch.cuda
import yaml
# Define the URLs for the model files
MODEL_URLS = {
'live_portrait': {
'grid_sample_3d': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/libgrid_sample_3d_plugin.so?download=true',
'F_onnx': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/appearance_feature_extractor.onnx?download=true',
'M_onnx': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/motion_extractor.onnx?download=true',
'GW_onnx': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/generator_fix_grid.onnx?download=true',
'S_onnx': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/stitching.onnx?download=true',
'SE_onnx': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/stitching_eye.onnx?download=true',
'SL_onnx': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/stitching_lip.onnx?download=true',
# TensorRT FP32
'F_rt': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP32/resolve/main/appearance_feature_extractor_fp32.engine?download=true',
'M_rt': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP32/resolve/main/motion_extractor_fp32.engine?download=true',
'GW_rt': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP32/resolve/main/generator_fp32.engine?download=true',
'S_rt': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP32/resolve/main/stitching_fp32.engine?download=true',
'SE_rt': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP32/resolve/main/stitching_eye_fp32.engine?download=true',
'SL_rt': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP32/resolve/main/stitching_lip_fp32.engine?download=true',
# TensorRT FP16
'F_rt_half': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP16/resolve/main/appearance_feature_extractor_fp16.engine?download=true',
'M_rt_half': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP16/resolve/main/motion_extractor_fp16.engine?download=true',
'GW_rt_half': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP16/resolve/main/generator_fp16.engine?download=true',
'S_rt_half': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP16/resolve/main/stitching_fp16.engine?download=true',
'SE_rt_half': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP16/resolve/main/stitching_eye_fp16.engine?download=true',
'SL_rt_half': 'https://huggingface.co/myn0908/Live-Portrait-TensorRT-FP16/resolve/main/stitching_lip_fp16.engine?download=true'
},
'insightface': {
'arc_face': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/w600k_r50.onnx?download=true',
'2d106det': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/2d106det.onnx?download=true',
'det_10g': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/det_10g.onnx?download=true',
'landmark': 'https://huggingface.co/myn0908/Live-Portrait-ONNX/resolve/main/landmark.onnx?download=true'
}
}
# Function to download a file from a URL and save it locally
def downloading(url, outf):
if not os.path.exists(outf):
print(f"Downloading checkpoint to {outf}")
response = requests.get(url, stream=True)
total_size_in_bytes = int(response.headers.get('content-length', 0))
block_size = 1024 # 1 Kibibyte
progress_bar = tqdm(total=total_size_in_bytes, unit='iB', unit_scale=True)
with open(outf, 'wb') as file:
for data in response.iter_content(block_size):
progress_bar.update(len(data))
file.write(data)
progress_bar.close()
if total_size_in_bytes != 0 and progress_bar.n != total_size_in_bytes:
print("ERROR, something went wrong")
print(f"Downloaded successfully to {outf}")
else:
return outf
def get_efficient_live_portrait():
# Download the models and save them in the current working directory
current_dir = os.getcwd()
face_dir = os.path.join(current_dir, 'live_portrait_weights')
model_paths = {}
for main_key, sub_dict in MODEL_URLS.items():
dir_path = os.path.join(current_dir, 'live_portrait_weights', main_key)
os.makedirs(dir_path, exist_ok=True)
model_paths[main_key] = {}
for sub_key, url in sub_dict.items():
filename = url.split('/')[-1].split('?')[0]
save_path = os.path.join(dir_path, filename)
downloading(url, save_path)
model_paths[main_key][sub_key] = save_path
print('Downloaded successfully and already saved')
return model_paths, face_dir
@dataclass(repr=False) # use repr from PrintableConfig
class Config:
model_paths, face_dir = get_efficient_live_portrait()
grid_sample_3d: str = model_paths['live_portrait']['grid_sample_3d']
# ONNX
checkpoint_F: str = model_paths['live_portrait']['F_onnx'] # path to checkpoint
checkpoint_M: str = model_paths['live_portrait']['M_onnx'] # path to checkpoint
checkpoint_GW: str = model_paths['live_portrait']['GW_onnx']
checkpoint_S: str = model_paths['live_portrait']['S_onnx'] # path to checkpoint
checkpoint_SE: str = model_paths['live_portrait']['SE_onnx']
checkpoint_SL: str = model_paths['live_portrait']['SL_onnx']
# TensorRT FP32
F_rt: str = model_paths['live_portrait']['F_rt'] # path to checkpoint
M_rt: str = model_paths['live_portrait']['M_rt'] # path to checkpoint
GW_rt: str = model_paths['live_portrait']['GW_rt'] # path to checkpoint
S_rt: str = model_paths['live_portrait']['S_rt'] # path to checkpoint
SE_rt: str = model_paths['live_portrait']['SE_rt']
SL_rt: str = model_paths['live_portrait']['SL_rt']
# TensorRT FP16
F_rt_half: str = model_paths['live_portrait']['F_rt_half'] # path to checkpoint
M_rt_half: str = model_paths['live_portrait']['M_rt_half'] # path to checkpoint
GW_rt_half: str = model_paths['live_portrait']['GW_rt_half'] # path to checkpoint
S_rt_half: str = model_paths['live_portrait']['S_rt_half'] # path to checkpoint
SE_rt_half: str = model_paths['live_portrait']['SE_rt_half']
SL_rt_half: str = model_paths['live_portrait']['SL_rt_half']
flag_use_half_precision: bool = True # whether to use half precision
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
lip_zero_threshold: float = 0.03
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 motion
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 source portrait to the face-cropping space
flag_do_rot: bool = True # whether to conduct the rotation when flag_do_crop is True
flag_write_result: bool = True # whether to write output video
flag_write_gif: bool = False
anchor_frame: int = 0 # set this value if find_best_frame is True
input_shape: Tuple[int, int] = (256, 256) # input shape
output_format: Literal['mp4', 'gif'] = 'mp4' # output video format
output_fps: int = 30 # fps for output video
crf: int = 15 # crf for output video
mask_crop: str = 'None'
size_gif: int = 256
ref_max_shape: int = 1280
ref_shape_n: int = 2
device: str = 'cuda' if torch.cuda.is_available() else 'cpu'
# crop config
ckpt_landmark: str = model_paths['insightface']['landmark']
ckpt_arc_face: str = model_paths['insightface']['arc_face']
ckpt_landmark_106: str = model_paths['insightface']['2d106det']
ckpt_det: str = model_paths['insightface']['det_10g']
ckpt_face: str = face_dir
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
# Function to save the configuration to a YAML file
def save_config_to_yaml(filename="efficient-live-portrait.yaml"):
# Define the path where the YAML file will be saved
file_path = os.path.join(os.getcwd(), filename)
if not os.path.exists(file_path):
# Save the configuration to the YAML file
with open(file_path, 'w') as file:
yaml.safe_dump(asdict(Config()), file)
return file_path
+47
View File
@@ -0,0 +1,47 @@
from .utils.onnx_driver import ONNXEngine
import numpy as np
class EfficientLivePortraitPredictor:
def __init__(self, use_tensorrt=False, half=False, **kwargs):
super().__init__()
self.use_tensorrt = use_tensorrt
self.half = half
self.cfg = kwargs
if self.use_tensorrt:
from .utils.tensorrt_driver import TensorRTEngine
self.trt_engine = TensorRTEngine(self.half, **kwargs)
else:
self.onnx_engine = ONNXEngine().initialize_sessions(self.cfg)
def run_time(self, engine_name, task, inputs_onnx=None, inputs_tensorrt=None):
"""
Run inference using either TensorRT or ONNX Runtime based on the configuration.
Args:
- engine_name (str): Name of the engine/model.
- task (str): The task or model session name.
- inputs_onnx (dict): Input dict for inference.
- inputs_tensorrt(np.array or tensor): Input for inference TensorRT
Returns:
- The outputs from the inference.
"""
if self.use_tensorrt:
return self.trt_engine.inference_tensorrt(engine_name, inputs_tensorrt)
else:
return self.inference_onnx(task, inputs_onnx)
def inference_onnx(self, task, inputs):
"""
Perform inference using ONNX Runtime.
Args:
- task (str): The name of the task/model to use for inference.
- inputs (list or array): A list or array of input tensors.
Returns:
- List: The outputs of the inference.
"""
session = self.onnx_engine[task]
outputs = session.run(None, inputs)
return outputs
+1
View File
@@ -0,0 +1 @@
from .utils import *
@@ -0,0 +1,51 @@
import onnxruntime as ort
import torch
import numpy as np
from typing import Dict
class ONNXEngine:
def __init__(self):
pass
@staticmethod
def get_providers() -> list:
"""Returns the list of providers based on the current device."""
if ort.get_device() == 'GPU':
return ['CUDAExecutionProvider']
elif ort.get_device() == 'CPU':
return ['CPUExecutionProvider', 'CoreMLExecutionProvider']
else:
return []
def initialize_sessions(self, cfg) -> Dict[str, ort.InferenceSession]:
"""
Initialize ONNX InferenceSession instances for each model checkpoint.
Args:
- cfg (dict): Configuration dictionary containing checkpoint paths.
Returns:
- Dict[str, ort.InferenceSession]: Dictionary mapping session names to InferenceSession objects.
"""
#providers = self.get_providers()
providers = ['CUDAExecutionProvider']
# Initialize each session manually
gw_session = ort.InferenceSession("live_portrait_weights\\live_portrait\\generator_fix_grid.onnx", providers=providers)
# m_session = ort.InferenceSession(cfg.get("checkpoint_M"), providers=providers)
# f_session = ort.InferenceSession(cfg.get("checkpoint_F"), providers=providers)
# s_session = ort.InferenceSession(cfg.get("checkpoint_S"), providers=providers)
# se_session = ort.InferenceSession(cfg.get("checkpoint_SE"), providers=providers)
# sl_session = ort.InferenceSession(cfg.get("checkpoint_SL"), providers=providers)
# Return the sessions in a dictionary
return {
"gw_session": gw_session,
# "m_session": m_session,
# "f_session": f_session,
# "s_session": s_session,
# "se_session": se_session,
# "sl_session": sl_session
}
@@ -0,0 +1,231 @@
import tensorrt as trt
import pycuda.driver as cuda
import pycuda.gpuarray
import pycuda.autoinit
import numpy as np
import ctypes
from pathlib import Path
TRT_LOGGER = trt.Logger(trt.Logger.WARNING)
class Binding:
def __init__(self, engine, idx_or_name):
self.name = idx_or_name if isinstance(idx_or_name, str) else engine.get_tensor_name(idx_or_name)
if not self.name:
raise IndexError(f"Binding index out of range: {idx_or_name}")
self.is_input = engine.get_tensor_mode(self.name) == trt.TensorIOMode.INPUT
dtype = engine.get_tensor_dtype(self.name)
dtype_map = {
trt.DataType.FLOAT: np.float32,
trt.DataType.HALF: np.float16,
trt.DataType.INT8: np.int8,
trt.DataType.BOOL: np.bool_,
}
if hasattr(trt.DataType, 'INT32'):
dtype_map[trt.DataType.INT32] = np.int32
if hasattr(trt.DataType, 'INT64'):
dtype_map[trt.DataType.INT64] = np.int64
self.dtype = dtype_map[dtype]
self.shape = tuple(engine.get_tensor_shape(self.name))
self._host_buf = None
self._device_buf = None
@property
def host_buffer(self):
if self._host_buf is None:
self._host_buf = cuda.pagelocked_empty(self.shape, self.dtype)
return self._host_buf
@property
def device_buffer(self):
if self._device_buf is None:
self._device_buf = pycuda.gpuarray.empty(self.shape, self.dtype)
return self._device_buf
def get_async(self, stream):
self.device_buffer.get_async(stream, self.host_buffer)
return self.host_buffer
def cleanup(self):
if self._host_buf is not None:
del self._host_buf
if self._device_buf is not None:
del self._device_buf
class TensorRTEngine:
def __init__(self, half, **kwargs):
self.cfg = kwargs
self.cfx = None
if kwargs.get("cuda_ctx", None) is None:
cuda.init()
self.cfx = cuda.Device(0).make_context()
else:
self.cfx = kwargs.get("cuda_ctx")
if half:
self.model_paths = {
#'feature_extractor': self.cfg['F_rt_half'],
#'motion_extractor': self.cfg['M_rt_half'],
'generator': "live_portrait_weights\\live_portrait\\warping_spade-fix.engine",
#'stitching_retargeting': self.cfg['S_rt_half'],
#'stitching_retargeting_eye': self.cfg['SE_rt_half'],
#'stitching_retargeting_lip': self.cfg['SL_rt_half']
}
else:
self.model_paths = {
'feature_extractor': self.cfg['F_rt'],
'motion_extractor': self.cfg['M_rt'],
'generator': self.cfg['GW_rt'],
'stitching_retargeting': self.cfg['S_rt'],
'stitching_retargeting_eye': self.cfg['SE_rt'],
'stitching_retargeting_lip': self.cfg['SL_rt']
}
self.plugin_path = Path("N:\\AI\\ComfyUI\\live_portrait_weights\\live_portrait\\grid_sample_3d_plugin.dll")
self.load_plugins(TRT_LOGGER)
self.engines = {}
self.contexts = {}
self.bindings = {}
self.binding_addresses = {}
self.inputs = {}
self.outputs = {}
self.stream = cuda.Stream()
self.initialize_engines()
def load_plugins(self, logger: trt.Logger):
ctypes.CDLL(self.plugin_path, mode=ctypes.RTLD_GLOBAL)
trt.init_libnvinfer_plugins(logger, "")
def initialize_engines(self):
for model_name, model_path in self.model_paths.items():
engine = self.load_engine(model_path)
if engine is None:
raise RuntimeError(f"Failed to load engine for {model_name}")
context = engine.create_execution_context()
if context is None:
raise RuntimeError(f"Failed to create execution context for {model_name}")
bindings = [Binding(engine, i) for i in range(engine.num_io_tensors)]
self.engines[model_name] = engine
self.contexts[model_name] = context
self.bindings[model_name] = bindings
self.binding_addresses[model_name] = [b.device_buffer.ptr for b in bindings]
self.inputs[model_name] = [b for b in bindings if b.is_input]
self.outputs[model_name] = [b for b in bindings if not b.is_input]
self.prepare_buffers(model_name)
@staticmethod
def load_engine(engine_file_path):
with open(engine_file_path, "rb") as f, trt.Runtime(TRT_LOGGER) as runtime:
return runtime.deserialize_cuda_engine(f.read())
def prepare_buffers(self, model_name):
for binding in self.inputs[model_name] + self.outputs[model_name]:
_ = binding.device_buffer # Force buffer allocation
@staticmethod
def check_input_validity(input_idx, input_array, input_binding):
if input_array.shape != input_binding.shape:
if not (input_binding.shape == (1,) and input_array.shape == ()):
raise ValueError(
f"Wrong shape for input {input_idx}. Expected {input_binding.shape}, got {input_array.shape}.")
if input_array.dtype != input_binding.dtype:
if input_array.dtype == np.int64 and input_binding.dtype == np.int32:
input_array = input_array.astype(np.int32)
if not np.array_equal(input_array, input_array.astype(np.int64)):
raise TypeError(
f"Wrong dtype for input {input_idx}. Expected {input_binding.dtype}, got {input_array.dtype}. Cannot safely cast.")
else:
raise TypeError(
f"Wrong dtype for input {input_idx}. Expected {input_binding.dtype}, got {input_array.dtype}.")
return input_array
def run_sequential_tasks(self, model_name, inputs):
if model_name not in self.engines:
raise ValueError(f"Model name {model_name} not found in engines.")
engine = self.engines[model_name]
context = self.contexts[model_name]
binding_addresses = self.binding_addresses[model_name]
inputs_bindings = self.inputs[model_name]
outputs_bindings = self.outputs[model_name]
if isinstance(inputs, dict):
inputs = [inputs[b.name] for b in inputs_bindings]
if len(inputs) != len(inputs_bindings):
raise ValueError(f"Number of input arrays does not match number of input bindings for model {model_name}.")
self.cfx.push() # Push CUDA context
try:
for i, (input_array, input_binding) in enumerate(zip(inputs, inputs_bindings)):
input_array = self.check_input_validity(i, input_array, input_binding)
input_array = np.ascontiguousarray(input_array) # Ensure the input array is contiguous
cuda.memcpy_htod(input_binding.device_buffer.ptr, input_array)
for i in range(engine.num_io_tensors):
tensor_name = engine.get_tensor_name(i)
if i < len(inputs) and engine.is_shape_inference_io(tensor_name):
context.set_tensor_address(tensor_name, inputs[i].ctypes.data)
else:
context.set_tensor_address(tensor_name, binding_addresses[i])
context.execute_async_v3(self.stream.handle)
self.stream.synchronize()
outputs = []
for output in outputs_bindings:
host_output = np.empty(output.shape, dtype=output.dtype)
cuda.memcpy_dtoh(host_output, output.device_buffer.ptr)
outputs.append(host_output)
except Exception as e:
print(f"Error during inference for model {model_name}: {e}")
outputs = None
self.cfx.pop() # Pop CUDA context
return outputs
def inference_tensorrt(self, task, inputs):
if not isinstance(inputs, list):
raise TypeError("Inputs should be a list of numpy arrays or tensors.")
if task not in self.inputs:
raise ValueError(f"Task {task} not found in the model inputs.")
# Ensure all inputs are on the same memory type
if isinstance(inputs[0], pycuda.gpuarray.GPUArray):
# Ensure all inputs are on GPU
inputs = [cuda.to_gpu(input_array) if not isinstance(input_array, pycuda.gpuarray.GPUArray) else input_array
for input_array in inputs]
else:
# Ensure all inputs are on CPU
inputs = [input_array.get() if isinstance(input_array, pycuda.gpuarray.GPUArray) else input_array
for input_array in inputs]
inputs = [self.check_input_validity(i, np.array(input_array), self.inputs[task][i])
for i, input_array in enumerate(inputs)]
result = self.run_sequential_tasks(task, inputs)
return result
def __del__(self):
del self.engines
del self.contexts
del self.bindings
del self.binding_addresses
del self.inputs
del self.outputs
del self.stream
try:
if self.cfx is not None:
self.cfx.pop()
del self.cfx
except Exception as e:
print(f"Error during cleanup: {e}")
# Example usage
# engine = TensorRTEngine(half=True, F_rt_half="path/to/F_rt_half", M_rt_half="path/to/M_rt_half",
# GW_rt_half="path/to/GW_rt_half", S_rt_half="path/to/S_rt_half",
# SE_rt_half="path/to/SE_rt_half", SL_rt_half="path/to/SL_rt_half",
# grid_sample_3d="path/to/grid_sample_3d.so")
+202
View File
@@ -0,0 +1,202 @@
# coding: utf-8
"""
utility functions and classes to handle feature extraction and model loading
"""
import torch
import os
from glob import glob
import os.path as osp
import imageio
import numpy as np
import cv2
from rich.progress import track
cv2.setNumThreads(0)
cv2.ocl.setUseOpenCL(False)
def suffix(filename):
"""a.jpg -> jpg"""
pos = filename.rfind(".")
if pos == -1:
return ""
return filename[pos + 1:]
def prefix(filename):
"""a.jpg -> a"""
pos = filename.rfind(".")
if pos == -1:
return filename
return filename[:pos]
def basename(filename):
"""a/b/c.jpg -> c"""
return prefix(osp.basename(filename))
def is_video(file_path):
if file_path.lower().endswith((".mp4", ".mov", ".avi", ".webm")) or osp.isdir(file_path):
return True
return False
def is_template(file_path):
if file_path.endswith(".pkl"):
return True
return False
def mkdir(d, log=False):
# return self-assined `d`, for one line code
if not osp.exists(d):
os.makedirs(d, exist_ok=True)
if log:
print(f"Make dir: {d}")
return d
def squeeze_tensor_to_numpy(tensor):
out = tensor.data.squeeze(0).cpu().numpy()
return out
def dct2cuda(dct: dict, device: str):
for key in dct:
dct[key] = torch.tensor(dct[key]).to(device)
return dct
def concat_feat(kp_source: torch.Tensor, kp_driving: torch.Tensor) -> torch.Tensor:
"""
kp_source: (bs, k, 3)
kp_driving: (bs, k, 3)
Return: (bs, 2k*3)
"""
bs_src = kp_source.shape[0]
bs_dri = kp_driving.shape[0]
assert bs_src == bs_dri, 'batch size must be equal'
feat = torch.cat([kp_source.view(bs_src, -1), kp_driving.view(bs_dri, -1)], dim=1)
return feat
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}")
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')
return wfp
+50 -73
View File
@@ -8,9 +8,7 @@ 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
@@ -58,42 +56,36 @@ class LivePortraitPipeline(object):
inference_cfg = self.live_portrait_wrapper.cfg
device = inference_cfg.device_id
cropped_image_list = []
composited_image_list = []
out_mask_list = []
out_list = []
R_d_0, x_d_0_info = None, None
if mismatch_method == "cut":
if mismatch_method == "cut" or relative_motion_mode == "source_video_smoothed":
total_frames = source_np.shape[0]
else:
total_frames = driving_images.shape[0]
disable_progress_bar = True if relative_motion_mode == "single_frame" else False
source_info = []
source_rot_list = []
f_s_list = []
for i in tqdm(range(source_np.shape[0]), desc='Processing source images...', total=source_np.shape[0]):
#get source keypoints info
img_crop_256x256 = crop_info["crop_info_list"][i]["img_crop_256x256"]
I_s = self.live_portrait_wrapper.prepare_source(img_crop_256x256)
x_s_info = self.live_portrait_wrapper.get_kp_info(I_s)
f_s = self.live_portrait_wrapper.extract_feature_3d(I_s)
f_s_list.append(f_s)
source_info.append(x_s_info)
R_s = get_rotation_matrix(
x_s_info["pitch"], x_s_info["yaw"], x_s_info["roll"]
)
source_rot_list.append(R_s)
source_info = crop_info["source_info"]
source_rot_list = crop_info["source_rot_list"]
f_s_list = crop_info["f_s_list"]
x_s_list = crop_info["x_s_list"]
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]):
for i in tqdm(range(driving_images.shape[0]), desc='Processing driving images...', total=driving_images.shape[0], disable=disable_progress_bar):
#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 crop_info["crop_info_list"][safe_index] is None:
driving_info.append(None)
driving_rot_list.append(None)
driving_exp_list.append(None)
continue
x_d_info = self.live_portrait_wrapper.get_kp_info(driving_images[i].unsqueeze(0).to(device))
if i == 0:
first = x_d_info
@@ -111,6 +103,9 @@ class LivePortraitPipeline(object):
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]):
if driving_rot_list[i] is None:
x_d_r_lst.append(None)
continue
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
@@ -120,28 +115,29 @@ class LivePortraitPipeline(object):
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):
for i in tqdm(range(total_frames), desc='Animating...', total=total_frames, disable=disable_progress_bar):
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))
if crop_info["crop_info_list"][safe_index] is None:
out_list.append({})
pbar.update(1)
continue
source_lmk = crop_info["crop_info_list"][safe_index]["lmk_crop"]
x_d_info = driving_info[i]
R_d = driving_rot_list[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)
x_s = x_s_list[safe_index]
x_c_s = x_s_info["kp"]
#lip zero
if inference_cfg.flag_lip_zero:
@@ -153,13 +149,10 @@ class LivePortraitPipeline(object):
else:
lip_delta_before_animation = (self.live_portrait_wrapper.retarget_lip(x_s, combined_lip_ratio_tensor_before_animation))
R_d = driving_rot_list[i]
if i == 0:
R_d_0 = R_d
x_d_0_info = x_d_info
if relative_motion_mode == "relative":
if i == 0:
R_d_0 = R_d
x_d_0_info = x_d_info
R_new = (R_d @ R_d_0.permute(0, 2, 1)) @ R_s
delta_new = x_s_info["exp"] + (x_d_info["exp"] - x_d_0_info["exp"])
scale_new = x_s_info["scale"] * (x_d_info["scale"] / x_d_0_info["scale"])
@@ -174,6 +167,11 @@ class LivePortraitPipeline(object):
delta_new = x_s_info['exp']
scale_new = x_s_info["scale"]
t_new = x_d_info["t"]
elif relative_motion_mode == "single_frame":
R_new = R_d
delta_new = x_d_info['exp']
scale_new = x_s_info["scale"]
t_new = x_d_info["t"]
else:
R_new = R_d
delta_new = x_s_info['exp']
@@ -275,39 +273,18 @@ class LivePortraitPipeline(object):
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)
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)
out = self.live_portrait_wrapper.warp_decode_tensorrt(f_s, x_s, x_d_i_new)
#out = self.live_portrait_wrapper.warp_decode(f_s, x_s, x_d_i_new)
# 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)
out_list.append(out)
pbar.update(1)
return cropped_image_list, composited_image_list, out_mask_list
out_dict = {
"out_list": out_list,
"crop_info": crop_info,
"mismatch_method": mismatch_method,
}
return out_dict
+19 -9
View File
@@ -15,6 +15,8 @@ from .utils.retargeting_utils import calc_eye_close_ratio, calc_lip_close_ratio
from .config.inference_config import InferenceConfig
from contextlib import nullcontext
from .efficient import EfficientLivePortraitPredictor
from comfy.model_management import get_autocast_device
class LivePortraitWrapper(object):
@@ -32,6 +34,8 @@ class LivePortraitWrapper(object):
self.device_id = cfg.device_id
self.timer = Timer()
self.predictor = EfficientLivePortraitPredictor(use_tensorrt = True, half = True)
def prepare_source(self, img: np.ndarray) -> torch.Tensor:
""" construct the input as standard
img: HxWx3, uint8, 256x256
@@ -255,17 +259,23 @@ class LivePortraitWrapper(object):
ret_dct[k] = v.float()
return ret_dct
def parse_output(self, out: torch.Tensor) -> np.ndarray:
""" construct the output as standard
return: 1xHxWx3, uint8
"""
out = np.transpose(out.data.cpu().numpy(), [0, 2, 3, 1]) # 1x3xHxW -> 1xHxWx3
def warp_decode_tensorrt(self, feature_3d, kp_source, kp_driving):
inputs = {
'feature_3d': np.array(feature_3d.cpu()),
'kp_driving': np.array(kp_driving.cpu()),
'kp_source': np.array(kp_source.cpu())
}
generator = self.predictor.run_time(engine_name='generator', task='gw_session',
inputs_onnx=inputs, inputs_tensorrt=[feature_3d.cpu(), kp_driving.cpu(), kp_source.cpu()])
out = np.transpose(generator[0], [0, 2, 3, 1]) # 1x3xHxW -> 1xHxWx3
out = np.clip(out, 0, 1) # clip to 0~1
out = np.clip(out * 255, 0, 255).astype(np.uint8) # 0~1 -> 0~255
return out
out = torch.from_numpy(out).permute(0, 3, 1, 2) / 255
return {'out': out}
def calc_retargeting_ratio(self, source_lmk, driving_lmk_lst):
input_eye_ratio_lst = []
input_lip_ratio_lst = []
+8 -2
View File
@@ -47,7 +47,13 @@ class DenseMotionNetwork(nn.Module):
feature_repeat = feature.unsqueeze(1).unsqueeze(1).repeat(1, self.num_kp+1, 1, 1, 1, 1, 1) # (bs, num_kp+1, 1, c, d, h, w)
feature_repeat = feature_repeat.view(bs * (self.num_kp+1), -1, d, h, w) # (bs*(num_kp+1), c, d, h, w)
sparse_motions = sparse_motions.view((bs * (self.num_kp+1), d, h, w, -1)) # (bs*(num_kp+1), d, h, w, 3)
sparse_deformed = F.grid_sample(feature_repeat, sparse_motions, align_corners=False)
try:
sparse_deformed = F.grid_sample(feature_repeat, sparse_motions, align_corners=False)
except NotImplementedError: #MPS fallback
out_device = feature_repeat.device # Store input device
feature_repeat = feature_repeat.to('cpu')
sparse_motions = sparse_motions.to('cpu')
sparse_deformed = F.grid_sample(feature_repeat, sparse_motions, align_corners=False).to(out_device)
sparse_deformed = sparse_deformed.view((bs, self.num_kp+1, -1, d, h, w)) # (bs, num_kp+1, c, d, h, w)
return sparse_deformed
@@ -61,7 +67,7 @@ class DenseMotionNetwork(nn.Module):
# adding background feature
try:
zeros = torch.zeros(heatmap.shape[0], 1, spatial_size[0], spatial_size[1], spatial_size[2]).type(heatmap.type()).to(heatmap.device)
except:
except ValueError:
zeros = torch.zeros(heatmap.shape[0], 1, spatial_size[0], spatial_size[1], spatial_size[2]).to(heatmap.device)
heatmap = torch.cat([zeros, heatmap], dim=1)
heatmap = heatmap.unsqueeze(2) # (bs, 1+num_kp, 1, d, h, w)
+5 -1
View File
@@ -158,7 +158,11 @@ class DownBlock3d(nn.Module):
out = self.conv(x)
out = self.norm(out)
out = F.relu(out)
out = self.pool(out)
try:
out = self.pool(out)
except NotImplementedError:
out_device = out.device # Store input device
out = self.pool(out.to('cpu')).to(out_device)
return out
+5 -1
View File
@@ -44,7 +44,11 @@ class WarpingNetwork(nn.Module):
self.estimate_occlusion_map = estimate_occlusion_map
def deform_input(self, inp, deformation):
return F.grid_sample(inp, deformation, align_corners=False)
try:
return F.grid_sample(inp, deformation, align_corners=False)
except NotImplementedError:
out_device = inp.device # Store input device
return F.grid_sample(inp.to('cpu'), deformation.to('cpu'), align_corners=False).to(out_device)
def forward(self, feature_3d, kp_driving, kp_source):
if self.dense_motion_network is not None:
+51 -19
View File
@@ -30,14 +30,14 @@ def _transform_img(img, M, dsize, flags=CV2_INTERP, borderMode=None):
import torch
import kornia.geometry.transform as KGT
def _transform_img_kornia(img, M, dsize, flags='bilinear', borderMode='zeros'):
def _transform_img_kornia(img, M, dsize, device, 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)
@@ -124,29 +124,60 @@ def parse_pt2_from_pt203(pt203, use_lip=True):
pt2 = np.stack([pt_left_eye, pt_right_eye], axis=0)
return pt2
def parse_pt2_from_pt9(pt9, use_lip=True):
'''
animal_face = {"keypoints": ['right eye right', 'right eye left', 'left eye right', 'left eye left', 'nose tip', 'lip right', 'lip left', 'upper lip', 'lower lip'], "skeleton": []}
def parse_pt2_from_pt68(pt68, use_lip=True):
"""
parsing the 2 points according to the 68 points, which cancels the roll
"""
lm_idx = np.array([31, 37, 40, 43, 46, 49, 55], dtype=np.int32) - 1
'''
if use_lip:
pt5 = np.stack([
np.mean(pt68[lm_idx[[1, 2]], :], 0), # left eye
np.mean(pt68[lm_idx[[3, 4]], :], 0), # right eye
pt68[lm_idx[0], :], # nose
pt68[lm_idx[5], :], # lip
pt68[lm_idx[6], :] # lip
pt9 = np.stack([
(pt9[2]+pt9[3])/2, # left eye
(pt9[0]+pt9[1])/2, # right eye
pt9[4],
# (pt9[5]+pt9[6]+pt9[7]+pt9[8])/4 # lip
(pt9[5] + pt9[6] ) / 2 # lip
], axis=0)
pt2 = np.stack([
(pt5[0] + pt5[1]) / 2,
(pt5[3] + pt5[4]) / 2
(pt9[0] + pt9[1]) / 2, # eye
pt9[3] # lip
], axis=0)
else:
pt2 = np.stack([
np.mean(pt68[lm_idx[[1, 2]], :], 0), # left eye
np.mean(pt68[lm_idx[[3, 4]], :], 0), # right eye
(pt9[2] + pt9[3]) / 2,
(pt9[0] + pt9[1]) / 2,
], axis=0)
return pt2
def parse_pt2_from_pt68(pt68, use_lip=True):
'''
face = {"keypoints": ['right cheekbone 1', 'right cheekbone 2', 'right cheek 1', 'right cheek 2', 'right cheek 3', 'right cheek 4', 'right cheek 5', 'right chin', 'chin center',
'left chin', 'left cheek 5', 'left cheek 4', 'left cheek 3', 'left cheek 2', 'left cheek 1', 'left cheekbone 2', 'left cheekbone 1', 'right eyebrow 1', 'right eyebrow 2', 'right eyebrow 3',
'right eyebrow 4', 'right eyebrow 5', 'left eyebrow 1', 'left eyebrow 2', 'left eyebrow 3', 'left eyebrow 4', 'left eyebrow 5', 'nasal bridge 1', 'nasal bridge 2', 'nasal bridge 3', 'nasal bridge 4',
'right nasal wing 1', 'right nasal wing 2', 'nasal wing center', 'left nasal wing 1', 'left nasal wing 2', 'right eye eye corner 1', 'right eye upper eyelid 1', 'right eye upper eyelid 2',
'right eye eye corner 2', 'right eye lower eyelid 2', 'right eye lower eyelid 1', 'left eye eye corner 1', 'left eye upper eyelid 1', 'left eye upper eyelid 2', 'left eye eye corner 2', 'left eye lower eyelid 2',
'left eye lower eyelid 1', 'right mouth corner', 'upper lip outer edge 1', 'upper lip outer edge 2', 'upper lip outer edge 3', 'upper lip outer edge 4', 'upper lip outer edge 5', 'left mouth corner',
'lower lip outer edge 5', 'lower lip outer edge 4', 'lower lip outer edge 3', 'lower lip outer edge 2', 'lower lip outer edge 1', 'upper lip inter edge 1', 'upper lip inter edge 2', 'upper lip inter edge 3',
'upper lip inter edge 4', 'upper lip inter edge 5', 'lower lip inter edge 3', 'lower lip inter edge 2', 'lower lip inter edge 1'], "skeleton": []}
'''
if use_lip:
pt68 = np.stack([
(pt68[42] + pt68[43] + pt68[44] + pt68[45] + pt68[46]+ pt68[47])/6, # left eye
(pt68[36] + pt68[37] + pt68[38] + pt68[39] + pt68[40] + pt68[41]) / 6, # right eye
(pt68[48] + pt68[54])/2
], axis=0)
pt2 = np.stack([
(pt68[0] + pt68[1]) / 2,
pt68[2]
], axis=0)
else:
pt2 = np.stack([
(pt68[42] + pt68[43] + pt68[44] + pt68[45] + pt68[46] + pt68[47]) / 6, # left eye
(pt68[36] + pt68[37] + pt68[38] + pt68[39] + pt68[40] + pt68[41]) / 6, # right eye
], axis=0)
return pt2
@@ -183,6 +214,8 @@ def parse_pt2_from_pt_x(pts, use_lip=True):
elif pts.shape[0] > 101:
# take the first 101 points
pt2 = parse_pt2_from_pt101(pts[:101], use_lip=use_lip)
elif pts.shape[0] == 9:
pt2 = parse_pt2_from_pt9(pts, use_lip=use_lip)
else:
raise Exception(f'Unknow shape: {pts.shape}')
@@ -427,4 +460,3 @@ def average_bbox_lst(bbox_lst):
return None
bbox_arr = np.array(bbox_lst)
return np.mean(bbox_arr, axis=0).tolist()
+15 -2
View File
@@ -4,14 +4,27 @@ 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]
# Reshape x_d_lst, skipping None values
x_d_lst_reshape = [x.reshape(-1) for x in x_d_lst if x is not None]
if not x_d_lst_reshape: # Check if x_d_lst_reshape is empty after filtering
return [None] * len(x_d_lst) # Return a list of Nones with the same length as 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]
# Initialize an iterator for smoothed_state_means
smoothed_states_iter = iter(smoothed_state_means)
# Create x_d_lst_smooth, inserting None for each None encountered in the original list
x_d_lst_smooth = [torch.tensor(next(smoothed_states_iter).reshape(shape[-2:]), dtype=torch.float32, device=device) if x is not None else None for x in x_d_lst]
return x_d_lst_smooth
+187 -44
View File
@@ -21,6 +21,8 @@ from .liveportrait.modules.appearance_feature_extractor import (
from .liveportrait.modules.stitching_retargeting_network import (
StitchingRetargetingNetwork,
)
from .liveportrait.utils.camera import get_rotation_matrix
from .liveportrait.utils.crop import _transform_img_kornia
import logging
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
@@ -29,7 +31,6 @@ log = logging.getLogger(__name__)
class InferenceConfig:
def __init__(
self,
mask_crop=None,
flag_use_half_precision=True,
flag_lip_zero=True,
lip_zero_threshold=0.03,
@@ -37,9 +38,7 @@ class InferenceConfig:
flag_lip_retargeting=False,
flag_stitching=True,
flag_relative=True,
flag_relative_rotation_only=False,
input_shape=(256, 256),
flag_pasteback=True,
device_id=0,
flag_do_crop=True,
flag_do_rot=True,
@@ -51,14 +50,11 @@ class InferenceConfig:
self.flag_lip_retargeting = flag_lip_retargeting
self.flag_stitching = flag_stitching
self.flag_relative = flag_relative
self.flag_relative_rotation_only = flag_relative_rotation_only
self.input_shape = input_shape
self.flag_pasteback = flag_pasteback
self.device_id = device_id
self.flag_do_crop = flag_do_crop
self.flag_do_rot = flag_do_rot
self.mask_crop = mask_crop
class DownloadAndLoadLivePortraitModels:
@classmethod
def INPUT_TYPES(s):
@@ -88,19 +84,19 @@ class DownloadAndLoadLivePortraitModels:
if precision == 'auto':
try:
if mm.is_device_mps(device):
print("LivePortrait using fp32 for MPS")
log.info("LivePortrait using fp32 for MPS")
dtype = 'fp32'
elif mm.should_use_fp16():
print("LivePortrait using fp16")
log.info("LivePortrait using fp16")
dtype = 'fp16'
else:
print("LivePortrait using fp32")
log.info("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}")
log.info(f"LivePortrait using {dtype}")
pbar = comfy.utils.ProgressBar(3)
@@ -258,6 +254,7 @@ class LivePortraitProcess:
"relative",
"source_video_smoothed",
"relative_rotation_only",
"single_frame",
"off"
],
),
@@ -265,20 +262,17 @@ class LivePortraitProcess:
},
"optional": {
"mask": ("MASK", {"default": None}),
"opt_retargeting_info": ("RETARGETINGINFO", {"default": None}),
}
}
RETURN_TYPES = (
"IMAGE",
"IMAGE",
"MASK",
"LP_OUT",
)
RETURN_NAMES = (
"cropped_images",
"full_images",
"mask",
"cropped_image",
"output",
)
FUNCTION = "process"
CATEGORY = "LivePortrait"
@@ -296,7 +290,6 @@ class LivePortraitProcess:
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]:
@@ -327,22 +320,16 @@ class LivePortraitProcess:
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
if driving_images.shape[1] != 256 or driving_images.shape[2] != 256:
driving_images_256 = comfy.utils.common_upscale(driving_images.permute(0, 3, 1, 2), 256, 256, "lanczos", "disabled")
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 = driving_images.permute(0, 3, 1, 2)
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(
out = pipeline.execute(
source_np,
driving_images_256,
crop_info,
@@ -352,20 +339,137 @@ class LivePortraitProcess:
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)
total_frames = len(out["out_list"])
if total_frames > 1:
cropped_image_list = []
for i in (range(total_frames)):
if not out["out_list"][i]:
cropped_image_list.append(torch.zeros(1, 512, 512, 3, dtype=torch.float32, device = "cpu"))
else:
cropped_image = torch.clamp(out["out_list"][i]["out"], 0, 1).permute(0, 2, 3, 1).cpu()
cropped_image_list.append(cropped_image)
cropped_out_tensors = torch.cat(cropped_image_list, dim=0)
else:
cropped_out_tensors = torch.clamp(out["out_list"][0]["out"], 0, 1).permute(0, 2, 3, 1)
return (cropped_out_tensors, out,)
class LivePortraitComposite:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"source_image": ("IMAGE",),
"cropped_image": ("IMAGE",),
"liveportrait_out": ("LP_OUT", ),
},
"optional": {
"mask": ("MASK", {"default": None}),
}
}
RETURN_TYPES = (
"IMAGE",
"MASK",
)
RETURN_NAMES = (
"full_images",
"mask",
)
FUNCTION = "process"
CATEGORY = "LivePortrait"
def process(self, source_image, cropped_image, liveportrait_out, mask=None):
mm.soft_empty_cache()
device = mm.get_torch_device()
if mm.is_device_mps(device):
device = torch.device('cpu') #this function returns NaNs on MPS, defaulting to CPU
B, H, W, C = source_image.shape
source_image = source_image.permute(0, 3, 1, 2) # B,H,W,C -> B,C,H,W
cropped_image = cropped_image.permute(0, 3, 1, 2)
if mask is not None:
crop_mask = mask.unsqueeze(-1).expand(-1, -1, -1, 3)
else:
log.info("Using default mask template")
crop_mask = cv2.imread(os.path.join(script_directory, "liveportrait", "utils", "resources", "mask_template.png"), cv2.IMREAD_COLOR)
crop_mask = torch.from_numpy(crop_mask)
crop_mask = crop_mask.unsqueeze(0).float() / 255.0
crop_info = liveportrait_out["crop_info"]
composited_image_list = []
out_mask_list = []
total_frames = len(liveportrait_out["out_list"])
log.info(f"Total frames: {total_frames}")
pbar = comfy.utils.ProgressBar(total_frames)
for i in tqdm(range(total_frames), desc='Compositing..', total=total_frames):
safe_index = min(i, len(crop_info["crop_info_list"]) - 1)
if liveportrait_out["mismatch_method"] == "cut":
source_frame = source_image[safe_index].unsqueeze(0).to(device)
else:
source_frame = _get_source_frame(source_image, i, liveportrait_out["mismatch_method"]).unsqueeze(0).to(device)
if not liveportrait_out["out_list"][i]:
composited_image_list.append(source_frame)
out_mask_list.append(torch.zeros((1, 3, H, W), device=device))
else:
cropped_image = torch.clamp(liveportrait_out["out_list"][i]["out"], 0, 1).permute(0, 2, 3, 1)
# Transform and blend
cropped_image_to_original = _transform_img_kornia(
cropped_image,
crop_info["crop_info_list"][safe_index]["M_c2o"],
dsize=(W, H),
device=device
)
mask_ori = _transform_img_kornia(
crop_mask,
crop_info["crop_info_list"][safe_index]["M_c2o"],
dsize=(W, H),
device=device
)
cropped_image_to_original_blend = torch.clip(
mask_ori * cropped_image_to_original + (1 - mask_ori) * source_frame, 0, 1
)
composited_image_list.append(cropped_image_to_original_blend)
out_mask_list.append(mask_ori)
pbar.update(1)
full_tensors_out = torch.cat(composited_image_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]
mask_tensors_out = mask_tensors_out[:, 0, :, :]
return (
cropped_out_tensors.cpu().float(),
full_tensors_out.cpu().float(),
mask_tensors_out.cpu().float()
)
def _get_source_frame(source, idx, method):
if source.shape[0] == 1:
return source[0]
if method == "constant":
return source[min(idx, source.shape[0] - 1)]
elif method == "cycle":
return source[idx % source.shape[0]]
elif method == "mirror":
cycle_length = 2 * source.shape[0] - 2
mirror_idx = idx % cycle_length
if mirror_idx >= source.shape[0]:
mirror_idx = cycle_length - mirror_idx
return source[mirror_idx]
class LivePortraitLoadCropper:
@classmethod
@@ -401,6 +505,7 @@ class LivePortraitCropper:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"pipeline": ("LIVEPORTRAITPIPE",),
"cropper": ("LPCROPPER",),
"source_image": ("IMAGE",),
"dsize": ("INT", {"default": 512, "min": 64, "max": 2048}),
@@ -428,31 +533,67 @@ class LivePortraitCropper:
FUNCTION = "process"
CATEGORY = "LivePortrait"
def process(self, cropper, source_image, dsize, scale, vx_ratio, vy_ratio, face_index, face_index_order, rotate):
def process(self, pipeline, cropper, source_image, dsize, scale, vx_ratio, vy_ratio, face_index, face_index_order, rotate):
source_image_np = (source_image * 255).byte().numpy()
# Initialize lists
crop_info_list = []
cropped_images_list = []
source_info = []
source_rot_list = []
f_s_list = []
x_s_list = []
# Initialize a progress bar for the combined operation
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)):
for i in tqdm(range(len(source_image_np)), desc='Detecting, cropping, and processing..', total=len(source_image_np)):
# Cropping operation
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)
# Processing source images
if crop_info:
crop_info_list.append(crop_info)
cropped_image = crop_info['img_crop_256x256']
cropped_images_list.append(cropped_image)
I_s = pipeline.live_portrait_wrapper.prepare_source(cropped_image)
x_s_info = pipeline.live_portrait_wrapper.get_kp_info(I_s)
source_info.append(x_s_info)
x_s = pipeline.live_portrait_wrapper.transform_keypoint(x_s_info)
x_s_list.append(x_s)
R_s = get_rotation_matrix(x_s_info["pitch"], x_s_info["yaw"], x_s_info["roll"])
source_rot_list.append(R_s)
f_s = pipeline.live_portrait_wrapper.extract_feature_3d(I_s)
f_s_list.append(f_s)
else:
log.warning(f"Warning: No face detected on frame {str(i)}, skipping")
cropped_image = np.zeros((256, 256, 3), dtype=np.uint8)
cropped_images_list.append(cropped_image)
crop_info_list.append(None)
f_s_list.append(None)
x_s_list.append(None)
source_info.append(None)
source_rot_list.append(None)
# Update progress bar
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
'crop_info_list': crop_info_list,
'source_rot_list': source_rot_list,
'f_s_list': f_s_list,
'x_s_list': x_s_list,
'source_info': source_info
}
return (cropped_tensors_out, crop_info_dict)
@@ -587,7 +728,8 @@ NODE_CLASS_MAPPINGS = {
"LivePortraitRetargeting": LivePortraitRetargeting,
#"KeypointScaler": KeypointScaler,
"KeypointsToImage": KeypointsToImage,
"LivePortraitLoadCropper": LivePortraitLoadCropper
"LivePortraitLoadCropper": LivePortraitLoadCropper,
"LivePortraitComposite": LivePortraitComposite,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DownloadAndLoadLivePortraitModels": "(Down)Load LivePortraitModels",
@@ -596,5 +738,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"LivePortraitRetargeting": "LivePortraitRetargeting",
#"KeypointScaler": "KeypointScaler",
"KeypointsToImage": "LivePortrait KeypointsToImage",
"LivePortraitLoadCropper": "LivePortrait LoadCropper"
"LivePortraitLoadCropper": "LivePortrait LoadCropper",
"LivePortraitComposite": "LivePortrait Composite",
}
+6 -1
View File
@@ -1,7 +1,12 @@
# ComfyUI nodes to use [LivePortrait](https://github.com/KwaiVGI/LivePortrait)
## Update
Rework of almost the whole thing that's been in develop is now merged into main, this means old workflows will not work, but everything should be faster and there's lots of new features.
For legacy purposes the old main branch is moved to the legacy -branch
https://github.com/kijai/ComfyUI-LivePortrait/assets/40791699/e55e10f6-af61-4d73-b162-af29eb847516
I have converted all the pickle files to safetensors: https://huggingface.co/Kijai/LivePortrait_safetensors/tree/main
+8
View File
@@ -0,0 +1,8 @@
pyyaml
numpy
opencv-python
onnxruntime-gpu
pykalman
tensorrt
pycuda
ctypes
+2 -1
View File
@@ -1,4 +1,5 @@
pyyaml
numpy
opencv-python
onnxruntime-gpu
onnxruntime-gpu
pykalman