From 38a79c6b4ce4acc26c08273272dc9433461b7335 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E8=83=A1=E5=85=B6=E6=96=8C?= Date: Tue, 30 Jul 2024 14:31:03 +0800 Subject: [PATCH] =?UTF-8?q?=E5=85=BC=E5=AE=B9cuda\mps\cpu=E5=B9=B3?= =?UTF-8?q?=E5=8F=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- nodes/LivePortrait/deviceutils.py | 18 ++++ nodes/LivePortrait/speed.py | 83 ++++++++++++++++--- .../LivePortrait/src/live_portrait_wrapper.py | 48 ++++++++--- .../LivePortrait/src/modules/dense_motion.py | 5 +- nodes/LivePortrait/src/utils/cropper.py | 11 ++- nodes/LivePortrait/src/utils/helper.py | 39 +++++++-- .../LivePortrait/src/utils/landmark_runner.py | 8 +- .../src/utils/retargeting_utils.py | 23 ++++- 8 files changed, 193 insertions(+), 42 deletions(-) create mode 100644 nodes/LivePortrait/deviceutils.py diff --git a/nodes/LivePortrait/deviceutils.py b/nodes/LivePortrait/deviceutils.py new file mode 100644 index 0000000..fb9ea6c --- /dev/null +++ b/nodes/LivePortrait/deviceutils.py @@ -0,0 +1,18 @@ +import torch + +device_name = ( + "cuda" + if torch.cuda.is_available() + else "mps" + if torch.backends.mps.is_available() + else "cpu" +) +print(f'@@device:{device_name}') + +device_providers = ( + ["CUDAExecutionProvider", "CPUExecutionProvider"] + if torch.cuda.is_available() + else ["MPSExecutionProvider", "CPUExecutionProvider"] + if torch.backends.mps.is_available() + else ["CPUExecutionProvider"] +) \ No newline at end of file diff --git a/nodes/LivePortrait/speed.py b/nodes/LivePortrait/speed.py index a13651e..74211c8 100644 --- a/nodes/LivePortrait/speed.py +++ b/nodes/LivePortrait/speed.py @@ -12,19 +12,45 @@ import time import numpy as np from LivePortrait.src.utils.helper import load_model, concat_feat from LivePortrait.src.config.inference_config import InferenceConfig +import deviceutils def initialize_inputs(batch_size=1): """ Generate random input tensors and move them to GPU """ - feature_3d = torch.randn(batch_size, 32, 16, 64, 64).cuda().half() - kp_source = torch.randn(batch_size, 21, 3).cuda().half() - kp_driving = torch.randn(batch_size, 21, 3).cuda().half() - source_image = torch.randn(batch_size, 3, 256, 256).cuda().half() - generator_input = torch.randn(batch_size, 256, 64, 64).cuda().half() - eye_close_ratio = torch.randn(batch_size, 3).cuda().half() - lip_close_ratio = torch.randn(batch_size, 2).cuda().half() + + if deviceutils.device_name != "cuda": + feature_3d = torch.randn(batch_size, 32, 16, 64, 64).cuda().half() + kp_source = torch.randn(batch_size, 21, 3).cuda().half() + kp_driving = torch.randn(batch_size, 21, 3).cuda().half() + source_image = torch.randn(batch_size, 3, 256, 256).cuda().half() + generator_input = torch.randn(batch_size, 256, 64, 64).cuda().half() + eye_close_ratio = torch.randn(batch_size, 3).cuda().half() + lip_close_ratio = torch.randn(batch_size, 2).cuda().half() + elif deviceutils.device_name == "mps": + feature_3d = torch.randn(batch_size, 32, 16, 64, 64).half() + feature_3d = feature_3d.to(deviceutils.device_name) + kp_source = torch.randn(batch_size, 21, 3).half() + kp_source = kp_source.to(deviceutils.device_name) + kp_driving = torch.randn(batch_size, 21, 3).half() + kp_driving = kp_driving.to(deviceutils.device_name) + source_image = torch.randn(batch_size, 3, 256, 256).half() + source_image = source_image.to(deviceutils.device_name) + generator_input = torch.randn(batch_size, 256, 64, 64).half() + generator_input = generator_input.to(deviceutils.device_name) + eye_close_ratio = torch.randn(batch_size, 3).half() + eye_close_ratio = eye_close_ratio.to(deviceutils.device_name) + lip_close_ratio = torch.randn(batch_size, 2).half() + lip_close_ratio = lip_close_ratio.to(deviceutils.device_name) + else: + feature_3d = torch.randn(batch_size, 32, 16, 64, 64).half() + kp_source = torch.randn(batch_size, 21, 3).half() + kp_driving = torch.randn(batch_size, 21, 3).half() + source_image = torch.randn(batch_size, 3, 256, 256).half() + generator_input = torch.randn(batch_size, 256, 64, 64).half() + eye_close_ratio = torch.randn(batch_size, 3).half() + lip_close_ratio = torch.randn(batch_size, 2).half() feat_stitching = concat_feat(kp_source, kp_driving).half() feat_eye = concat_feat(kp_source, eye_close_ratio).half() feat_lip = concat_feat(kp_source, lip_close_ratio).half() @@ -105,34 +131,65 @@ def measure_inference_times(compiled_models, stitching_retargeting_module, input with torch.no_grad(): for _ in range(100): - torch.cuda.synchronize() + if deviceutils.device_name == "cuda": + torch.cuda.synchronize() + elif deviceutils.device_name == "mps": + torch.mps.synchronize() + else: + torch.cpu.synchronize() + overall_start = time.time() start = time.time() compiled_models['Appearance Feature Extractor'](inputs['source_image']) - torch.cuda.synchronize() + if deviceutils.device_name == "cuda": + torch.cuda.synchronize() + elif deviceutils.device_name == "mps": + torch.mps.synchronize() + else: + torch.cpu.synchronize() times['Appearance Feature Extractor'].append(time.time() - start) start = time.time() compiled_models['Motion Extractor'](inputs['source_image']) - torch.cuda.synchronize() + if deviceutils.device_name == "cuda": + torch.cuda.synchronize() + elif deviceutils.device_name == "mps": + torch.mps.synchronize() + else: + torch.cpu.synchronize() times['Motion Extractor'].append(time.time() - start) start = time.time() compiled_models['Warping Network'](inputs['feature_3d'], inputs['kp_driving'], inputs['kp_source']) - torch.cuda.synchronize() + if deviceutils.device_name == "cuda": + torch.cuda.synchronize() + elif deviceutils.device_name == "mps": + torch.mps.synchronize() + else: + torch.cpu.synchronize() times['Warping Network'].append(time.time() - start) start = time.time() compiled_models['SPADE Decoder'](inputs['generator_input']) # Adjust input as required - torch.cuda.synchronize() + if deviceutils.device_name == "cuda": + torch.cuda.synchronize() + elif deviceutils.device_name == "mps": + torch.mps.synchronize() + else: + torch.cpu.synchronize() times['SPADE Decoder'].append(time.time() - start) start = time.time() stitching_retargeting_module['stitching'](inputs['feat_stitching']) stitching_retargeting_module['eye'](inputs['feat_eye']) stitching_retargeting_module['lip'](inputs['feat_lip']) - torch.cuda.synchronize() + if deviceutils.device_name == "cuda": + torch.cuda.synchronize() + elif deviceutils.device_name == "mps": + torch.mps.synchronize() + else: + torch.cpu.synchronize() times['Retargeting Models'].append(time.time() - start) overall_times.append(time.time() - overall_start) diff --git a/nodes/LivePortrait/src/live_portrait_wrapper.py b/nodes/LivePortrait/src/live_portrait_wrapper.py index 465cbe9..94ba7c7 100644 --- a/nodes/LivePortrait/src/live_portrait_wrapper.py +++ b/nodes/LivePortrait/src/live_portrait_wrapper.py @@ -18,6 +18,8 @@ from .utils.retargeting_utils import calc_eye_close_ratio, calc_lip_close_ratio from LivePortrait.src.config.inference_config import InferenceConfig from LivePortrait.src.utils.rprint import rlog as log +from .. import deviceutils + class LivePortraitWrapper(object): @@ -71,7 +73,12 @@ class LivePortraitWrapper(object): raise ValueError(f'img ndim should be 3 or 4: {x.ndim}') x = np.clip(x, 0, 1) # clip to 0~1 x = torch.from_numpy(x).permute(0, 3, 1, 2) # 1xHxWx3 -> 1x3xHxW - x = x.cuda(self.device_id) + if deviceutils.device_name == "cuda": + x = x.cuda(self.device_id) + elif deviceutils.device_name == "mps": + x = x.to('mps') + else: + x = x.to('cpu') return x def prepare_driving_videos(self, imgs) -> torch.Tensor: @@ -88,7 +95,14 @@ class LivePortraitWrapper(object): y = _imgs.astype(np.float32) / 255. y = np.clip(y, 0, 1) # clip to 0~1 y = torch.from_numpy(y).permute(0, 4, 3, 1, 2) # TxHxWx3x1 -> Tx1x3xHxW - y = y.cuda(self.device_id) + + if deviceutils.device_name == "cuda": + y = y.cuda(self.device_id) + elif deviceutils.device_name == "mps": + y = y.to('mps') + else: + y = y.to('cpu') + return y @@ -97,7 +111,7 @@ class LivePortraitWrapper(object): x: Bx3xHxW, normalized to 0~1 """ with torch.no_grad(): - with torch.autocast(device_type='cuda', dtype=torch.float16, enabled=self.cfg.flag_use_half_precision): + with torch.autocast(device_type="cuda" if deviceutils.device_name is "cuda" else "cpu", dtype=torch.float16, enabled=self.cfg.flag_use_half_precision): feature_3d = self.appearance_feature_extractor(x) return feature_3d.float() @@ -108,8 +122,8 @@ class LivePortraitWrapper(object): flag_refine_info: whether to trandform the pose to degrees and the dimention of the reshape return: A dict contains keys: 'pitch', 'yaw', 'roll', 't', 'exp', 'scale', 'kp' """ - with torch.no_grad(): - with torch.autocast(device_type='cuda', dtype=torch.float16, enabled=self.cfg.flag_use_half_precision): + with torch.no_grad():#mps不支持autocast + with torch.autocast(device_type="cuda" if deviceutils.device_name is "cuda" else "cpu", dtype=torch.float16, enabled=self.cfg.flag_use_half_precision): kp_info = self.motion_extractor(x) if self.cfg.flag_use_half_precision: @@ -282,7 +296,7 @@ class LivePortraitWrapper(object): """ # The line 18 in Algorithm 1: D(W(f_s; x_s, x′_d,i)) with torch.no_grad(): - with torch.autocast(device_type='cuda', dtype=torch.float16, enabled=self.cfg.flag_use_half_precision): + with torch.autocast(device_type="cuda" if deviceutils.device_name is "cuda" else "cpu", dtype=torch.float16, enabled=self.cfg.flag_use_half_precision): # get decoder input ret_dct = self.warping_module(feature_3d, kp_source=kp_source, kp_driving=kp_driving) # decode @@ -318,17 +332,31 @@ class LivePortraitWrapper(object): def calc_combined_eye_ratio(self, input_eye_ratio, source_lmk): eye_close_ratio = calc_eye_close_ratio(source_lmk[None]) - eye_close_ratio_tensor = torch.from_numpy(eye_close_ratio).float().cuda(self.device_id) - input_eye_ratio_tensor = torch.Tensor([input_eye_ratio[0][0]]).reshape(1, 1).cuda(self.device_id) + eye_close_ratio_tensor = torch.from_numpy(eye_close_ratio).float() + input_eye_ratio_tensor = torch.Tensor([input_eye_ratio[0][0]]).reshape(1, 1) + if deviceutils.device_name == "cuda": + eye_close_ratio_tensor = eye_close_ratio_tensor.to(f'cuda:{self.device_id}') + input_eye_ratio_tensor = input_eye_ratio_tensor.to(f'cuda:{self.device_id}') + else: + eye_close_ratio_tensor = eye_close_ratio_tensor.to(deviceutils.device_name) + input_eye_ratio_tensor = input_eye_ratio_tensor.to(deviceutils.device_name) + + # [c_s,eyes, c_d,eyes,i] combined_eye_ratio_tensor = torch.cat([eye_close_ratio_tensor, input_eye_ratio_tensor], dim=1) return combined_eye_ratio_tensor def calc_combined_lip_ratio(self, input_lip_ratio, source_lmk): lip_close_ratio = calc_lip_close_ratio(source_lmk[None]) - lip_close_ratio_tensor = torch.from_numpy(lip_close_ratio).float().cuda(self.device_id) + lip_close_ratio_tensor = torch.from_numpy(lip_close_ratio).float() # [c_s,lip, c_d,lip,i] - input_lip_ratio_tensor = torch.Tensor([input_lip_ratio[0]]).cuda(self.device_id) + input_lip_ratio_tensor = torch.Tensor([input_lip_ratio[0]]) + if deviceutils.device_name == "cuda": + lip_close_ratio_tensor = lip_close_ratio_tensor.to(f'cuda:{self.device_id}') + input_lip_ratio_tensor = input_lip_ratio_tensor.to(f'cuda:{self.device_id}') + else: + lip_close_ratio_tensor = lip_close_ratio_tensor.to(deviceutils.device_name) + input_lip_ratio_tensor = input_lip_ratio_tensor.to(deviceutils.device_name) if input_lip_ratio_tensor.shape != [1, 1]: input_lip_ratio_tensor = input_lip_ratio_tensor.reshape(1, 1) combined_lip_ratio_tensor = torch.cat([lip_close_ratio_tensor, input_lip_ratio_tensor], dim=1) diff --git a/nodes/LivePortrait/src/modules/dense_motion.py b/nodes/LivePortrait/src/modules/dense_motion.py index 0eec0c4..e234cb0 100644 --- a/nodes/LivePortrait/src/modules/dense_motion.py +++ b/nodes/LivePortrait/src/modules/dense_motion.py @@ -8,6 +8,7 @@ from torch import nn import torch.nn.functional as F import torch from .util import Hourglass, make_coordinate_grid, kp2gaussian +from ... import deviceutils class DenseMotionNetwork(nn.Module): @@ -59,7 +60,9 @@ class DenseMotionNetwork(nn.Module): heatmap = gaussian_driving - gaussian_source # (bs, num_kp, d, h, w) # adding background feature - zeros = torch.zeros(heatmap.shape[0], 1, spatial_size[0], spatial_size[1], spatial_size[2]).type(heatmap.type()).to(heatmap.device) + print(f'heatmap.type():{heatmap.type()},heatmap.device:{heatmap.device}') + ty = heatmap.type() if deviceutils.device_name is 'cuda' else 'torch.FloatTensor' + zeros = torch.zeros(heatmap.shape[0], 1, spatial_size[0], spatial_size[1], spatial_size[2]).type(ty).to(heatmap.device) heatmap = torch.cat([zeros, heatmap], dim=1) heatmap = heatmap.unsqueeze(2) # (bs, 1+num_kp, 1, d, h, w) return heatmap diff --git a/nodes/LivePortrait/src/utils/cropper.py b/nodes/LivePortrait/src/utils/cropper.py index 6d9d174..6fc32a1 100644 --- a/nodes/LivePortrait/src/utils/cropper.py +++ b/nodes/LivePortrait/src/utils/cropper.py @@ -5,7 +5,11 @@ from PIL import Image import os.path as osp from typing import List, Union, Tuple from dataclasses import dataclass, field -import cv2; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False) +import cv2; + +from ... import deviceutils + +cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False) from .landmark_runner import LandmarkRunner from .face_analysis_diy import FaceAnalysisDIY @@ -30,7 +34,6 @@ class Trajectory: frame_rgb_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # frame list frame_rgb_crop_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # frame crop list - class Cropper(object): def __init__(self, **kwargs) -> None: device_id = kwargs.get('device_id', 0) @@ -41,7 +44,7 @@ class Cropper(object): self.landmark_runner = LandmarkRunner( # ckpt_path=make_abs_path('../../pretrained_weights/liveportrait/landmark.onnx'), ckpt_path=landmark_runner_ckpt, - onnx_provider='cuda', + onnx_provider=deviceutils.device_name, device_id=device_id ) self.landmark_runner.warmup() @@ -50,7 +53,7 @@ class Cropper(object): name='buffalo_l', # root=make_abs_path('../../pretrained_weights/insightface'), root=insightface_pretrained_weights, - providers=["CUDAExecutionProvider"] + providers=deviceutils.device_providers ) self.face_analysis_wrapper.prepare(ctx_id=device_id, det_size=(512, 512)) self.face_analysis_wrapper.warmup() diff --git a/nodes/LivePortrait/src/utils/helper.py b/nodes/LivePortrait/src/utils/helper.py index 009c3df..cad0edc 100644 --- a/nodes/LivePortrait/src/utils/helper.py +++ b/nodes/LivePortrait/src/utils/helper.py @@ -17,6 +17,7 @@ from LivePortrait.src.modules.motion_extractor import MotionExtractor from LivePortrait.src.modules.appearance_feature_extractor import AppearanceFeatureExtractor from LivePortrait.src.modules.stitching_retargeting_network import StitchingRetargetingNetwork from .rprint import rlog as log +from ... import deviceutils def suffix(filename): @@ -67,7 +68,15 @@ def squeeze_tensor_to_numpy(tensor): def dct2cuda(dct: dict, device_id: int): for key in dct: - dct[key] = torch.tensor(dct[key]).cuda(device_id) + tensor = torch.tensor(dct[key]) + if deviceutils.device_name == "cuda": + tensor = tensor.to(f'cuda:{device_id}') + elif deviceutils.device_name == "mps": + tensor = tensor.to('mps') + else: + tensor = tensor.to('cpu') + dct[key] = tensor + # dct[key] = torch.tensor(dct[key]).cuda(device_id) return dct @@ -96,13 +105,13 @@ def load_model(ckpt_path, model_config, device, model_type): model_params = model_config['model_params'][f'{model_type}_params'] if model_type == 'appearance_feature_extractor': - model = AppearanceFeatureExtractor(**model_params).cuda(device) + model = AppearanceFeatureExtractor(**model_params) elif model_type == 'motion_extractor': - model = MotionExtractor(**model_params).cuda(device) + model = MotionExtractor(**model_params) elif model_type == 'warping_module': - model = WarpingNetwork(**model_params).cuda(device) + model = WarpingNetwork(**model_params) elif model_type == 'spade_generator': - model = SPADEDecoder(**model_params).cuda(device) + model = SPADEDecoder(**model_params) elif model_type == 'stitching_retargeting_module': # Special handling for stitching and retargeting module config = model_config['model_params']['stitching_retargeting_module_params'] @@ -110,17 +119,26 @@ def load_model(ckpt_path, model_config, device, model_type): stitcher = StitchingRetargetingNetwork(**config.get('stitching')) stitcher.load_state_dict(remove_ddp_dumplicate_key(checkpoint['retarget_shoulder'])) - stitcher = stitcher.cuda(device) + if deviceutils.device_name == "cuda": + stitcher = stitcher.cuda(device) + else: + stitcher = stitcher.to(deviceutils.device_name) stitcher.eval() retargetor_lip = StitchingRetargetingNetwork(**config.get('lip')) retargetor_lip.load_state_dict(remove_ddp_dumplicate_key(checkpoint['retarget_mouth'])) - retargetor_lip = retargetor_lip.cuda(device) + if deviceutils.device_name == "cuda": + retargetor_lip = retargetor_lip.cuda(device) + else: + retargetor_lip = retargetor_lip.to(deviceutils.device_name) retargetor_lip.eval() retargetor_eye = StitchingRetargetingNetwork(**config.get('eye')) retargetor_eye.load_state_dict(remove_ddp_dumplicate_key(checkpoint['retarget_eye'])) - retargetor_eye = retargetor_eye.cuda(device) + if deviceutils.device_name == "cuda": + retargetor_eye = retargetor_eye.cuda(device) + else: + retargetor_eye = retargetor_eye.to(deviceutils.device_name) retargetor_eye.eval() return { @@ -131,6 +149,11 @@ def load_model(ckpt_path, model_config, device, model_type): else: raise ValueError(f"Unknown model type: {model_type}") + if deviceutils.device_name == "cuda": + model = model.to(f'cuda:{device}') + else: + model = model.to(deviceutils.device_name) + model.load_state_dict(torch.load(ckpt_path, map_location=lambda storage, loc: storage)) model.eval() return model diff --git a/nodes/LivePortrait/src/utils/landmark_runner.py b/nodes/LivePortrait/src/utils/landmark_runner.py index 419eff1..5a1894d 100644 --- a/nodes/LivePortrait/src/utils/landmark_runner.py +++ b/nodes/LivePortrait/src/utils/landmark_runner.py @@ -1,7 +1,11 @@ # coding: utf-8 import os.path as osp -import cv2; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False) +import cv2; + +from ... import deviceutils + +cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False) import torch import numpy as np import onnxruntime @@ -32,7 +36,7 @@ class LandmarkRunner(object): self.dsize = kwargs.get('dsize', 224) self.timer = Timer() # print('---------------------------------#onnx_provider',ckpt_path) - if onnx_provider.lower() == 'cuda': + if deviceutils.device_name == 'cuda': self.session = onnxruntime.InferenceSession( ckpt_path, providers=[ ('CUDAExecutionProvider', {'device_id': device_id}) diff --git a/nodes/LivePortrait/src/utils/retargeting_utils.py b/nodes/LivePortrait/src/utils/retargeting_utils.py index 2028590..90a5975 100644 --- a/nodes/LivePortrait/src/utils/retargeting_utils.py +++ b/nodes/LivePortrait/src/utils/retargeting_utils.py @@ -6,6 +6,8 @@ Functions to compute distance ratios between specific pairs of facial landmarks import numpy as np import torch +from LivePortrait import deviceutils + def calculate_distance_ratio(lmk: np.ndarray, idx1: int, idx2: int, idx3: int, idx4: int, eps: float = 1e-6) -> np.ndarray: """ @@ -58,8 +60,15 @@ def calc_lip_close_ratio(lmk: np.ndarray) -> np.ndarray: def compute_eye_delta(frame_idx, input_eye_ratios, source_landmarks, portrait_wrapper, kp_source): input_eye_ratio = input_eye_ratios[frame_idx][0][0] eye_close_ratio = calc_eye_close_ratio(source_landmarks[None]) - eye_close_ratio_tensor = torch.from_numpy(eye_close_ratio).float().cuda(portrait_wrapper.device_id) - input_eye_ratio_tensor = torch.Tensor([input_eye_ratio]).reshape(1, 1).cuda(portrait_wrapper.device_id) + eye_close_ratio_tensor = torch.from_numpy(eye_close_ratio).float() + input_eye_ratio_tensor = torch.Tensor([input_eye_ratio]).reshape(1, 1) + if deviceutils.device_name == "cuda": + eye_close_ratio_tensor = eye_close_ratio_tensor.to(f'cuda:{portrait_wrapper.device_id}') + input_eye_ratio_tensor = input_eye_ratio_tensor.to(f'cuda:{portrait_wrapper.device_id}') + else: + eye_close_ratio_tensor = eye_close_ratio_tensor.to(deviceutils.device_name) + input_eye_ratio_tensor = input_eye_ratio_tensor.to(deviceutils.device_name) + combined_eye_ratio_tensor = torch.cat([eye_close_ratio_tensor, input_eye_ratio_tensor], dim=1) # print(combined_eye_ratio_tensor.mean()) eye_delta = portrait_wrapper.retarget_eye(kp_source, combined_eye_ratio_tensor) @@ -69,8 +78,14 @@ def compute_eye_delta(frame_idx, input_eye_ratios, source_landmarks, portrait_wr def compute_lip_delta(frame_idx, input_lip_ratios, source_landmarks, portrait_wrapper, kp_source): input_lip_ratio = input_lip_ratios[frame_idx][0] lip_close_ratio = calc_lip_close_ratio(source_landmarks[None]) - lip_close_ratio_tensor = torch.from_numpy(lip_close_ratio).float().cuda(portrait_wrapper.device_id) - input_lip_ratio_tensor = torch.Tensor([input_lip_ratio]).cuda(portrait_wrapper.device_id) + lip_close_ratio_tensor = torch.from_numpy(lip_close_ratio).float() + input_lip_ratio_tensor = torch.Tensor([input_lip_ratio]) + if deviceutils.device_name == "cuda": + lip_close_ratio_tensor = lip_close_ratio_tensor.to(f'cuda:{portrait_wrapper.device_id}') + input_lip_ratio_tensor = input_lip_ratio_tensor.to(f'cuda:{portrait_wrapper.device_id}') + else: + lip_close_ratio_tensor = lip_close_ratio_tensor.to(deviceutils.device_name) + input_lip_ratio_tensor = input_lip_ratio_tensor.to(deviceutils.device_name) combined_lip_ratio_tensor = torch.cat([lip_close_ratio_tensor, input_lip_ratio_tensor], dim=1) lip_delta = portrait_wrapper.retarget_lip(kp_source, combined_lip_ratio_tensor) return lip_delta