remove reference frame attempt

also use batches for driving when either retargetting is enabled
This commit is contained in:
Mel Massadian
2024-07-08 21:22:50 +02:00
parent eb5fddf4de
commit 68d0ddf72a
+99 -23
View File
@@ -73,24 +73,17 @@ class LivePortraitPipeline(object):
self, source_np, driving_images_np, mismatch_method="repeat", reference_frame=0
):
inference_cfg = self.live_portrait_wrapper.cfg
is_video = source_np.shape[0] > 1
I_p_lst = []
I_p_paste_lst = []
driving_lmk_lst = []
total_frames = driving_images_np.shape[0]
pbar = comfy.utils.ProgressBar(total_frames)
ref_frame = self._get_source_frame(
source_np, reference_frame, total_frames, mismatch_method
)
rcrop_info = self.cropper.crop_single_image(ref_frame)
rsource_lmk = rcrop_info["lmk_crop"]
rimg_crop, ref_crop_256x256 = (
rcrop_info["img_crop"],
rcrop_info["img_crop_256x256"],
)
if inference_cfg.flag_eye_retargeting or inference_cfg.flag_lip_retargeting:
driving_lmk_lst = self.cropper.get_retargeting_lmk_info(driving_images_np)
for i in range(total_frames):
source_frame_rgb = self._get_source_frame(
@@ -100,7 +93,7 @@ class LivePortraitPipeline(object):
crop_info = self.cropper.crop_single_image(source_frame_rgb)
source_lmk = crop_info["lmk_crop"]
img_crop, img_crop_256x256 = (
_, img_crop_256x256 = (
crop_info["img_crop"],
crop_info["img_crop_256x256"],
)
@@ -110,9 +103,6 @@ class LivePortraitPipeline(object):
else:
I_s = self.live_portrait_wrapper.prepare_source(source_frame_rgb)
rel_s_info = self.live_portrait_wrapper.get_kp_info(
self.live_portrait_wrapper.prepare_source(ref_crop_256x256)
)
x_s_info = self.live_portrait_wrapper.get_kp_info(I_s)
x_c_s = x_s_info["kp"]
R_s = get_rotation_matrix(
@@ -148,7 +138,7 @@ class LivePortraitPipeline(object):
)[0]
if inference_cfg.flag_eye_retargeting or inference_cfg.flag_lip_retargeting:
driving_lmk_lst = self.cropper.get_retargeting_lmk_info([driving_frame])
# driving_lmk_lst = self.cropper.get_retargeting_lmk_info([driving_frame])
input_eye_ratio_lst, input_lip_ratio_lst = (
self.live_portrait_wrapper.calc_retargeting_ratio(
source_lmk, driving_lmk_lst
@@ -162,11 +152,9 @@ class LivePortraitPipeline(object):
if inference_cfg.flag_relative:
R_new = R_d @ R_s
delta_new = rel_s_info["exp"] + (x_d_info["exp"] - rel_s_info["exp"])
scale_new = rel_s_info["scale"] * (
x_d_info["scale"] / rel_s_info["scale"]
)
t_new = rel_s_info["t"] + (x_d_info["t"] - rel_s_info["t"])
delta_new = x_s_info["exp"] + (x_d_info["exp"] - x_s_info["exp"])
scale_new = x_s_info["scale"] * (x_d_info["scale"] / x_s_info["scale"])
t_new = x_s_info["t"] + (x_d_info["t"] - x_s_info["t"])
else:
R_new = R_d
delta_new = x_d_info["exp"]
@@ -174,12 +162,100 @@ class LivePortraitPipeline(object):
t_new = x_d_info["t"]
t_new[..., 2].fill_(0) # zero tz
x_d_new = scale_new * (x_c_s @ R_new + delta_new) + t_new
x_d_i_new = scale_new * (x_c_s @ R_new + delta_new) + t_new
if (
not inference_cfg.flag_stitching
and not inference_cfg.flag_eye_retargeting
and not inference_cfg.flag_lip_retargeting
):
# without stitching or retargeting
if inference_cfg.flag_lip_zero:
x_d_i_new += lip_delta_before_animation.reshape(-1, x_s.shape[1], 3)
else:
pass
elif (
inference_cfg.flag_stitching
and not inference_cfg.flag_eye_retargeting
and not inference_cfg.flag_lip_retargeting
):
# with stitching and without retargeting
if inference_cfg.flag_lip_zero:
x_d_i_new = self.live_portrait_wrapper.stitching(
x_s, x_d_i_new
) + lip_delta_before_animation.reshape(-1, x_s.shape[1], 3)
else:
x_d_i_new = self.live_portrait_wrapper.stitching(x_s, x_d_i_new)
else:
eyes_delta, lip_delta = None, None
if inference_cfg.flag_eye_retargeting:
c_d_eyes_i = input_eye_ratio_lst[i]
combined_eye_ratio_tensor = (
self.live_portrait_wrapper.calc_combined_eye_ratio(
c_d_eyes_i, source_lmk
)
)
combined_eye_ratio_tensor = (
combined_eye_ratio_tensor
* inference_cfg.eyes_retargeting_multiplier
)
# ∆_eyes,i = R_eyes(x_s; c_s,eyes, c_d,eyes,i)
eyes_delta = self.live_portrait_wrapper.retarget_eye(
x_s, combined_eye_ratio_tensor
)
if inference_cfg.flag_lip_retargeting:
c_d_lip_i = input_lip_ratio_lst[i]
combined_lip_ratio_tensor = (
self.live_portrait_wrapper.calc_combined_lip_ratio(
c_d_lip_i, source_lmk
)
)
combined_lip_ratio_tensor = (
combined_lip_ratio_tensor
* inference_cfg.lip_retargeting_multiplier
)
# ∆_lip,i = R_lip(x_s; c_s,lip, c_d,lip,i)
lip_delta = self.live_portrait_wrapper.retarget_lip(
x_s, combined_lip_ratio_tensor
)
if inference_cfg.flag_relative: # use x_s
x_d_i_new = (
x_s
+ (
eyes_delta.reshape(-1, x_s.shape[1], 3)
if eyes_delta is not None
else 0
)
+ (
lip_delta.reshape(-1, x_s.shape[1], 3)
if lip_delta is not None
else 0
)
)
else: # use x_d,i
x_d_i_new = (
x_d_i_new
+ (
eyes_delta.reshape(-1, x_s.shape[1], 3)
if eyes_delta is not None
else 0
)
+ (
lip_delta.reshape(-1, x_s.shape[1], 3)
if lip_delta is not None
else 0
)
)
if inference_cfg.flag_stitching:
x_d_i_new = self.live_portrait_wrapper.stitching(x_s, x_d_i_new)
out = self.live_portrait_wrapper.warp_decode(f_s, x_s, x_d_i_new)
if inference_cfg.flag_stitching:
x_d_new = self.live_portrait_wrapper.stitching(x_s, x_d_new)
x_d_i_new = self.live_portrait_wrapper.stitching(x_s, x_d_i_new)
out = self.live_portrait_wrapper.warp_decode(f_s, x_s, x_d_new)
out = self.live_portrait_wrapper.warp_decode(f_s, x_s, x_d_i_new)
I_p_i = self.live_portrait_wrapper.parse_output(out["out"])[0]
I_p_lst.append(I_p_i)