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" "Node name for S&R": "DownloadAndLoadLivePortraitModels"
}, },
"widgets_values": [ "widgets_values": [
"fp16" "auto"
] ]
}, },
{ {
+2 -2
View File
@@ -7,7 +7,7 @@ Pipeline of LivePortrait
import cv2 import cv2
import numpy as np import numpy as np
import os.path as osp import os.path as osp
from rich.progress import track from tqdm import tqdm
from .config.inference_config import InferenceConfig from .config.inference_config import InferenceConfig
@@ -103,7 +103,7 @@ class LivePortraitPipeline(object):
I_p_lst = [] I_p_lst = []
R_d_0, x_d_0_info = None, None R_d_0, x_d_0_info = None, None
pbar = comfy.utils.ProgressBar(n_frames) 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): #if is_video(args.driving_info):
# extract kp info by M # extract kp info by M
I_d_i = I_d_lst[i] 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.camera import headpose_pred_to_degree, get_rotation_matrix
from .utils.retargeting_utils import calc_eye_close_ratio, calc_lip_close_ratio from .utils.retargeting_utils import calc_eye_close_ratio, calc_lip_close_ratio
from .config.inference_config import InferenceConfig from .config.inference_config import InferenceConfig
from contextlib import nullcontext
from comfy.model_management import get_autocast_device from comfy.model_management import get_autocast_device
@@ -79,9 +80,8 @@ class LivePortraitWrapper(object):
""" get the appearance feature of the image by F """ get the appearance feature of the image by F
x: Bx3xHxW, normalized to 0~1 x: Bx3xHxW, normalized to 0~1
""" """
with torch.no_grad(): with torch.autocast(get_autocast_device(self.device_id), dtype=torch.float16) if self.cfg.flag_use_half_precision else nullcontext():
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)
feature_3d = self.appearance_feature_extractor(x)
return feature_3d.float() 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 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' return: A dict contains keys: 'pitch', 'yaw', 'roll', 't', 'exp', 'scale', 'kp'
""" """
with torch.no_grad(): with torch.autocast(get_autocast_device(self.device_id), dtype=torch.float16) if self.cfg.flag_use_half_precision else nullcontext():
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)
kp_info = self.motion_extractor(x)
if self.cfg.flag_use_half_precision: if self.cfg.flag_use_half_precision:
# float the dict # float the dict
for k, v in kp_info.items(): for k, v in kp_info.items():
if isinstance(v, torch.Tensor): if isinstance(v, torch.Tensor):
kp_info[k] = v.float() kp_info[k] = v.float()
flag_refine_info: bool = kwargs.get('flag_refine_info', True) flag_refine_info: bool = kwargs.get('flag_refine_info', True)
if flag_refine_info: if flag_refine_info:
@@ -266,18 +265,17 @@ class LivePortraitWrapper(object):
kp_driving: BxNx3 kp_driving: BxNx3
""" """
# The line 18 in Algorithm 1: D(W(f_s; x_s, x′_d,i)) # The line 18 in Algorithm 1: D(W(f_s; x_s, x′_d,i))
with torch.no_grad(): with torch.autocast(get_autocast_device(self.device_id), dtype=torch.float16) if self.cfg.flag_use_half_precision else nullcontext():
with torch.autocast(device_type=get_autocast_device(self.device_id), dtype=torch.float16, enabled=self.cfg.flag_use_half_precision): # get decoder input
# get decoder input ret_dct = self.warping_module(feature_3d, kp_source=kp_source, kp_driving=kp_driving)
ret_dct = self.warping_module(feature_3d, kp_source=kp_source, kp_driving=kp_driving) # decode
# decode ret_dct['out'] = self.spade_generator(feature=ret_dct['out'])
ret_dct['out'] = self.spade_generator(feature=ret_dct['out'])
# float the dict # float the dict
if self.cfg.flag_use_half_precision: if self.cfg.flag_use_half_precision:
for k, v in ret_dct.items(): for k, v in ret_dct.items():
if isinstance(v, torch.Tensor): if isinstance(v, torch.Tensor):
ret_dct[k] = v.float() ret_dct[k] = v.float()
return ret_dct return ret_dct
+2 -2
View File
@@ -8,7 +8,7 @@ import os
import cv2 import cv2
import numpy as np import numpy as np
import pickle import pickle
from rich.progress import track from tqdm import tqdm
from .utils.cropper import Cropper from .utils.cropper import Cropper
from .utils.io import load_driving_info from .utils.io import load_driving_info
@@ -41,7 +41,7 @@ class TemplateMaker:
templates = [] 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] I_d_i = I_d_lst[i]
x_d_i_info = self.live_portrait_wrapper.get_kp_info(I_d_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']) 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 imageio
import cv2 import cv2
from rich.progress import track from tqdm import tqdm
from .helper import prefix from .helper import prefix
from .rprint import rprint as print from .rprint import rprint as print
@@ -35,7 +35,7 @@ def images2video(images, wfp, **kwargs):
) )
n = len(images) 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': if image_mode.lower() == 'bgr':
writer.append_data(images[i][..., ::-1]) writer.append_data(images[i][..., ::-1])
else: 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): def concat_frames(I_p_lst, driving_rgb_lst, img_rgb):
# TODO: add more concat style, e.g., left-down corner driving # TODO: add more concat style, e.g., left-down corner driving
out_lst = [] 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] source_image_drived = I_p_lst[idx]
image_drive = driving_rgb_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__)) 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.live_portrait_pipeline import LivePortraitPipeline
from .liveportrait.utils.cropper import Cropper from .liveportrait.utils.cropper import Cropper
from .liveportrait.modules.spade_generator import SPADEDecoder from .liveportrait.modules.spade_generator import SPADEDecoder
@@ -98,10 +97,11 @@ class DownloadAndLoadLivePortraitModels:
"optional": { "optional": {
"precision": ( "precision": (
[ [
'auto',
'fp16', 'fp16',
'fp32', 'fp32',
], { ], {
"default": 'fp16' "default": 'auto'
}), }),
} }
} }
@@ -111,10 +111,27 @@ class DownloadAndLoadLivePortraitModels:
FUNCTION = "loadmodel" FUNCTION = "loadmodel"
CATEGORY = "LivePortrait" CATEGORY = "LivePortrait"
def loadmodel(self, precision='fp16'): def loadmodel(self, precision='auto'):
device = mm.get_torch_device() device = mm.get_torch_device()
mm.soft_empty_cache() 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) pbar = comfy.utils.ProgressBar(3)
download_path = os.path.join(folder_paths.models_dir, "liveportrait") download_path = os.path.join(folder_paths.models_dir, "liveportrait")
@@ -211,7 +228,7 @@ class DownloadAndLoadLivePortraitModels:
self.stich_retargeting_module, self.stich_retargeting_module,
InferenceConfig( InferenceConfig(
device_id=device, device_id=device,
flag_use_half_precision = True if precision == 'fp16' else False flag_use_half_precision = True if dtype == 'fp16' else False
) )
) )