Autodetect dtype, use tqdm progress bars
This commit is contained in:
@@ -263,7 +263,7 @@
|
||||
"Node name for S&R": "DownloadAndLoadLivePortraitModels"
|
||||
},
|
||||
"widgets_values": [
|
||||
"fp16"
|
||||
"auto"
|
||||
]
|
||||
},
|
||||
{
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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'])
|
||||
|
||||
@@ -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]
|
||||
|
||||
|
||||
@@ -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
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user