From a5e1130088935c072719005bfd2a26ef56f11bc3 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 23 Jul 2024 21:13:47 +0300 Subject: [PATCH] testing --- liveportrait/efficient/__init__.py | 3 + liveportrait/efficient/config/__init__.py | 1 + liveportrait/efficient/config/base_config.py | 29 +++ liveportrait/efficient/config/config.py | 155 ++++++++++++ liveportrait/efficient/predictor.py | 47 ++++ liveportrait/efficient/utils/__init__.py | 1 + liveportrait/efficient/utils/onnx_driver.py | 51 ++++ .../efficient/utils/tensorrt_driver.py | 231 ++++++++++++++++++ liveportrait/efficient/utils/utils.py | 202 +++++++++++++++ liveportrait/live_portrait_pipeline.py | 8 +- liveportrait/live_portrait_wrapper.py | 22 +- requirements-tensorrt.txt | 8 + 12 files changed, 755 insertions(+), 3 deletions(-) create mode 100644 liveportrait/efficient/__init__.py create mode 100644 liveportrait/efficient/config/__init__.py create mode 100644 liveportrait/efficient/config/base_config.py create mode 100644 liveportrait/efficient/config/config.py create mode 100644 liveportrait/efficient/predictor.py create mode 100644 liveportrait/efficient/utils/__init__.py create mode 100644 liveportrait/efficient/utils/onnx_driver.py create mode 100644 liveportrait/efficient/utils/tensorrt_driver.py create mode 100644 liveportrait/efficient/utils/utils.py create mode 100644 requirements-tensorrt.txt diff --git a/liveportrait/efficient/__init__.py b/liveportrait/efficient/__init__.py new file mode 100644 index 0000000..69d7d6a --- /dev/null +++ b/liveportrait/efficient/__init__.py @@ -0,0 +1,3 @@ +from .predictor import EfficientLivePortraitPredictor +from .config.config import save_config_to_yaml +from .utils import * \ No newline at end of file diff --git a/liveportrait/efficient/config/__init__.py b/liveportrait/efficient/config/__init__.py new file mode 100644 index 0000000..3558f42 --- /dev/null +++ b/liveportrait/efficient/config/__init__.py @@ -0,0 +1 @@ +from .config import Config \ No newline at end of file diff --git a/liveportrait/efficient/config/base_config.py b/liveportrait/efficient/config/base_config.py new file mode 100644 index 0000000..216b8be --- /dev/null +++ b/liveportrait/efficient/config/base_config.py @@ -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) diff --git a/liveportrait/efficient/config/config.py b/liveportrait/efficient/config/config.py new file mode 100644 index 0000000..f89e859 --- /dev/null +++ b/liveportrait/efficient/config/config.py @@ -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 diff --git a/liveportrait/efficient/predictor.py b/liveportrait/efficient/predictor.py new file mode 100644 index 0000000..1a90075 --- /dev/null +++ b/liveportrait/efficient/predictor.py @@ -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 diff --git a/liveportrait/efficient/utils/__init__.py b/liveportrait/efficient/utils/__init__.py new file mode 100644 index 0000000..90f60fd --- /dev/null +++ b/liveportrait/efficient/utils/__init__.py @@ -0,0 +1 @@ +from .utils import * \ No newline at end of file diff --git a/liveportrait/efficient/utils/onnx_driver.py b/liveportrait/efficient/utils/onnx_driver.py new file mode 100644 index 0000000..73e9463 --- /dev/null +++ b/liveportrait/efficient/utils/onnx_driver.py @@ -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 + } + diff --git a/liveportrait/efficient/utils/tensorrt_driver.py b/liveportrait/efficient/utils/tensorrt_driver.py new file mode 100644 index 0000000..2b3527e --- /dev/null +++ b/liveportrait/efficient/utils/tensorrt_driver.py @@ -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") diff --git a/liveportrait/efficient/utils/utils.py b/liveportrait/efficient/utils/utils.py new file mode 100644 index 0000000..66a1ea5 --- /dev/null +++ b/liveportrait/efficient/utils/utils.py @@ -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 diff --git a/liveportrait/live_portrait_pipeline.py b/liveportrait/live_portrait_pipeline.py index cafc3d7..cb18be9 100644 --- a/liveportrait/live_portrait_pipeline.py +++ b/liveportrait/live_portrait_pipeline.py @@ -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, diff --git a/liveportrait/live_portrait_wrapper.py b/liveportrait/live_portrait_wrapper.py index 9752e23..2c9be40 100644 --- a/liveportrait/live_portrait_wrapper.py +++ b/liveportrait/live_portrait_wrapper.py @@ -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 = [] diff --git a/requirements-tensorrt.txt b/requirements-tensorrt.txt new file mode 100644 index 0000000..c146226 --- /dev/null +++ b/requirements-tensorrt.txt @@ -0,0 +1,8 @@ +pyyaml +numpy +opencv-python +onnxruntime-gpu +pykalman +tensorrt +pycuda +ctypes \ No newline at end of file