diff --git a/nodes.py b/nodes.py index 215a676..e263318 100644 --- a/nodes.py +++ b/nodes.py @@ -2643,7 +2643,8 @@ class WanVideoSampler: fresca_freq_cutoff = experimental_args.get("fresca_freq_cutoff", 20) #region model pred - def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, control_latents=None, vace_data=None, teacache_state=None): + def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, + control_latents=None, vace_data=None, unianim_data=None, teacache_state=None): with torch.autocast(device_type=mm.get_autocast_device(device), dtype=model["dtype"], enabled=True): if use_cfg_zero_star and (idx <= zero_star_steps) and use_zero_init: @@ -3015,11 +3016,22 @@ class WanVideoSampler: partial_vace_context = [partial_vace_context] partial_latent_model_input = latent_model_input[:, c, :, :] + partial_unianim_data = None + if unianim_data is not None: + partial_dwpose = unianim_data["dwpose"][:, c, :, :] + partial_unianim_data = { + "dwpose": partial_dwpose, + "random_ref": unianim_data["random_ref"], + "strength": unianimate_poses["strength"], + "start_percent": unianimate_poses["start_percent"], + "end_percent": unianimate_poses["end_percent"] + } + noise_pred_context, new_teacache = predict_with_cfg( partial_latent_model_input, cfg[idx], positive, text_embeds["negative_prompt_embeds"], - timestep, idx, partial_img_emb, clip_fea, partial_control_latents, partial_vace_context, + timestep, idx, partial_img_emb, clip_fea, partial_control_latents, partial_vace_context, partial_unianim_data, current_teacache) # if callback is not None: @@ -3042,7 +3054,7 @@ class WanVideoSampler: cfg[idx], text_embeds["prompt_embeds"], text_embeds["negative_prompt_embeds"], - timestep, idx, image_cond, clip_fea, control_latents, vace_data, + timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, teacache_state=self.teacache_state) if latent_shift_loop: diff --git a/unianimate/dwpose/util.py b/unianimate/dwpose/util.py index 2f83229..bc8a6e3 100644 --- a/unianimate/dwpose/util.py +++ b/unianimate/dwpose/util.py @@ -109,51 +109,56 @@ def draw_bodypose(canvas, candidate, subset): return canvas -def draw_body_and_foot(canvas, candidate, subset): +def draw_body_and_foot(canvas, candidate, subset, stick_width=4, draw_body=True, draw_feet=True, draw_body_keypoints=True): H, W, C = canvas.shape candidate = np.array(candidate) subset = np.array(subset) - stickwidth = 4 - - limbSeq = [[2, 3], [2, 6], [3, 4], [4, 5], [6, 7], [7, 8], [2, 9], [9, 10], \ - [10, 11], [2, 12], [12, 13], [13, 14], [2, 1], [1, 15], [15, 17], \ - [1, 16], [16, 18], [14,19], [11, 20]] - + if draw_feet: + limbSeq = [[2, 3], [2, 6], [3, 4], [4, 5], [6, 7], [7, 8], [2, 9], [9, 10], \ + [10, 11], [2, 12], [12, 13], [13, 14], [2, 1], [1, 15], [15, 17], \ + [1, 16], [16, 18], [14,19], [11, 20]] + else: + limbSeq = [[2, 3], [2, 6], [3, 4], [4, 5], [6, 7], [7, 8], [2, 9], [9, 10], \ + [10, 11], [2, 12], [12, 13], [13, 14], [2, 1], [1, 15], [15, 17], \ + [1, 16], [16, 18]] + colors = [[255, 0, 0], [255, 85, 0], [255, 170, 0], [255, 255, 0], [170, 255, 0], [85, 255, 0], [0, 255, 0], \ [0, 255, 85], [0, 255, 170], [0, 255, 255], [0, 170, 255], [0, 85, 255], [0, 0, 255], [85, 0, 255], \ [170, 0, 255], [255, 0, 255], [255, 0, 170], [255, 0, 85], [170, 255, 255], [255, 255, 0]] - - for i in range(19): - for n in range(len(subset)): - index = subset[n][np.array(limbSeq[i]) - 1] - if -1 in index: - continue - Y = candidate[index.astype(int), 0] * float(W) - X = candidate[index.astype(int), 1] * float(H) - mX = np.mean(X) - mY = np.mean(Y) - length = ((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2) ** 0.5 - angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1])) - polygon = cv2.ellipse2Poly((int(mY), int(mX)), (int(length / 2), stickwidth), int(angle), 0, 360, 1) - cv2.fillConvexPoly(canvas, polygon, colors[i]) + + if draw_body: + for i in range(len(limbSeq)): + for n in range(len(subset)): + index = subset[n][np.array(limbSeq[i]) - 1] + if -1 in index: + continue + Y = candidate[index.astype(int), 0] * float(W) + X = candidate[index.astype(int), 1] * float(H) + mX = np.mean(X) + mY = np.mean(Y) + length = ((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2) ** 0.5 + angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1])) + polygon = cv2.ellipse2Poly((int(mY), int(mX)), (int(length / 2), stick_width), int(angle), 0, 360, 1) + cv2.fillConvexPoly(canvas, polygon, colors[i]) canvas = (canvas * 0.6).astype(np.uint8) - for i in range(20): - for n in range(len(subset)): - index = int(subset[n][i]) - if index == -1: - continue - x, y = candidate[index][0:2] - x = int(x * W) - y = int(y * H) - cv2.circle(canvas, (int(x), int(y)), 4, colors[i], thickness=-1) + if draw_body_keypoints: + for i in range(len(limbSeq)+1): + for n in range(len(subset)): + index = int(subset[n][i]) + if index == -1: + continue + x, y = candidate[index][0:2] + x = int(x * W) + y = int(y * H) + cv2.circle(canvas, (int(x), int(y)), 4, colors[i], thickness=-1) return canvas -def draw_handpose(canvas, all_hand_peaks): +def draw_handpose(canvas, all_hand_peaks, draw_hands=True, draw_hand_keypoints=True): H, W, C = canvas.shape edges = [[0, 1], [1, 2], [2, 3], [3, 4], [0, 5], [5, 6], [6, 7], [7, 8], [0, 9], [9, 10], \ @@ -162,22 +167,24 @@ def draw_handpose(canvas, all_hand_peaks): for peaks in all_hand_peaks: peaks = np.array(peaks) - for ie, e in enumerate(edges): - x1, y1 = peaks[e[0]] - x2, y2 = peaks[e[1]] - x1 = int(x1 * W) - y1 = int(y1 * H) - x2 = int(x2 * W) - y2 = int(y2 * H) - if x1 > eps and y1 > eps and x2 > eps and y2 > eps: - cv2.line(canvas, (x1, y1), (x2, y2), matplotlib.colors.hsv_to_rgb([ie / float(len(edges)), 1.0, 1.0]) * 255, thickness=2) + if draw_hands: + for ie, e in enumerate(edges): + x1, y1 = peaks[e[0]] + x2, y2 = peaks[e[1]] + x1 = int(x1 * W) + y1 = int(y1 * H) + x2 = int(x2 * W) + y2 = int(y2 * H) + if x1 > eps and y1 > eps and x2 > eps and y2 > eps: + cv2.line(canvas, (x1, y1), (x2, y2), matplotlib.colors.hsv_to_rgb([ie / float(len(edges)), 1.0, 1.0]) * 255, thickness=2) - for i, keyponit in enumerate(peaks): - x, y = keyponit - x = int(x * W) - y = int(y * H) - if x > eps and y > eps: - cv2.circle(canvas, (x, y), 4, (0, 0, 255), thickness=-1) + if draw_hand_keypoints: + for i, keypoint in enumerate(peaks): + x, y = keypoint + x = int(x * W) + y = int(y * H) + if x > eps and y > eps: + cv2.circle(canvas, (x, y), 4, (0, 0, 255), thickness=-1) return canvas diff --git a/unianimate/nodes.py b/unianimate/nodes.py index be1e2ff..e12129e 100644 --- a/unianimate/nodes.py +++ b/unianimate/nodes.py @@ -184,7 +184,8 @@ class DWposeDetector: # return draw_pose(pose, H, W) return pose -def draw_pose(pose, H, W): +def draw_pose(pose, H, W, stick_width=4,draw_body=True, draw_hands=True, draw_feet=True, + draw_body_keypoints=True, draw_hand_keypoints=True): bodies = pose['bodies'] faces = pose['faces'] hands = pose['hands'] @@ -192,15 +193,17 @@ def draw_pose(pose, H, W): subset = bodies['subset'] canvas = np.zeros(shape=(H, W, 3), dtype=np.uint8) - canvas = draw_body_and_foot(canvas, candidate, subset) - canvas = draw_handpose(canvas, hands) + canvas = draw_body_and_foot(canvas, candidate, subset, draw_body=draw_body, stick_width=stick_width, draw_feet=draw_feet, draw_body_keypoints=draw_body_keypoints) + canvas = draw_handpose(canvas, hands, draw_hands=draw_hands, draw_hand_keypoints=draw_hand_keypoints) canvas_without_face = copy.deepcopy(canvas) canvas = draw_facepose(canvas, faces) return canvas_without_face, canvas -def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_threshold): +def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_threshold, stick_width, + draw_body=True, draw_hands=True, draw_hand_keypoints=True, draw_feet=True, + draw_body_keypoints=True): results_vis = [] comfy_pbar = ProgressBar(len(pose_images)) @@ -674,7 +677,9 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre dwpose_woface_list = [] for i in range(len(results_vis)): try: - dwpose_woface, dwpose_wface = draw_pose(results_vis[i], H=height, W=width) + dwpose_woface, dwpose_wface = draw_pose(results_vis[i], H=height, W=width, stick_width=stick_width, + draw_body=draw_body, draw_hands=draw_hands, draw_hand_keypoints=draw_hand_keypoints, + draw_feet=draw_feet, draw_body_keypoints=draw_body_keypoints) result = torch.from_numpy(dwpose_woface) except: result = torch.zeros((height, width, 3), dtype=torch.uint8) @@ -683,7 +688,9 @@ def pose_extract(pose_images, ref_image, dwpose_model, height, width, score_thre dwpose_woface_ref_tensor = None if ref_image is not None: - dwpose_woface_ref, dwpose_wface_ref = draw_pose(pose_ref, H=height, W=width) + dwpose_woface_ref, dwpose_wface_ref = draw_pose(pose_ref, H=height, W=width, stick_width=stick_width, + draw_body=draw_body, draw_hands=draw_hands, draw_hand_keypoints=draw_hand_keypoints, + draw_feet=draw_feet, draw_body_keypoints=draw_body_keypoints) dwpose_woface_ref_tensor = torch.from_numpy(dwpose_woface_ref) return dwpose_woface_tensor, dwpose_woface_ref_tensor @@ -692,8 +699,14 @@ class WanVideoUniAnimateDWPoseDetector: @classmethod def INPUT_TYPES(s): return {"required": { - "pose_images": ("IMAGE", {"tooltip": "Pose images"}), - "score_threshold": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Score threshold for pose detection"}), + "pose_images": ("IMAGE", {"tooltip": "Pose images"}), + "score_threshold": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Score threshold for pose detection"}), + "stick_width": ("INT", {"default": 4, "min": 1, "max": 100, "step": 1, "tooltip": "Stick width for drawing keypoints"}), + "draw_body": ("BOOLEAN", {"default": True, "tooltip": "Draw body keypoints"}), + "draw_body_keypoints": ("BOOLEAN", {"default": True, "tooltip": "Draw body keypoints"}), + "draw_feet": ("BOOLEAN", {"default": True, "tooltip": "Draw feet keypoints"}), + "draw_hands": ("BOOLEAN", {"default": True, "tooltip": "Draw hand keypoints"}), + "draw_hand_keypoints": ("BOOLEAN", {"default": True, "tooltip": "Draw hand keypoints"}), }, "optional": { "reference_pose_image": ("IMAGE", {"tooltip": "Reference pose image"}), @@ -705,7 +718,7 @@ class WanVideoUniAnimateDWPoseDetector: FUNCTION = "process" CATEGORY = "WanVideoWrapper" - def process(self, pose_images, score_threshold, reference_pose_image=None): + def process(self, pose_images, score_threshold, stick_width, reference_pose_image=None, draw_body=True, draw_body_keypoints=True, draw_feet=True, draw_hands=True, draw_hand_keypoints=True): device = mm.get_torch_device() @@ -749,7 +762,9 @@ class WanVideoUniAnimateDWPoseDetector: ref = reference_pose_image ref_np = ref.cpu().numpy() * 255 - poses, reference_pose = pose_extract(pose_np, ref_np, self.dwpose_detector, height, width, score_threshold) + poses, reference_pose = pose_extract(pose_np, ref_np, self.dwpose_detector, height, width, score_threshold, stick_width=stick_width, + draw_body=draw_body, draw_body_keypoints=draw_body_keypoints, draw_feet=draw_feet, + draw_hands=draw_hands, draw_hand_keypoints=draw_hand_keypoints) poses = poses / 255.0 if reference_pose_image is not None: reference_pose = reference_pose.unsqueeze(0) / 255.0