Do video cropping on the cropped node too
This commit is contained in:
@@ -61,9 +61,9 @@ class LivePortraitPipeline(object):
|
||||
):
|
||||
inference_cfg = self.live_portrait_wrapper.cfg
|
||||
|
||||
I_p_lst = []
|
||||
I_p_paste_lst = []
|
||||
driving_lmk_lst = []
|
||||
cropped_image_list = []
|
||||
composited_image_list = []
|
||||
driving_landmark_list = []
|
||||
out_mask_list = []
|
||||
R_d_0, x_d_0_info = None, None
|
||||
|
||||
@@ -72,7 +72,7 @@ class LivePortraitPipeline(object):
|
||||
pbar = comfy.utils.ProgressBar(total_frames)
|
||||
|
||||
if inference_cfg.flag_eye_retargeting or inference_cfg.flag_lip_retargeting:
|
||||
driving_lmk_lst = self.cropper.get_retargeting_lmk_info(driving_images_np)
|
||||
driving_landmark_list = self.cropper.get_retargeting_lmk_info(driving_images_np)
|
||||
|
||||
for i in tqdm(range(total_frames), desc='Animating...', total=total_frames):
|
||||
source_frame_rgb = self._get_source_frame(
|
||||
@@ -80,12 +80,8 @@ class LivePortraitPipeline(object):
|
||||
)
|
||||
driving_frame = driving_images_np[i]
|
||||
|
||||
crop_info = self.cropper.crop_single_image(source_frame_rgb, draw_keypoints=False)
|
||||
source_lmk = crop_info["lmk_crop"]
|
||||
_, img_crop_256x256 = (
|
||||
crop_info["img_crop"],
|
||||
crop_info["img_crop_256x256"],
|
||||
)
|
||||
source_lmk = crop_info[i]["lmk_crop"]
|
||||
img_crop_256x256 = crop_info[i]["img_crop_256x256"]
|
||||
|
||||
if inference_cfg.flag_do_crop:
|
||||
I_s = self.live_portrait_wrapper.prepare_source(img_crop_256x256)
|
||||
@@ -126,10 +122,10 @@ 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_landmark_list = 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
|
||||
source_lmk, driving_landmark_list
|
||||
)
|
||||
)
|
||||
|
||||
@@ -251,12 +247,12 @@ class LivePortraitPipeline(object):
|
||||
|
||||
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)
|
||||
cropped_image_list.append(I_p_i)
|
||||
|
||||
# Transform and blend
|
||||
I_p_i_to_ori = _transform_img(
|
||||
I_p_i,
|
||||
crop_info["M_c2o"],
|
||||
crop_info[i]["M_c2o"],
|
||||
dsize=(source_frame_rgb.shape[1], source_frame_rgb.shape[0]),
|
||||
)
|
||||
|
||||
@@ -265,18 +261,16 @@ class LivePortraitPipeline(object):
|
||||
inference_cfg.mask_crop = cv2.imread(os.path.join(script_directory, "utils", "resources", "mask_template.png"), cv2.IMREAD_COLOR)
|
||||
mask_ori = _transform_img(
|
||||
inference_cfg.mask_crop,
|
||||
crop_info["M_c2o"],
|
||||
crop_info[i]["M_c2o"],
|
||||
dsize=(source_frame_rgb.shape[1], source_frame_rgb.shape[0]),
|
||||
)
|
||||
mask_ori = mask_ori.astype(np.float32) / 255.0
|
||||
I_p_i_to_ori_blend = np.clip(
|
||||
mask_ori * I_p_i_to_ori + (1 - mask_ori) * source_frame_rgb, 0, 255
|
||||
).astype(np.uint8)
|
||||
else:
|
||||
I_p_i_to_ori_blend = I_p_i_to_ori
|
||||
|
||||
I_p_paste_lst.append(I_p_i_to_ori_blend)
|
||||
composited_image_list.append(I_p_i_to_ori_blend)
|
||||
out_mask_list.append(mask_ori)
|
||||
pbar.update(1)
|
||||
|
||||
return I_p_lst, I_p_paste_lst, out_mask_list
|
||||
return cropped_image_list, composited_image_list, out_mask_list
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
cropping function and the related preprocess functions for cropping
|
||||
"""
|
||||
|
||||
import cv2; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False) # NOTE: enforce single thread
|
||||
import cv2#; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False) # NOTE: enforce single thread
|
||||
import numpy as np
|
||||
from math import sin, cos, acos, degrees
|
||||
|
||||
|
||||
@@ -3,11 +3,11 @@
|
||||
import numpy as np
|
||||
from typing import List, Union, Tuple
|
||||
from dataclasses import dataclass, field
|
||||
import cv2; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False)
|
||||
import cv2#; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False)
|
||||
|
||||
from .landmark_runner import LandmarkRunner
|
||||
from .face_analysis_diy import FaceAnalysisDIY
|
||||
from .crop import crop_image, crop_image_by_bbox, parse_bbox_from_landmark, average_bbox_lst
|
||||
from .crop import crop_image
|
||||
|
||||
import folder_paths
|
||||
import os
|
||||
@@ -15,8 +15,8 @@ script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
@dataclass
|
||||
class Trajectory:
|
||||
start: int = -1 # 起始帧 闭区间
|
||||
end: int = -1 # 结束帧 闭区间
|
||||
start: int = -1
|
||||
end: int = -1
|
||||
lmk_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # lmk list
|
||||
bbox_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # bbox list
|
||||
frame_rgb_lst: Union[Tuple, List, np.ndarray] = field(default_factory=list) # frame list
|
||||
@@ -48,9 +48,11 @@ class Cropper(object):
|
||||
if hasattr(self.crop_cfg, k):
|
||||
setattr(self.crop_cfg, k, v)
|
||||
|
||||
def crop_single_image(self, img_rgb, draw_keypoints, **kwargs):
|
||||
def crop_single_image(self, img_rgb, **kwargs):
|
||||
direction = kwargs.get('direction', 'large-small')
|
||||
|
||||
|
||||
|
||||
src_face = self.face_analysis_wrapper.get(
|
||||
img_rgb,
|
||||
flag_do_landmark_2d_106=True,
|
||||
@@ -62,10 +64,9 @@ class Cropper(object):
|
||||
#elif len(src_face) > 1:
|
||||
# print(f'More than one face detected in the image, only pick one face by rule {direction}.')
|
||||
|
||||
src_face = src_face[self.crop_cfg.face_index]
|
||||
src_face = src_face[self.crop_cfg.face_index] # choose the index if multiple faces detected
|
||||
pts = src_face.landmark_2d_106
|
||||
|
||||
|
||||
# crop the face
|
||||
ret_dct = crop_image(
|
||||
img_rgb, # ndarray
|
||||
@@ -78,67 +79,19 @@ class Cropper(object):
|
||||
ret_dct['img_crop_256x256'] = cv2.resize(ret_dct['img_crop'], (256, 256), interpolation=cv2.INTER_AREA)
|
||||
ret_dct['pt_crop_256x256'] = ret_dct['pt_crop'] * 256 / kwargs.get('dsize', 512)
|
||||
|
||||
input_image_size = img_rgb.shape[:2]
|
||||
ret_dct['input_image_size'] = input_image_size
|
||||
|
||||
recon_ret = self.landmark_runner.run(img_rgb, pts)
|
||||
lmk = recon_ret['pts']
|
||||
ret_dct['lmk_crop'] = lmk
|
||||
|
||||
# Draw each landmark as a circle
|
||||
if draw_keypoints:
|
||||
print("Drawing keypoints...")
|
||||
height, width = img_rgb.shape[:2]
|
||||
blank_image = np.zeros((height, width, 3), dtype=np.uint8) * 255
|
||||
for (x, y) in lmk:
|
||||
# Ensure the coordinates are within the dimensions of the blank image
|
||||
if 0 <= x < width and 0 <= y < height:
|
||||
cv2.circle(blank_image, (int(x), int(y)), radius=2, color=(0, 0, 255))
|
||||
|
||||
keypoints_image = cv2.cvtColor(blank_image, cv2.COLOR_BGR2RGB)
|
||||
|
||||
return ret_dct, keypoints_image
|
||||
else:
|
||||
return ret_dct
|
||||
|
||||
def get_retargeting_lmk_info(self, driving_rgb_lst):
|
||||
# TODO: implement a tracking-based version
|
||||
driving_lmk_lst = []
|
||||
for driving_image in driving_rgb_lst:
|
||||
ret_dct = self.crop_single_image(driving_image, draw_keypoints=False)
|
||||
ret_dct = self.crop_single_image(driving_image)
|
||||
driving_lmk_lst.append(ret_dct['lmk_crop'])
|
||||
return driving_lmk_lst
|
||||
|
||||
def make_video_clip(self, driving_rgb_lst, output_path, output_fps=30, **kwargs):
|
||||
trajectory = Trajectory()
|
||||
direction = kwargs.get('direction', 'large-small')
|
||||
for idx, driving_image in enumerate(driving_rgb_lst):
|
||||
if idx == 0 or trajectory.start == -1:
|
||||
src_face = self.face_analysis_wrapper.get(
|
||||
driving_image,
|
||||
flag_do_landmark_2d_106=True,
|
||||
direction=direction
|
||||
)
|
||||
if len(src_face) == 0:
|
||||
# No face detected in the driving_image
|
||||
continue
|
||||
elif len(src_face) > 1:
|
||||
print(f'More than one face detected in the driving frame_{idx}, only pick one face by rule {direction}.')
|
||||
src_face = src_face[0]
|
||||
pts = src_face.landmark_2d_106
|
||||
lmk_203 = self.landmark_runner(driving_image, pts)['pts']
|
||||
trajectory.start, trajectory.end = idx, idx
|
||||
else:
|
||||
lmk_203 = self.face_recon_wrapper(driving_image, trajectory.lmk_lst[-1])['pts']
|
||||
trajectory.end = idx
|
||||
|
||||
trajectory.lmk_lst.append(lmk_203)
|
||||
ret_bbox = parse_bbox_from_landmark(lmk_203, scale=self.crop_cfg.globalscale, vy_ratio=elf.crop_cfg.vy_ratio)['bbox']
|
||||
bbox = [ret_bbox[0, 0], ret_bbox[0, 1], ret_bbox[2, 0], ret_bbox[2, 1]] # 4,
|
||||
trajectory.bbox_lst.append(bbox) # bbox
|
||||
trajectory.frame_rgb_lst.append(driving_image)
|
||||
|
||||
global_bbox = average_bbox_lst(trajectory.bbox_lst)
|
||||
for idx, (frame_rgb, lmk) in enumerate(zip(trajectory.frame_rgb_lst, trajectory.lmk_lst)):
|
||||
ret_dct = crop_image_by_bbox(
|
||||
frame_rgb, global_bbox, lmk=lmk,
|
||||
dsize=self.video_crop_cfg.dsize, flag_rot=self.video_crop_cfg.flag_rot, borderValue=self.video_crop_cfg.borderValue
|
||||
)
|
||||
frame_rgb_crop = ret_dct['img_crop']
|
||||
|
||||
@@ -1,15 +1,30 @@
|
||||
# coding: utf-8
|
||||
|
||||
"""
|
||||
face detectoin and alignment using InsightFace
|
||||
face detection and alignment using InsightFace
|
||||
"""
|
||||
from insightface.utils import transform
|
||||
|
||||
#patch Insightface function to get rid of the annoying warnings
|
||||
def patched_estimate_affine_matrix_3d23d(X, Y):
|
||||
''' Using least-squares solution
|
||||
Args:
|
||||
X: [n, 3]. 3d points(fixed)
|
||||
Y: [n, 3]. corresponding 3d points(moving). Y = PX
|
||||
Returns:
|
||||
P_Affine: (3, 4). Affine camera matrix (the third row is [0, 0, 0, 1]).
|
||||
'''
|
||||
X_homo = np.hstack((X, np.ones([X.shape[0],1]))) # n x 4
|
||||
P = np.linalg.lstsq(X_homo, Y, rcond=None)[0].T # Affine matrix. 3 x 4
|
||||
return P
|
||||
|
||||
transform.estimate_affine_matrix_3d23d = patched_estimate_affine_matrix_3d23d
|
||||
|
||||
import numpy as np
|
||||
from insightface.app import FaceAnalysis
|
||||
from insightface.app.common import Face
|
||||
from .timer import Timer
|
||||
|
||||
|
||||
def sort_by_direction(faces, direction: str = 'large-small', face_center=None):
|
||||
if len(faces) <= 0:
|
||||
return faces
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# coding: utf-8
|
||||
|
||||
import cv2; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False)
|
||||
import cv2#; cv2.setNumThreads(0); cv2.ocl.setUseOpenCL(False)
|
||||
import torch
|
||||
import numpy as np
|
||||
import onnxruntime
|
||||
|
||||
@@ -6,6 +6,7 @@ import comfy.model_management as mm
|
||||
import comfy.utils
|
||||
import numpy as np
|
||||
import cv2
|
||||
from tqdm import tqdm
|
||||
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
@@ -224,7 +225,7 @@ class LivePortraitProcess:
|
||||
return {"required": {
|
||||
|
||||
"pipeline": ("LIVEPORTRAITPIPE",),
|
||||
"crop_info": ("CROPINFO", {"default": {}}),
|
||||
"crop_info": ("CROPINFO", {"default": []}),
|
||||
"source_image": ("IMAGE",),
|
||||
"driving_images": ("IMAGE",),
|
||||
"lip_zero": ("BOOLEAN", {"default": True}),
|
||||
@@ -267,7 +268,7 @@ class LivePortraitProcess:
|
||||
self,
|
||||
source_image: torch.Tensor,
|
||||
driving_images: torch.Tensor,
|
||||
crop_info: dict,
|
||||
crop_info: list,
|
||||
pipeline: LivePortraitPipeline,
|
||||
lip_zero: bool,
|
||||
lip_zero_threshold: float,
|
||||
@@ -305,13 +306,11 @@ class LivePortraitProcess:
|
||||
crop_mask = np.repeat(np.atleast_3d(crop_mask), 3, axis=2)
|
||||
pipeline.live_portrait_wrapper.cfg.mask_crop = crop_mask
|
||||
|
||||
pipeline.cropper = crop_info['cropper']
|
||||
|
||||
cropped_out_list = []
|
||||
full_out_list = []
|
||||
|
||||
cropped_out_list, full_out_list, out_mask_list = pipeline.execute(
|
||||
source_np, driving_images_np, crop_info['crop_info'], mismatch_method
|
||||
source_np, driving_images_np, crop_info, mismatch_method
|
||||
)
|
||||
cropped_tensors_out = (
|
||||
torch.stack([torch.from_numpy(np_array) for np_array in cropped_out_list])
|
||||
@@ -343,20 +342,15 @@ class LivePortraitCropper:
|
||||
"vx_ratio": ("FLOAT", {"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.01}),
|
||||
"vy_ratio": ("FLOAT", {"default": -0.125, "min": -1.0, "max": 1.0, "step": 0.01}),
|
||||
"face_index": ("INT", {"default": 0, "min": 0, "max": 100}),
|
||||
},
|
||||
"optional": {
|
||||
"onnx_device": (
|
||||
[
|
||||
'CPU',
|
||||
'CUDA',
|
||||
], {
|
||||
['CPU', 'CUDA', 'ROCM'], {
|
||||
"default": 'CPU'
|
||||
}),
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "CROPINFO", "IMAGE",)
|
||||
RETURN_NAMES = ("cropped_image", "crop_info", "keypoints_image",)
|
||||
RETURN_TYPES = ("IMAGE", "CROPINFO",)
|
||||
RETURN_NAMES = ("cropped_image", "crop_info",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "LivePortrait"
|
||||
|
||||
@@ -372,21 +366,56 @@ class LivePortraitCropper:
|
||||
)
|
||||
|
||||
cropper = Cropper(crop_cfg=crop_cfg, provider=onnx_device)
|
||||
crop_info, keypoints_img = cropper.crop_single_image(source_image_np[0], draw_keypoints=True)
|
||||
|
||||
keypoints_image_tensor = torch.from_numpy(keypoints_img) / 255
|
||||
keypoints_image_tensor = keypoints_image_tensor.unsqueeze(0).cpu().float()
|
||||
crop_info_list = []
|
||||
|
||||
pbar = comfy.utils.ProgressBar(len(source_image_np))
|
||||
for i in tqdm(range(len(source_image_np)), desc='Detecting and cropping..', total=len(source_image_np)):
|
||||
crop_info = cropper.crop_single_image(source_image_np[i])
|
||||
crop_info_list.append(crop_info)
|
||||
pbar.update(1)
|
||||
|
||||
cropped_image = crop_info['img_crop_256x256']
|
||||
cropped_tensors = torch.from_numpy(cropped_image) / 255
|
||||
cropped_tensors = cropped_tensors.unsqueeze(0).cpu().float()
|
||||
|
||||
cropper_dict = {
|
||||
"cropper": cropper,
|
||||
"crop_info": crop_info,
|
||||
return (cropped_tensors, crop_info_list)
|
||||
|
||||
class KeypointsToImage:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"crop_info": ("CROPINFO", {"default": []}),
|
||||
},
|
||||
}
|
||||
|
||||
return (cropped_tensors, cropper_dict, keypoints_image_tensor)
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("keypoints_image",)
|
||||
FUNCTION = "drawkeypoints"
|
||||
CATEGORY = "LivePortrait"
|
||||
|
||||
def drawkeypoints(self, crop_info):
|
||||
height, width = crop_info[0]['input_image_size']
|
||||
keypoints_img_list = []
|
||||
pbar = comfy.utils.ProgressBar(len(crop_info))
|
||||
for crop in crop_info:
|
||||
keypoints = crop['lmk_crop'].copy()
|
||||
# Draw each landmark as a circle
|
||||
blank_image = np.zeros((height, width, 3), dtype=np.uint8) * 255
|
||||
for (x, y) in keypoints:
|
||||
# Ensure the coordinates are within the dimensions of the blank image
|
||||
if 0 <= x < width and 0 <= y < height:
|
||||
cv2.circle(blank_image, (int(x), int(y)), radius=2, color=(0, 0, 255))
|
||||
|
||||
keypoints_image = cv2.cvtColor(blank_image, cv2.COLOR_BGR2RGB)
|
||||
keypoints_img_list.append(keypoints_image)
|
||||
pbar.update(1)
|
||||
|
||||
keypoints_img_tensor = (
|
||||
torch.stack([torch.from_numpy(np_array) for np_array in keypoints_img_list]) / 255).float()
|
||||
|
||||
|
||||
return (keypoints_img_tensor,)
|
||||
|
||||
class KeypointScaler:
|
||||
@classmethod
|
||||
@@ -442,11 +471,13 @@ NODE_CLASS_MAPPINGS = {
|
||||
"DownloadAndLoadLivePortraitModels": DownloadAndLoadLivePortraitModels,
|
||||
"LivePortraitProcess": LivePortraitProcess,
|
||||
"LivePortraitCropper": LivePortraitCropper,
|
||||
"KeypointScaler": KeypointScaler
|
||||
"KeypointScaler": KeypointScaler,
|
||||
"KeypointsToImage": KeypointsToImage
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"DownloadAndLoadLivePortraitModels": "(Down)Load LivePortraitModels",
|
||||
"LivePortraitProcess": "LivePortraitProcess",
|
||||
"LivePortraitCropper": "LivePortraitCropper",
|
||||
"KeypointScaler": "KeypointScaler"
|
||||
"KeypointScaler": "KeypointScaler",
|
||||
"KeypointsToImage": "LivePortrait KeypointsToImage"
|
||||
}
|
||||
Reference in New Issue
Block a user