less cuda reliant
This commit is contained in:
@@ -45,7 +45,7 @@ class LivePortraitPipeline(object):
|
||||
#log(f"Load source image from {args.source_image}")
|
||||
crop_info = self.cropper.crop_single_image(img_rgb)
|
||||
source_lmk = crop_info['lmk_crop']
|
||||
img_crop, img_crop_256x256 = crop_info['img_crop'], crop_info['img_crop_256x256']
|
||||
_, img_crop_256x256 = crop_info['img_crop'], crop_info['img_crop_256x256']
|
||||
if inference_cfg.flag_do_crop:
|
||||
I_s = self.live_portrait_wrapper.prepare_source(img_crop_256x256)
|
||||
else:
|
||||
|
||||
@@ -3,21 +3,18 @@
|
||||
"""
|
||||
Wrapper for LivePortrait core functions
|
||||
"""
|
||||
|
||||
import os.path as osp
|
||||
import numpy as np
|
||||
import cv2
|
||||
import torch
|
||||
import yaml
|
||||
|
||||
from .utils.timer import Timer
|
||||
from .utils.helper import load_model, concat_feat
|
||||
from .utils.helper import concat_feat
|
||||
from .utils.retargeting_utils import compute_eye_delta, compute_lip_delta
|
||||
from .utils.camera import headpose_pred_to_degree, get_rotation_matrix
|
||||
from .utils.retargeting_utils import calc_eye_close_ratio, calc_lip_close_ratio
|
||||
from .config.inference_config import InferenceConfig
|
||||
from .utils.rprint import rlog as log
|
||||
|
||||
from comfy.model_management import get_autocast_device
|
||||
|
||||
class LivePortraitWrapper(object):
|
||||
|
||||
@@ -57,7 +54,7 @@ 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)
|
||||
x = x.to(self.device_id)
|
||||
return x
|
||||
|
||||
def prepare_driving_videos(self, imgs) -> torch.Tensor:
|
||||
@@ -74,7 +71,7 @@ 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)
|
||||
y = y.to(self.device_id)
|
||||
|
||||
return y
|
||||
|
||||
@@ -83,7 +80,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=get_autocast_device(self.device_id), dtype=torch.float16, enabled=self.cfg.flag_use_half_precision):
|
||||
feature_3d = self.appearance_feature_extractor(x)
|
||||
|
||||
return feature_3d.float()
|
||||
@@ -95,7 +92,7 @@ class LivePortraitWrapper(object):
|
||||
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.autocast(device_type=get_autocast_device(self.device_id), dtype=torch.float16, enabled=self.cfg.flag_use_half_precision):
|
||||
kp_info = self.motion_extractor(x)
|
||||
|
||||
if self.cfg.flag_use_half_precision:
|
||||
@@ -270,7 +267,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=get_autocast_device(self.device_id), 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
|
||||
@@ -306,17 +303,17 @@ 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().to(self.device_id)
|
||||
input_eye_ratio_tensor = torch.Tensor([input_eye_ratio[0][0]]).reshape(1, 1).to(self.device_id)
|
||||
# [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().to(self.device_id)
|
||||
# [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]]).to(self.device_id)
|
||||
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)
|
||||
|
||||
@@ -59,7 +59,10 @@ 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)
|
||||
try:
|
||||
zeros = torch.zeros(heatmap.shape[0], 1, spatial_size[0], spatial_size[1], spatial_size[2]).type(heatmap.type()).to(heatmap.device)
|
||||
except:
|
||||
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)
|
||||
return heatmap
|
||||
|
||||
@@ -8,17 +8,8 @@ import os
|
||||
import os.path as osp
|
||||
import cv2
|
||||
import torch
|
||||
from rich.console import Console
|
||||
from collections import OrderedDict
|
||||
|
||||
from ..modules.spade_generator import SPADEDecoder
|
||||
from ..modules.warping_network import WarpingNetwork
|
||||
from ..modules.motion_extractor import MotionExtractor
|
||||
from ..modules.appearance_feature_extractor import AppearanceFeatureExtractor
|
||||
from ..modules.stitching_retargeting_network import StitchingRetargetingNetwork
|
||||
from .rprint import rlog as log
|
||||
|
||||
|
||||
def suffix(filename):
|
||||
"""a.jpg -> jpg"""
|
||||
pos = filename.rfind(".")
|
||||
@@ -67,7 +58,7 @@ 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)
|
||||
dct[key] = torch.tensor(dct[key]).to(device_id)
|
||||
return dct
|
||||
|
||||
|
||||
@@ -91,51 +82,6 @@ def remove_ddp_dumplicate_key(state_dict):
|
||||
state_dict_new[key.replace('module.', '')] = state_dict[key]
|
||||
return state_dict_new
|
||||
|
||||
|
||||
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)
|
||||
elif model_type == 'motion_extractor':
|
||||
model = MotionExtractor(**model_params).cuda(device)
|
||||
elif model_type == 'warping_module':
|
||||
model = WarpingNetwork(**model_params).cuda(device)
|
||||
elif model_type == 'spade_generator':
|
||||
model = SPADEDecoder(**model_params).cuda(device)
|
||||
elif model_type == 'stitching_retargeting_module':
|
||||
# Special handling for stitching and retargeting module
|
||||
config = model_config['model_params']['stitching_retargeting_module_params']
|
||||
checkpoint = torch.load(ckpt_path, map_location=lambda storage, loc: storage)
|
||||
|
||||
stitcher = StitchingRetargetingNetwork(**config.get('stitching'))
|
||||
stitcher.load_state_dict(remove_ddp_dumplicate_key(checkpoint['retarget_shoulder']))
|
||||
stitcher = stitcher.cuda(device)
|
||||
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)
|
||||
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)
|
||||
retargetor_eye.eval()
|
||||
|
||||
return {
|
||||
'stitching': stitcher,
|
||||
'lip': retargetor_lip,
|
||||
'eye': retargetor_eye
|
||||
}
|
||||
else:
|
||||
raise ValueError(f"Unknown model type: {model_type}")
|
||||
|
||||
model.load_state_dict(torch.load(ckpt_path, map_location=lambda storage, loc: storage))
|
||||
model.eval()
|
||||
return model
|
||||
|
||||
|
||||
# get coefficients of Eqn. 7
|
||||
def calculate_transformation(config, s_kp_info, t_0_kp_info, t_i_kp_info, R_s, R_t_0, R_t_i):
|
||||
if config.relative:
|
||||
|
||||
@@ -58,8 +58,8 @@ 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().to(portrait_wrapper.device_id)
|
||||
input_eye_ratio_tensor = torch.Tensor([input_eye_ratio]).reshape(1, 1).to(portrait_wrapper.device_id)
|
||||
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 +69,8 @@ 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().to(portrait_wrapper.device_id)
|
||||
input_lip_ratio_tensor = torch.Tensor([input_lip_ratio]).to(portrait_wrapper.device_id)
|
||||
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
|
||||
|
||||
Reference in New Issue
Block a user