From 36e46cb5d30599cdfb6873277c6193643a2d3077 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 9 Jul 2024 11:46:12 +0300 Subject: [PATCH] Autodetect dtype, use tqdm progress bars --- examples/liveportrait_example_01.json | 2 +- liveportrait/live_portrait_pipeline.py | 4 ++-- liveportrait/template_maker.py | 4 ++-- liveportrait/utils/video.py | 6 +++--- nodes.py | 25 +++++++++++++++++++++---- 5 files changed, 29 insertions(+), 12 deletions(-) diff --git a/examples/liveportrait_example_01.json b/examples/liveportrait_example_01.json index def7705..40e3239 100644 --- a/examples/liveportrait_example_01.json +++ b/examples/liveportrait_example_01.json @@ -263,7 +263,7 @@ "Node name for S&R": "DownloadAndLoadLivePortraitModels" }, "widgets_values": [ - "fp16" + "auto" ] }, { diff --git a/liveportrait/live_portrait_pipeline.py b/liveportrait/live_portrait_pipeline.py index 3f544cd..06b5e71 100644 --- a/liveportrait/live_portrait_pipeline.py +++ b/liveportrait/live_portrait_pipeline.py @@ -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] diff --git a/liveportrait/template_maker.py b/liveportrait/template_maker.py index 7f3ce06..8d21bf5 100644 --- a/liveportrait/template_maker.py +++ b/liveportrait/template_maker.py @@ -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']) diff --git a/liveportrait/utils/video.py b/liveportrait/utils/video.py index 720e082..2df7aee 100644 --- a/liveportrait/utils/video.py +++ b/liveportrait/utils/video.py @@ -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] diff --git a/nodes.py b/nodes.py index b9bf637..29d718e 100644 --- a/nodes.py +++ b/nodes.py @@ -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 ) )