2 Commits
Author SHA1 Message Date
kijai 36e46cb5d3 Autodetect dtype, use tqdm progress bars 2024-07-09 11:46:12 +03:00
Jukka Seppänen 3508b80c8b skip autocast if not needed 2024-07-09 02:51:24 +03:00
6 changed files with 49 additions and 34 deletions
+1 -1
View File
@@ -263,7 +263,7 @@
"Node name for S&R": "DownloadAndLoadLivePortraitModels"
},
"widgets_values": [
"fp16"
"auto"
]
},
{
+2 -2
View File
@@ -7,7 +7,7 @@ Pipeline of LivePortrait
import cv2
import numpy as np
import os.path as osp
from rich.progress import track
from tqdm import tqdm
from .config.inference_config import InferenceConfig
@@ -103,7 +103,7 @@ class LivePortraitPipeline(object):
I_p_lst = []
R_d_0, x_d_0_info = None, None
pbar = comfy.utils.ProgressBar(n_frames)
for i in track(range(n_frames), description='Animating...', total=n_frames):
for i in tqdm(range(n_frames), desc='Animating...', total=n_frames):
#if is_video(args.driving_info):
# extract kp info by M
I_d_i = I_d_lst[i]
+20 -22
View File
@@ -13,6 +13,7 @@ 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 contextlib import nullcontext
from comfy.model_management import get_autocast_device
@@ -79,9 +80,8 @@ class LivePortraitWrapper(object):
""" get the appearance feature of the image by F
x: Bx3xHxW, normalized to 0~1
"""
with torch.no_grad():
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)
with torch.autocast(get_autocast_device(self.device_id), dtype=torch.float16) if self.cfg.flag_use_half_precision else nullcontext():
feature_3d = self.appearance_feature_extractor(x)
return feature_3d.float()
@@ -91,15 +91,14 @@ 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=get_autocast_device(self.device_id), dtype=torch.float16, enabled=self.cfg.flag_use_half_precision):
kp_info = self.motion_extractor(x)
with torch.autocast(get_autocast_device(self.device_id), dtype=torch.float16) if self.cfg.flag_use_half_precision else nullcontext():
kp_info = self.motion_extractor(x)
if self.cfg.flag_use_half_precision:
# float the dict
for k, v in kp_info.items():
if isinstance(v, torch.Tensor):
kp_info[k] = v.float()
if self.cfg.flag_use_half_precision:
# float the dict
for k, v in kp_info.items():
if isinstance(v, torch.Tensor):
kp_info[k] = v.float()
flag_refine_info: bool = kwargs.get('flag_refine_info', True)
if flag_refine_info:
@@ -266,18 +265,17 @@ class LivePortraitWrapper(object):
kp_driving: BxNx3
"""
# The line 18 in Algorithm 1: D(W(f_s; x_s, x′_d,i))
with torch.no_grad():
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
ret_dct['out'] = self.spade_generator(feature=ret_dct['out'])
with torch.autocast(get_autocast_device(self.device_id), dtype=torch.float16) if self.cfg.flag_use_half_precision else nullcontext():
# get decoder input
ret_dct = self.warping_module(feature_3d, kp_source=kp_source, kp_driving=kp_driving)
# decode
ret_dct['out'] = self.spade_generator(feature=ret_dct['out'])
# float the dict
if self.cfg.flag_use_half_precision:
for k, v in ret_dct.items():
if isinstance(v, torch.Tensor):
ret_dct[k] = v.float()
# float the dict
if self.cfg.flag_use_half_precision:
for k, v in ret_dct.items():
if isinstance(v, torch.Tensor):
ret_dct[k] = v.float()
return ret_dct
+2 -2
View File
@@ -8,7 +8,7 @@ import os
import cv2
import numpy as np
import pickle
from rich.progress import track
from tqdm import tqdm
from .utils.cropper import Cropper
from .utils.io import load_driving_info
@@ -41,7 +41,7 @@ class TemplateMaker:
templates = []
for i in track(range(n_frames), description='Making templates...', total=n_frames):
for i in tqdm(range(n_frames), desc='Making templates...', total=n_frames):
I_d_i = I_d_lst[i]
x_d_i_info = self.live_portrait_wrapper.get_kp_info(I_d_i)
R_d_i = get_rotation_matrix(x_d_i_info['pitch'], x_d_i_info['yaw'], x_d_i_info['roll'])
+3 -3
View File
@@ -10,7 +10,7 @@ import subprocess
import imageio
import cv2
from rich.progress import track
from tqdm import tqdm
from .helper import prefix
from .rprint import rprint as print
@@ -35,7 +35,7 @@ def images2video(images, wfp, **kwargs):
)
n = len(images)
for i in track(range(n), description='writing', transient=True):
for i in tqdm(range(n), desc='writing', transient=True):
if image_mode.lower() == 'bgr':
writer.append_data(images[i][..., ::-1])
else:
@@ -83,7 +83,7 @@ def blend(img: np.ndarray, mask: np.ndarray, background_color=(255, 255, 255)):
def concat_frames(I_p_lst, driving_rgb_lst, img_rgb):
# TODO: add more concat style, e.g., left-down corner driving
out_lst = []
for idx, _ in track(enumerate(I_p_lst), total=len(I_p_lst), description='Concatenating result...'):
for idx, _ in tqdm(enumerate(I_p_lst), total=len(I_p_lst), desc='Concatenating result...'):
source_image_drived = I_p_lst[idx]
image_drive = driving_rgb_lst[idx]
+21 -4
View File
@@ -7,7 +7,6 @@ import comfy.utils
script_directory = os.path.dirname(os.path.abspath(__file__))
from .liveportrait.config.argument_config import ArgumentConfig
from .liveportrait.live_portrait_pipeline import LivePortraitPipeline
from .liveportrait.utils.cropper import Cropper
from .liveportrait.modules.spade_generator import SPADEDecoder
@@ -98,10 +97,11 @@ class DownloadAndLoadLivePortraitModels:
"optional": {
"precision": (
[
'auto',
'fp16',
'fp32',
], {
"default": 'fp16'
"default": 'auto'
}),
}
}
@@ -111,10 +111,27 @@ class DownloadAndLoadLivePortraitModels:
FUNCTION = "loadmodel"
CATEGORY = "LivePortrait"
def loadmodel(self, precision='fp16'):
def loadmodel(self, precision='auto'):
device = mm.get_torch_device()
mm.soft_empty_cache()
if precision == 'auto':
try:
if mm.is_device_mps(device):
print("LivePortrait using fp32 for MPS")
dtype = 'fp32'
elif mm.should_use_fp16():
print("LivePortrait using fp16")
dtype = 'fp16'
else:
print("LivePortrait using fp32")
dtype = 'fp32'
except:
raise AttributeError("ComfyUI version too old, can't autodetect properly. Set your dtypes manually.")
else:
dtype = precision
print(f"LivePortrait using {dtype}")
pbar = comfy.utils.ProgressBar(3)
download_path = os.path.join(folder_paths.models_dir, "liveportrait")
@@ -211,7 +228,7 @@ class DownloadAndLoadLivePortraitModels:
self.stich_retargeting_module,
InferenceConfig(
device_id=device,
flag_use_half_precision = True if precision == 'fp16' else False
flag_use_half_precision = True if dtype == 'fp16' else False
)
)