testing
This commit is contained in:
@@ -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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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")
|
||||
@@ -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
|
||||
@@ -115,8 +115,9 @@ 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, disable=disable_progress_bar):
|
||||
|
||||
|
||||
safe_index = min(i, len(crop_info["crop_info_list"]) - 1)
|
||||
|
||||
@@ -272,11 +273,14 @@ 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)
|
||||
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)
|
||||
|
||||
out_list.append(out)
|
||||
|
||||
pbar.update(1)
|
||||
|
||||
|
||||
out_dict = {
|
||||
"out_list": out_list,
|
||||
"crop_info": crop_info,
|
||||
|
||||
@@ -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,7 +259,23 @@ class LivePortraitWrapper(object):
|
||||
ret_dct[k] = v.float()
|
||||
|
||||
return ret_dct
|
||||
|
||||
|
||||
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
|
||||
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 = []
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
pyyaml
|
||||
numpy
|
||||
opencv-python
|
||||
onnxruntime-gpu
|
||||
pykalman
|
||||
tensorrt
|
||||
pycuda
|
||||
ctypes
|
||||
Reference in New Issue
Block a user