bring KJ edits
This commit is contained in:
@@ -60,7 +60,7 @@ class LivePortraitPipeline(object):
|
||||
]
|
||||
|
||||
def execute(
|
||||
self, source_np, driving_images_np, mismatch_method="repeat", reference_frame=0
|
||||
self, source_np, driving_images_np, crop_info, mismatch_method="repeat", reference_frame=0
|
||||
):
|
||||
inference_cfg = self.live_portrait_wrapper.cfg
|
||||
|
||||
@@ -82,7 +82,7 @@ class LivePortraitPipeline(object):
|
||||
)
|
||||
driving_frame = driving_images_np[i]
|
||||
|
||||
crop_info = self.cropper.crop_single_image(source_frame_rgb)
|
||||
crop_info, _ = self.cropper.crop_single_image(source_frame_rgb)
|
||||
source_lmk = crop_info["lmk_crop"]
|
||||
_, img_crop_256x256 = (
|
||||
crop_info["img_crop"],
|
||||
@@ -106,7 +106,7 @@ class LivePortraitPipeline(object):
|
||||
c_d_lip_before_animation = [0.0]
|
||||
combined_lip_ratio_tensor_before_animation = (
|
||||
self.live_portrait_wrapper.calc_combined_lip_ratio(
|
||||
c_d_lip_before_animation, source_lmk
|
||||
c_d_lip_before_animation, source_lmk, inference_cfg
|
||||
)
|
||||
)
|
||||
# TODO: expose lip_zero_threshold
|
||||
|
||||
@@ -301,20 +301,20 @@ class LivePortraitWrapper(object):
|
||||
input_lip_ratio_lst.append(calc_lip_close_ratio(lmk[None]))
|
||||
return input_eye_ratio_lst, input_lip_ratio_lst
|
||||
|
||||
def calc_combined_eye_ratio(self, input_eye_ratio, source_lmk):
|
||||
def calc_combined_eye_ratio(self, input_eye_ratio, source_lmk, inference_cfg):
|
||||
eye_close_ratio = calc_eye_close_ratio(source_lmk[None])
|
||||
eye_close_ratio_tensor = torch.from_numpy(eye_close_ratio).float().to(self.device_id)
|
||||
input_eye_ratio_tensor = torch.Tensor([input_eye_ratio[0][0]]).reshape(1, 1).to(self.device_id)
|
||||
input_eye_ratio_tensor = torch.Tensor([input_eye_ratio[0][0]]).reshape(1, 1).to(self.device_id) * inference_cfg.eyes_retargeting_multiplier
|
||||
# [c_s,eyes, c_d,eyes,i]
|
||||
combined_eye_ratio_tensor = torch.cat([eye_close_ratio_tensor, input_eye_ratio_tensor], dim=1)
|
||||
return combined_eye_ratio_tensor
|
||||
|
||||
def calc_combined_lip_ratio(self, input_lip_ratio, source_lmk):
|
||||
def calc_combined_lip_ratio(self, input_lip_ratio, source_lmk, inference_cfg):
|
||||
lip_close_ratio = calc_lip_close_ratio(source_lmk[None])
|
||||
lip_close_ratio_tensor = torch.from_numpy(lip_close_ratio).float().to(self.device_id)
|
||||
# [c_s,lip, c_d,lip,i]
|
||||
input_lip_ratio_tensor = torch.Tensor([input_lip_ratio[0]]).to(self.device_id)
|
||||
if input_lip_ratio_tensor.shape != [1, 1]:
|
||||
input_lip_ratio_tensor = input_lip_ratio_tensor.reshape(1, 1)
|
||||
combined_lip_ratio_tensor = torch.cat([lip_close_ratio_tensor, input_lip_ratio_tensor], dim=1)
|
||||
combined_lip_ratio_tensor = torch.cat([lip_close_ratio_tensor, input_lip_ratio_tensor], dim=1) * inference_cfg.lip_retargeting_multiplier
|
||||
return combined_lip_ratio_tensor
|
||||
|
||||
@@ -83,6 +83,7 @@ class Cropper(object):
|
||||
src_face = src_face[0]
|
||||
pts = src_face.landmark_2d_106
|
||||
|
||||
|
||||
# crop the face
|
||||
ret_dct = crop_image(
|
||||
img_rgb, # ndarray
|
||||
@@ -99,7 +100,17 @@ class Cropper(object):
|
||||
lmk = recon_ret['pts']
|
||||
ret_dct['lmk_crop'] = lmk
|
||||
|
||||
return ret_dct
|
||||
# Draw each landmark as a circle
|
||||
width, height = 512, 512
|
||||
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
|
||||
|
||||
def get_retargeting_lmk_info(self, driving_rgb_lst):
|
||||
# TODO: implement a tracking-based version
|
||||
|
||||
@@ -4,6 +4,8 @@ import yaml
|
||||
import folder_paths
|
||||
import comfy.model_management as mm
|
||||
import comfy.utils
|
||||
import numpy as np
|
||||
import cv2
|
||||
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
@@ -253,54 +255,22 @@ class DownloadAndLoadLivePortraitModels:
|
||||
return (pipeline,)
|
||||
|
||||
|
||||
# OUR CURRENT NODE
|
||||
class LivePortraitProcess:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"pipeline": ("LIVEPORTRAITPIPE",),
|
||||
"source_image": ("IMAGE",),
|
||||
"driving_images": ("IMAGE",),
|
||||
"dsize": ("INT", {"default": 512, "min": 64, "max": 2048}),
|
||||
"scale": (
|
||||
"FLOAT",
|
||||
{"default": 2.3, "min": 1.0, "max": 4.0, "step": 0.01},
|
||||
),
|
||||
"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},
|
||||
),
|
||||
"lip_zero": ("BOOLEAN", {"default": True}),
|
||||
"eye_retargeting": ("BOOLEAN", {"default": False}),
|
||||
"eyes_retargeting_multiplier": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001},
|
||||
),
|
||||
"lip_retargeting": ("BOOLEAN", {"default": False}),
|
||||
"lip_retargeting_multiplier": (
|
||||
"FLOAT",
|
||||
{"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001},
|
||||
),
|
||||
"stitching": ("BOOLEAN", {"default": True}),
|
||||
"relative": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
"optional": {
|
||||
"mismatch_method": (
|
||||
["repeat", "cycle", "mirror", "nearest"],
|
||||
{"default": "repeat"},
|
||||
),
|
||||
"onnx_device": (
|
||||
[
|
||||
"CPU",
|
||||
"CUDA",
|
||||
],
|
||||
{"default": "CPU"},
|
||||
),
|
||||
return {"required": {
|
||||
|
||||
"pipeline": ("LIVEPORTRAITPIPE",),
|
||||
"crop_info": ("CROPINFO", {"default": {}}),
|
||||
"source_image": ("IMAGE",),
|
||||
"driving_images": ("IMAGE",),
|
||||
"lip_zero": ("BOOLEAN", {"default": True}),
|
||||
"eye_retargeting": ("BOOLEAN", {"default": False}),
|
||||
"eyes_retargeting_multiplier": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001}),
|
||||
"lip_retargeting": ("BOOLEAN", {"default": False}),
|
||||
"lip_retargeting_multiplier": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001}),
|
||||
"stitching": ("BOOLEAN", {"default": True}),
|
||||
"relative": ("BOOLEAN", {"default": True}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -319,10 +289,7 @@ class LivePortraitProcess:
|
||||
self,
|
||||
source_image: torch.Tensor,
|
||||
driving_images: torch.Tensor,
|
||||
dsize: int,
|
||||
scale: float,
|
||||
vx_ratio: float,
|
||||
vy_ratio: float,
|
||||
crop_info: dict,
|
||||
pipeline: LivePortraitPipeline,
|
||||
lip_zero: bool,
|
||||
eye_retargeting: bool,
|
||||
@@ -332,20 +299,10 @@ class LivePortraitProcess:
|
||||
eyes_retargeting_multiplier: float,
|
||||
lip_retargeting_multiplier: float,
|
||||
mismatch_method: str = "repeat",
|
||||
onnx_device="CUDA",
|
||||
):
|
||||
source_np = (source_image * 255).byte().numpy()
|
||||
driving_images_np = (driving_images * 255).byte().numpy()
|
||||
|
||||
crop_cfg = CropConfig(
|
||||
dsize=dsize,
|
||||
scale=scale,
|
||||
vx_ratio=vx_ratio,
|
||||
vy_ratio=vy_ratio,
|
||||
)
|
||||
|
||||
cropper = Cropper(crop_cfg=crop_cfg, provider=onnx_device)
|
||||
pipeline.cropper = cropper
|
||||
pipeline.live_portrait_wrapper.cfg.flag_eye_retargeting = eye_retargeting
|
||||
pipeline.live_portrait_wrapper.cfg.eyes_retargeting_multiplier = (
|
||||
eyes_retargeting_multiplier
|
||||
@@ -358,11 +315,13 @@ class LivePortraitProcess:
|
||||
pipeline.live_portrait_wrapper.cfg.flag_relative = relative
|
||||
pipeline.live_portrait_wrapper.cfg.flag_lip_zero = lip_zero
|
||||
|
||||
pipeline.cropper = crop_info['cropper']
|
||||
|
||||
cropped_out_list = []
|
||||
full_out_list = []
|
||||
|
||||
cropped_out_list, full_out_list = pipeline.execute(
|
||||
source_np, driving_images_np, mismatch_method
|
||||
source_np, driving_images_np, crop_info['crop_info'], mismatch_method
|
||||
)
|
||||
cropped_tensors_out = (
|
||||
torch.stack([torch.from_numpy(np_array) for np_array in cropped_out_list])
|
||||
@@ -376,10 +335,124 @@ class LivePortraitProcess:
|
||||
return (cropped_tensors_out.cpu().float(), full_tensors_out.cpu().float())
|
||||
|
||||
|
||||
class LivePortraitCropper:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
|
||||
"source_image": ("IMAGE",),
|
||||
"dsize": ("INT", {"default": 512, "min": 64, "max": 2048}),
|
||||
"scale": ("FLOAT", {"default": 2.3, "min": 1.0, "max": 4.0, "step": 0.01}),
|
||||
"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}),
|
||||
},
|
||||
"optional": {
|
||||
"onnx_device": (
|
||||
[
|
||||
'CPU',
|
||||
'CUDA',
|
||||
], {
|
||||
"default": 'CPU'
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "CROPINFO", "IMAGE",)
|
||||
RETURN_NAMES = ("cropped_image", "crop_info", "keypoints_image",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "LivePortrait"
|
||||
|
||||
def process(self, source_image, dsize, scale, vx_ratio, vy_ratio, onnx_device='CUDA'):
|
||||
source_image_np = (source_image * 255).byte().numpy()
|
||||
|
||||
crop_cfg = CropConfig(
|
||||
dsize = dsize,
|
||||
scale = scale,
|
||||
vx_ratio = vx_ratio,
|
||||
vy_ratio = vy_ratio,
|
||||
)
|
||||
|
||||
cropper = Cropper(crop_cfg=crop_cfg, provider=onnx_device)
|
||||
crop_info, keypoints_img = cropper.crop_single_image(source_image_np[0])
|
||||
|
||||
keypoints_image_tensor = torch.from_numpy(keypoints_img) / 255
|
||||
keypoints_image_tensor = keypoints_image_tensor.unsqueeze(0).cpu().float()
|
||||
print(keypoints_image_tensor.shape)
|
||||
|
||||
|
||||
cropped_image = crop_info['img_crop_256x256']
|
||||
cropped_tensors = torch.from_numpy(cropped_image) / 255
|
||||
cropped_tensors = cropped_tensors.unsqueeze(0).cpu().float()
|
||||
|
||||
print(cropped_tensors.shape)
|
||||
|
||||
cropper_dict = {
|
||||
"cropper": cropper,
|
||||
"crop_info": crop_info,
|
||||
}
|
||||
|
||||
return (cropped_tensors, cropper_dict, keypoints_image_tensor)
|
||||
|
||||
class KeypointScaler:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"crop_info": ("CROPINFO", {"default": {}}),
|
||||
"scale": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001}),
|
||||
"offset_x": ("INT", {"default": 0, "min": -1024, "max": 1024, "step": 1}),
|
||||
"offset_y": ("INT", {"default": 0, "min": -1024, "max": 1024, "step": 1}),
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CROPINFO", "IMAGE",)
|
||||
RETURN_NAMES = ("crop_info", "keypoints_image",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "LivePortrait"
|
||||
|
||||
def process(self, crop_info, offset_x, offset_y, scale):
|
||||
|
||||
keypoints = crop_info['crop_info']['lmk_crop'].copy()
|
||||
|
||||
# Create an offset array
|
||||
# Calculate the centroid of the keypoints
|
||||
centroid = keypoints.mean(axis=0)
|
||||
|
||||
# Translate keypoints to origin by subtracting the centroid
|
||||
translated_keypoints = keypoints - centroid
|
||||
|
||||
# Scale the translated keypoints
|
||||
scaled_keypoints = translated_keypoints * scale
|
||||
|
||||
# Translate scaled keypoints back to original position and then apply the offset
|
||||
final_keypoints = scaled_keypoints + centroid + np.array([offset_x, offset_y])
|
||||
|
||||
crop_info['crop_info']['lmk_crop'] = final_keypoints
|
||||
|
||||
# Draw each landmark as a circle
|
||||
width, height = 512, 512
|
||||
blank_image = np.zeros((height, width, 3), dtype=np.uint8) * 255
|
||||
for (x, y) in final_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_image_tensor = torch.from_numpy(keypoints_image) / 255
|
||||
keypoints_image_tensor = keypoints_image_tensor.unsqueeze(0).cpu().float()
|
||||
print(keypoints_image_tensor.shape)
|
||||
|
||||
return (crop_info, keypoints_image_tensor,)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"DownloadAndLoadLivePortraitModels": DownloadAndLoadLivePortraitModels,
|
||||
"LivePortraitProcess": LivePortraitProcess,
|
||||
"LivePortraitCropper": LivePortraitCropper,
|
||||
"KeypointScaler": KeypointScaler
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"DownloadAndLoadLivePortraitModels": "(Down)Load LivePortraitModels",
|
||||
}
|
||||
"LivePortraitProcess": "LivePortraitProcess",
|
||||
"LivePortraitCropper": "LivePortraitCropper",
|
||||
"KeypointScaler": "KeypointScaler"
|
||||
}
|
||||
Reference in New Issue
Block a user