Add options what to draw on dwpose detector

This commit is contained in:
kijai
2025-04-19 19:47:09 +03:00
parent b9b8ec7bf8
commit 7b6c34e26d
3 changed files with 93 additions and 59 deletions
+15 -3
View File
@@ -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:
+53 -46
View File
@@ -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
+25 -10
View File
@@ -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