Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
36e46cb5d3 | ||
|
|
3508b80c8b |
@@ -263,7 +263,7 @@
|
|||||||
"Node name for S&R": "DownloadAndLoadLivePortraitModels"
|
"Node name for S&R": "DownloadAndLoadLivePortraitModels"
|
||||||
},
|
},
|
||||||
"widgets_values": [
|
"widgets_values": [
|
||||||
"fp16"
|
"auto"
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -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]
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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'])
|
||||||
|
|||||||
@@ -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]
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user