Support UniAnimate-DiT

https://github.com/ali-vilab/UniAnimate-DiT
This commit is contained in:
kijai
2025-04-18 19:43:40 +03:00
parent 00de4f5e0e
commit 7462356743
13 changed files with 2265 additions and 11 deletions
+2 -1
View File
@@ -9,4 +9,5 @@ logs/
.idea
tools/
.vscode/
convert_*
convert_*
*.pt
+6 -2
View File
@@ -1,7 +1,11 @@
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
from .recammaster.nodes import NODE_CLASS_MAPPINGS as RECAM_MASTER_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS
from .unianimate.nodes import NODE_CLASS_MAPPINGS as UNIANIMATE_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS
NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS)
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+61 -6
View File
@@ -2,7 +2,7 @@ import os
import torch
import torch.nn.functional as F
import gc
from .utils import log, print_memory, apply_lora, clip_encode_image_tiled
from .utils import log, print_memory, apply_lora, clip_encode_image_tiled, fourier_filter
import numpy as np
import math
from tqdm import tqdm
@@ -205,6 +205,8 @@ def standardize_lora_key_format(lora_sd):
# Diffusers format
if k.startswith('transformer.'):
k = k.replace('transformer.', 'diffusion_model.')
if k.startswith('pipe.dit.'): #unianimate-dit/diffsynth
k = k.replace('pipe.dit.', 'diffusion_model.')
# Fun LoRA format
if k.startswith('lora_unet__'):
@@ -676,6 +678,11 @@ class WanVideoModelLoader:
lora_path = l["path"]
lora_strength = l["strength"]
lora_sd = load_torch_file(lora_path, safe_load=True)
if "dwpose_embedding.0.weight" in lora_sd: #unianimate
from .unianimate.nodes import update_transformer
log.info("Unianimate LoRA detected, patching model...")
transformer = update_transformer(transformer, lora_sd)
lora_sd = standardize_lora_key_format(lora_sd)
if l["blocks"]:
lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"])
@@ -1569,6 +1576,7 @@ class WanVideoImageToVideoEncode:
"fun_or_fl2v_model": ("BOOLEAN", {"default": True, "tooltip": "Enable when using official FLF2V or Fun model"}),
"temporal_mask": ("MASK", {"tooltip": "mask"}),
"extra_latents": ("LATENT", {"tooltip": "Extra latents to add to the input front, used for Skyreels A2 reference images"}),
"unianimate_poses": ("UNIANIMATEPOSES", {"tooltip": "Unianimate poses"}),
}
}
@@ -1578,7 +1586,8 @@ class WanVideoImageToVideoEncode:
CATEGORY = "WanVideoWrapper"
def process(self, vae, width, height, num_frames, clip_embeds, force_offload, noise_aug_strength,
start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_or_fl2v_model=False, temporal_mask=None, extra_latents=None):
start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_or_fl2v_model=False,
temporal_mask=None, extra_latents=None, unianimate_poses=None):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
@@ -1699,6 +1708,7 @@ class WanVideoImageToVideoEncode:
"end_image": resized_end_image if end_image is not None else None,
"fun_or_fl2v_model": fun_or_fl2v_model,
"has_ref": has_ref,
"unianimate_poses": unianimate_poses
}
return (image_embeds,)
@@ -2131,6 +2141,10 @@ class WanVideoExperimentalArgs:
"cfg_zero_star": ("BOOLEAN", {"default": False, "tooltip": "https://github.com/WeichenFan/CFG-Zero-star"}),
"use_zero_init": ("BOOLEAN", {"default": True}),
"zero_star_steps": ("INT", {"default": 0, "min": 0, "tooltip": "Steps to split self attention when using multiple prompts"}),
"use_fresca": ("BOOLEAN", {"default": False, "tooltip": "https://github.com/WikiChao/FreSca"}),
"fresca_scale_low": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
"fresca_scale_high": ("FLOAT", {"default": 1.25, "min": 0.0, "max": 10.0, "step": 0.01}),
"fresca_freq_cutoff": ("INT", {"default": 20, "min": 0, "max": 10000, "step": 1}),
},
}
@@ -2179,6 +2193,7 @@ class WanVideoSampler:
"loop_args": ("LOOPARGS", ),
"experimental_args": ("EXPERIMENTALARGS", ),
"sigmas": ("SIGMAS", ),
"unianimate_poses": ("UNIANIMATE_POSE", ),
}
}
@@ -2189,7 +2204,7 @@ class WanVideoSampler:
def process(self, model, text_embeds, image_embeds, shift, steps, cfg, seed, scheduler, riflex_freq_index,
force_offload=True, samples=None, feta_args=None, denoise_strength=1.0, context_options=None,
teacache_args=None, flowedit_args=None, batched_cfg=False, slg_args=None, rope_function="default", loop_args=None, experimental_args=None, sigmas=None):
teacache_args=None, flowedit_args=None, batched_cfg=False, slg_args=None, rope_function="default", loop_args=None, experimental_args=None, sigmas=None, unianimate_poses=None):
#assert not (context_options and teacache_args), "Context options cannot currently be used together with teacache."
patcher = model
model = model.model
@@ -2359,6 +2374,27 @@ class WanVideoSampler:
)
masked_video_latents_input = torch.zeros_like(noise)
image_cond = torch.cat([mask_latents, masked_video_latents_input], dim=0).to(device)
if unianimate_poses is not None:
transformer.dwpose_embedding.to(device)
transformer.randomref_embedding_pose.to(device)
dwpose_data = unianimate_poses["pose"]
dwpose_data = transformer.dwpose_embedding(
(torch.cat([dwpose_data[:,:,:1].repeat(1,1,3,1,1), dwpose_data], dim=2)
).to(device)).to(model["dtype"])
random_ref_dwpose_data = None
if image_cond is not None:
random_ref_dwpose = unianimate_poses["ref"]
random_ref_dwpose_data = transformer.randomref_embedding_pose(
random_ref_dwpose.to(device)#.permute(0,3,1,2)
).unsqueeze(2).to(model["dtype"]) # [1, 20, 104, 60]
image_cond += random_ref_dwpose_data.squeeze(0)
unianim_data = {
"dwpose": dwpose_data,
"random_ref": random_ref_dwpose_data
}
latent_video_length = noise.shape[1]
@@ -2571,7 +2607,7 @@ class WanVideoSampler:
drift_timesteps = torch.cat([drift_timesteps, torch.tensor([0]).to(drift_timesteps.device)]).to(drift_timesteps.device)
timesteps[-drift_steps:] = drift_timesteps[-drift_steps:]
use_cfg_zero_star = False
use_cfg_zero_star, use_fresca = False, False
if experimental_args is not None:
video_attention_split_steps = experimental_args.get("video_attention_split_steps", [])
if video_attention_split_steps:
@@ -2582,6 +2618,12 @@ class WanVideoSampler:
use_cfg_zero_star = experimental_args.get("cfg_zero_star", False)
zero_star_steps = experimental_args.get("zero_star_steps", 0)
use_fresca = experimental_args.get("use_fresca", False)
if use_fresca:
fresca_scale_low = experimental_args.get("fresca_scale_low", 1.0)
fresca_scale_high = experimental_args.get("fresca_scale_high", 1.25)
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):
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=model["dtype"], enabled=True):
@@ -2632,6 +2674,7 @@ class WanVideoSampler:
'control_lora_enabled': control_lora_enabled,
'vace_data': vace_data if vace_data is not None else None,
'camera_embed': camera_embed,
'unianim_data': unianim_data if unianimate_poses is not None else None,
}
batch_size = 1
@@ -2673,9 +2716,21 @@ class WanVideoSampler:
noise_pred_cond.view(batch_size, -1),
noise_pred_uncond.view(batch_size, -1)
).view(batch_size, 1, 1, 1)
noise_pred = noise_pred_uncond * alpha + cfg_scale * (noise_pred_cond - noise_pred_uncond * alpha)
else:
noise_pred = noise_pred_uncond + cfg_scale * (noise_pred_cond - noise_pred_uncond)
alpha = 1.0
#https://github.com/WikiChao/FreSca
if use_fresca:
filtered_cond = fourier_filter(
noise_pred_cond - noise_pred_uncond,
scale_low=fresca_scale_low,
scale_high=fresca_scale_high,
freq_cutoff=fresca_freq_cutoff,
)
noise_pred = noise_pred_uncond * alpha + cfg_scale * filtered_cond * alpha
else:
noise_pred = noise_pred_uncond * alpha + cfg_scale * (noise_pred_cond - noise_pred_uncond * alpha)
return noise_pred, [teacache_state_cond, teacache_state_uncond]
View File
+125
View File
@@ -0,0 +1,125 @@
import cv2
import numpy as np
import torch
def nms(boxes, scores, nms_thr):
"""Single class NMS implemented in Numpy."""
x1 = boxes[:, 0]
y1 = boxes[:, 1]
x2 = boxes[:, 2]
y2 = boxes[:, 3]
areas = (x2 - x1 + 1) * (y2 - y1 + 1)
order = scores.argsort()[::-1]
keep = []
while order.size > 0:
i = order[0]
keep.append(i)
xx1 = np.maximum(x1[i], x1[order[1:]])
yy1 = np.maximum(y1[i], y1[order[1:]])
xx2 = np.minimum(x2[i], x2[order[1:]])
yy2 = np.minimum(y2[i], y2[order[1:]])
w = np.maximum(0.0, xx2 - xx1 + 1)
h = np.maximum(0.0, yy2 - yy1 + 1)
inter = w * h
ovr = inter / (areas[i] + areas[order[1:]] - inter)
inds = np.where(ovr <= nms_thr)[0]
order = order[inds + 1]
return keep
def multiclass_nms(boxes, scores, nms_thr, score_thr):
"""Multiclass NMS implemented in Numpy. Class-aware version."""
final_dets = []
num_classes = scores.shape[1]
for cls_ind in range(num_classes):
cls_scores = scores[:, cls_ind]
valid_score_mask = cls_scores > score_thr
if valid_score_mask.sum() == 0:
continue
else:
valid_scores = cls_scores[valid_score_mask]
valid_boxes = boxes[valid_score_mask]
keep = nms(valid_boxes, valid_scores, nms_thr)
if len(keep) > 0:
cls_inds = np.ones((len(keep), 1)) * cls_ind
dets = np.concatenate(
[valid_boxes[keep], valid_scores[keep, None], cls_inds], 1
)
final_dets.append(dets)
if len(final_dets) == 0:
return None
return np.concatenate(final_dets, 0)
def demo_postprocess(outputs, img_size, p6=False):
grids = []
expanded_strides = []
strides = [8, 16, 32] if not p6 else [8, 16, 32, 64]
hsizes = [img_size[0] // stride for stride in strides]
wsizes = [img_size[1] // stride for stride in strides]
for hsize, wsize, stride in zip(hsizes, wsizes, strides):
xv, yv = np.meshgrid(np.arange(wsize), np.arange(hsize))
grid = np.stack((xv, yv), 2).reshape(1, -1, 2)
grids.append(grid)
shape = grid.shape[:2]
expanded_strides.append(np.full((*shape, 1), stride))
grids = np.concatenate(grids, 1)
expanded_strides = np.concatenate(expanded_strides, 1)
outputs[..., :2] = (outputs[..., :2] + grids) * expanded_strides
outputs[..., 2:4] = np.exp(outputs[..., 2:4]) * expanded_strides
return outputs
def preprocess(img, input_size, swap=(2, 0, 1)):
if len(img.shape) == 3:
padded_img = np.ones((input_size[0], input_size[1], 3), dtype=np.uint8) * 114
else:
padded_img = np.ones(input_size, dtype=np.uint8) * 114
r = min(input_size[0] / img.shape[0], input_size[1] / img.shape[1])
resized_img = cv2.resize(
img,
(int(img.shape[1] * r), int(img.shape[0] * r)),
interpolation=cv2.INTER_LINEAR,
).astype(np.uint8)
padded_img[: int(img.shape[0] * r), : int(img.shape[1] * r)] = resized_img
padded_img = padded_img.transpose(swap)
padded_img = np.ascontiguousarray(padded_img, dtype=np.float32)
return padded_img, r
def inference_detector(model, oriImg, detect_classes=[0]):
input_shape = (640,640)
img, ratio = preprocess(oriImg, input_shape)
device, dtype = next(model.parameters()).device, next(model.parameters()).dtype
input = img[None, :, :, :]
input = torch.from_numpy(input).to(device, dtype)
output = model(input).float().cpu().detach().numpy()
predictions = demo_postprocess(output[0], input_shape)
boxes = predictions[:, :4]
scores = predictions[:, 4:5] * predictions[:, 5:]
boxes_xyxy = np.ones_like(boxes)
boxes_xyxy[:, 0] = boxes[:, 0] - boxes[:, 2]/2.
boxes_xyxy[:, 1] = boxes[:, 1] - boxes[:, 3]/2.
boxes_xyxy[:, 2] = boxes[:, 0] + boxes[:, 2]/2.
boxes_xyxy[:, 3] = boxes[:, 1] + boxes[:, 3]/2.
boxes_xyxy /= ratio
dets = multiclass_nms(boxes_xyxy, scores, nms_thr=0.45, score_thr=0.1)
if dets is None:
return None
final_boxes, final_scores, final_cls_inds = dets[:, :4], dets[:, 4], dets[:, 5]
isscore = final_scores>0.3
iscat = np.isin(final_cls_inds, detect_classes)
isbbox = [ i and j for (i, j) in zip(isscore, iscat)]
final_boxes = final_boxes[isbbox]
return final_boxes
+363
View File
@@ -0,0 +1,363 @@
from typing import List, Tuple
import cv2
import numpy as np
import torch
def preprocess(
img: np.ndarray, out_bbox, input_size: Tuple[int, int] = (192, 256)
) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
"""Do preprocessing for DWPose model inference.
Args:
img (np.ndarray): Input image in shape.
input_size (tuple): Input image size in shape (w, h).
Returns:
tuple:
- resized_img (np.ndarray): Preprocessed image.
- center (np.ndarray): Center of image.
- scale (np.ndarray): Scale of image.
"""
# get shape of image
img_shape = img.shape[:2]
out_img, out_center, out_scale = [], [], []
if len(out_bbox) == 0:
out_bbox = [[0, 0, img_shape[1], img_shape[0]]]
for i in range(len(out_bbox)):
x0 = out_bbox[i][0]
y0 = out_bbox[i][1]
x1 = out_bbox[i][2]
y1 = out_bbox[i][3]
bbox = np.array([x0, y0, x1, y1])
# get center and scale
center, scale = bbox_xyxy2cs(bbox, padding=1.25)
# do affine transformation
resized_img, scale = top_down_affine(input_size, scale, center, img)
# normalize image
mean = np.array([123.675, 116.28, 103.53])
std = np.array([58.395, 57.12, 57.375])
resized_img = (resized_img - mean) / std
out_img.append(resized_img)
out_center.append(center)
out_scale.append(scale)
return out_img, out_center, out_scale
def inference(model, img, bs=5):
"""Inference DWPose model implemented in TorchScript.
Args:
model : TorchScript Model.
img : Input image in shape.
Returns:
outputs : Output of DWPose model.
"""
all_out = []
# build input
orig_img_count = len(img)
#Pad zeros to fit batch size
for _ in range(bs - (orig_img_count % bs)):
img.append(np.zeros_like(img[0]))
input = np.stack(img, axis=0).transpose(0, 3, 1, 2)
device, dtype = next(model.parameters()).device, next(model.parameters()).dtype
input = torch.from_numpy(input).to(device, dtype)
out1, out2 = [], []
for i in range(input.shape[0] // bs):
curr_batch_output = model(input[i*bs:(i+1)*bs])
out1.append(curr_batch_output[0].float())
out2.append(curr_batch_output[1].float())
out1, out2 = torch.cat(out1, dim=0)[:orig_img_count], torch.cat(out2, dim=0)[:orig_img_count]
out1, out2 = out1.float().cpu().detach().numpy(), out2.float().cpu().detach().numpy()
all_outputs = out1, out2
for batch_idx in range(len(all_outputs[0])):
outputs = [all_outputs[i][batch_idx:batch_idx+1,...] for i in range(len(all_outputs))]
all_out.append(outputs)
return all_out
def postprocess(outputs: List[np.ndarray],
model_input_size: Tuple[int, int],
center: Tuple[int, int],
scale: Tuple[int, int],
simcc_split_ratio: float = 2.0
) -> Tuple[np.ndarray, np.ndarray]:
"""Postprocess for DWPose model output.
Args:
outputs (np.ndarray): Output of RTMPose model.
model_input_size (tuple): RTMPose model Input image size.
center (tuple): Center of bbox in shape (x, y).
scale (tuple): Scale of bbox in shape (w, h).
simcc_split_ratio (float): Split ratio of simcc.
Returns:
tuple:
- keypoints (np.ndarray): Rescaled keypoints.
- scores (np.ndarray): Model predict scores.
"""
all_key = []
all_score = []
for i in range(len(outputs)):
# use simcc to decode
simcc_x, simcc_y = outputs[i]
keypoints, scores = decode(simcc_x, simcc_y, simcc_split_ratio)
# rescale keypoints
keypoints = keypoints / model_input_size * scale[i] + center[i] - scale[i] / 2
all_key.append(keypoints[0])
all_score.append(scores[0])
return np.array(all_key), np.array(all_score)
def bbox_xyxy2cs(bbox: np.ndarray,
padding: float = 1.) -> Tuple[np.ndarray, np.ndarray]:
"""Transform the bbox format from (x,y,w,h) into (center, scale)
Args:
bbox (ndarray): Bounding box(es) in shape (4,) or (n, 4), formatted
as (left, top, right, bottom)
padding (float): BBox padding factor that will be multilied to scale.
Default: 1.0
Returns:
tuple: A tuple containing center and scale.
- np.ndarray[float32]: Center (x, y) of the bbox in shape (2,) or
(n, 2)
- np.ndarray[float32]: Scale (w, h) of the bbox in shape (2,) or
(n, 2)
"""
# convert single bbox from (4, ) to (1, 4)
dim = bbox.ndim
if dim == 1:
bbox = bbox[None, :]
# get bbox center and scale
x1, y1, x2, y2 = np.hsplit(bbox, [1, 2, 3])
center = np.hstack([x1 + x2, y1 + y2]) * 0.5
scale = np.hstack([x2 - x1, y2 - y1]) * padding
if dim == 1:
center = center[0]
scale = scale[0]
return center, scale
def _fix_aspect_ratio(bbox_scale: np.ndarray,
aspect_ratio: float) -> np.ndarray:
"""Extend the scale to match the given aspect ratio.
Args:
scale (np.ndarray): The image scale (w, h) in shape (2, )
aspect_ratio (float): The ratio of ``w/h``
Returns:
np.ndarray: The reshaped image scale in (2, )
"""
w, h = np.hsplit(bbox_scale, [1])
bbox_scale = np.where(w > h * aspect_ratio,
np.hstack([w, w / aspect_ratio]),
np.hstack([h * aspect_ratio, h]))
return bbox_scale
def _rotate_point(pt: np.ndarray, angle_rad: float) -> np.ndarray:
"""Rotate a point by an angle.
Args:
pt (np.ndarray): 2D point coordinates (x, y) in shape (2, )
angle_rad (float): rotation angle in radian
Returns:
np.ndarray: Rotated point in shape (2, )
"""
sn, cs = np.sin(angle_rad), np.cos(angle_rad)
rot_mat = np.array([[cs, -sn], [sn, cs]])
return rot_mat @ pt
def _get_3rd_point(a: np.ndarray, b: np.ndarray) -> np.ndarray:
"""To calculate the affine matrix, three pairs of points are required. This
function is used to get the 3rd point, given 2D points a & b.
The 3rd point is defined by rotating vector `a - b` by 90 degrees
anticlockwise, using b as the rotation center.
Args:
a (np.ndarray): The 1st point (x,y) in shape (2, )
b (np.ndarray): The 2nd point (x,y) in shape (2, )
Returns:
np.ndarray: The 3rd point.
"""
direction = a - b
c = b + np.r_[-direction[1], direction[0]]
return c
def get_warp_matrix(center: np.ndarray,
scale: np.ndarray,
rot: float,
output_size: Tuple[int, int],
shift: Tuple[float, float] = (0., 0.),
inv: bool = False) -> np.ndarray:
"""Calculate the affine transformation matrix that can warp the bbox area
in the input image to the output size.
Args:
center (np.ndarray[2, ]): Center of the bounding box (x, y).
scale (np.ndarray[2, ]): Scale of the bounding box
wrt [width, height].
rot (float): Rotation angle (degree).
output_size (np.ndarray[2, ] | list(2,)): Size of the
destination heatmaps.
shift (0-100%): Shift translation ratio wrt the width/height.
Default (0., 0.).
inv (bool): Option to inverse the affine transform direction.
(inv=False: src->dst or inv=True: dst->src)
Returns:
np.ndarray: A 2x3 transformation matrix
"""
shift = np.array(shift)
src_w = scale[0]
dst_w = output_size[0]
dst_h = output_size[1]
# compute transformation matrix
rot_rad = np.deg2rad(rot)
src_dir = _rotate_point(np.array([0., src_w * -0.5]), rot_rad)
dst_dir = np.array([0., dst_w * -0.5])
# get four corners of the src rectangle in the original image
src = np.zeros((3, 2), dtype=np.float32)
src[0, :] = center + scale * shift
src[1, :] = center + src_dir + scale * shift
src[2, :] = _get_3rd_point(src[0, :], src[1, :])
# get four corners of the dst rectangle in the input image
dst = np.zeros((3, 2), dtype=np.float32)
dst[0, :] = [dst_w * 0.5, dst_h * 0.5]
dst[1, :] = np.array([dst_w * 0.5, dst_h * 0.5]) + dst_dir
dst[2, :] = _get_3rd_point(dst[0, :], dst[1, :])
if inv:
warp_mat = cv2.getAffineTransform(np.float32(dst), np.float32(src))
else:
warp_mat = cv2.getAffineTransform(np.float32(src), np.float32(dst))
return warp_mat
def top_down_affine(input_size: dict, bbox_scale: dict, bbox_center: dict,
img: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
"""Get the bbox image as the model input by affine transform.
Args:
input_size (dict): The input size of the model.
bbox_scale (dict): The bbox scale of the img.
bbox_center (dict): The bbox center of the img.
img (np.ndarray): The original image.
Returns:
tuple: A tuple containing center and scale.
- np.ndarray[float32]: img after affine transform.
- np.ndarray[float32]: bbox scale after affine transform.
"""
w, h = input_size
warp_size = (int(w), int(h))
# reshape bbox to fixed aspect ratio
bbox_scale = _fix_aspect_ratio(bbox_scale, aspect_ratio=w / h)
# get the affine matrix
center = bbox_center
scale = bbox_scale
rot = 0
warp_mat = get_warp_matrix(center, scale, rot, output_size=(w, h))
# do affine transform
img = cv2.warpAffine(img, warp_mat, warp_size, flags=cv2.INTER_LINEAR)
return img, bbox_scale
def get_simcc_maximum(simcc_x: np.ndarray,
simcc_y: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
"""Get maximum response location and value from simcc representations.
Note:
instance number: N
num_keypoints: K
heatmap height: H
heatmap width: W
Args:
simcc_x (np.ndarray): x-axis SimCC in shape (K, Wx) or (N, K, Wx)
simcc_y (np.ndarray): y-axis SimCC in shape (K, Wy) or (N, K, Wy)
Returns:
tuple:
- locs (np.ndarray): locations of maximum heatmap responses in shape
(K, 2) or (N, K, 2)
- vals (np.ndarray): values of maximum heatmap responses in shape
(K,) or (N, K)
"""
N, K, Wx = simcc_x.shape
simcc_x = simcc_x.reshape(N * K, -1)
simcc_y = simcc_y.reshape(N * K, -1)
# get maximum value locations
x_locs = np.argmax(simcc_x, axis=1)
y_locs = np.argmax(simcc_y, axis=1)
locs = np.stack((x_locs, y_locs), axis=-1).astype(np.float32)
max_val_x = np.amax(simcc_x, axis=1)
max_val_y = np.amax(simcc_y, axis=1)
# get maximum value across x and y axis
mask = max_val_x > max_val_y
max_val_x[mask] = max_val_y[mask]
vals = max_val_x
locs[vals <= 0.] = -1
# reshape
locs = locs.reshape(N, K, 2)
vals = vals.reshape(N, K)
return locs, vals
def decode(simcc_x: np.ndarray, simcc_y: np.ndarray,
simcc_split_ratio) -> Tuple[np.ndarray, np.ndarray]:
"""Modulate simcc distribution with Gaussian.
Args:
simcc_x (np.ndarray[K, Wx]): model predicted simcc in x.
simcc_y (np.ndarray[K, Wy]): model predicted simcc in y.
simcc_split_ratio (int): The split ratio of simcc.
Returns:
tuple: A tuple containing center and scale.
- np.ndarray[float32]: keypoints in shape (K, 2) or (n, K, 2)
- np.ndarray[float32]: scores in shape (K,) or (n, K)
"""
keypoints, scores = get_simcc_maximum(simcc_x, simcc_y)
keypoints /= simcc_split_ratio
return keypoints, scores
def inference_pose(model, out_bbox, oriImg, model_input_size=(288, 384)):
resized_img, center, scale = preprocess(oriImg, out_bbox, model_input_size)
#outputs = inference(session, resized_img, dtype)
outputs = inference(model, resized_img)
keypoints, scores = postprocess(outputs, model_input_size, center, scale)
return keypoints, scores
+127
View File
@@ -0,0 +1,127 @@
import cv2
import numpy as np
import onnxruntime
def nms(boxes, scores, nms_thr):
"""Single class NMS implemented in Numpy."""
x1 = boxes[:, 0]
y1 = boxes[:, 1]
x2 = boxes[:, 2]
y2 = boxes[:, 3]
areas = (x2 - x1 + 1) * (y2 - y1 + 1)
order = scores.argsort()[::-1]
keep = []
while order.size > 0:
i = order[0]
keep.append(i)
xx1 = np.maximum(x1[i], x1[order[1:]])
yy1 = np.maximum(y1[i], y1[order[1:]])
xx2 = np.minimum(x2[i], x2[order[1:]])
yy2 = np.minimum(y2[i], y2[order[1:]])
w = np.maximum(0.0, xx2 - xx1 + 1)
h = np.maximum(0.0, yy2 - yy1 + 1)
inter = w * h
ovr = inter / (areas[i] + areas[order[1:]] - inter)
inds = np.where(ovr <= nms_thr)[0]
order = order[inds + 1]
return keep
def multiclass_nms(boxes, scores, nms_thr, score_thr):
"""Multiclass NMS implemented in Numpy. Class-aware version."""
final_dets = []
num_classes = scores.shape[1]
for cls_ind in range(num_classes):
cls_scores = scores[:, cls_ind]
valid_score_mask = cls_scores > score_thr
if valid_score_mask.sum() == 0:
continue
else:
valid_scores = cls_scores[valid_score_mask]
valid_boxes = boxes[valid_score_mask]
keep = nms(valid_boxes, valid_scores, nms_thr)
if len(keep) > 0:
cls_inds = np.ones((len(keep), 1)) * cls_ind
dets = np.concatenate(
[valid_boxes[keep], valid_scores[keep, None], cls_inds], 1
)
final_dets.append(dets)
if len(final_dets) == 0:
return None
return np.concatenate(final_dets, 0)
def demo_postprocess(outputs, img_size, p6=False):
grids = []
expanded_strides = []
strides = [8, 16, 32] if not p6 else [8, 16, 32, 64]
hsizes = [img_size[0] // stride for stride in strides]
wsizes = [img_size[1] // stride for stride in strides]
for hsize, wsize, stride in zip(hsizes, wsizes, strides):
xv, yv = np.meshgrid(np.arange(wsize), np.arange(hsize))
grid = np.stack((xv, yv), 2).reshape(1, -1, 2)
grids.append(grid)
shape = grid.shape[:2]
expanded_strides.append(np.full((*shape, 1), stride))
grids = np.concatenate(grids, 1)
expanded_strides = np.concatenate(expanded_strides, 1)
outputs[..., :2] = (outputs[..., :2] + grids) * expanded_strides
outputs[..., 2:4] = np.exp(outputs[..., 2:4]) * expanded_strides
return outputs
def preprocess(img, input_size, swap=(2, 0, 1)):
if len(img.shape) == 3:
padded_img = np.ones((input_size[0], input_size[1], 3), dtype=np.uint8) * 114
else:
padded_img = np.ones(input_size, dtype=np.uint8) * 114
r = min(input_size[0] / img.shape[0], input_size[1] / img.shape[1])
resized_img = cv2.resize(
img,
(int(img.shape[1] * r), int(img.shape[0] * r)),
interpolation=cv2.INTER_LINEAR,
).astype(np.uint8)
padded_img[: int(img.shape[0] * r), : int(img.shape[1] * r)] = resized_img
padded_img = padded_img.transpose(swap)
padded_img = np.ascontiguousarray(padded_img, dtype=np.float32)
return padded_img, r
def inference_detector(session, oriImg):
input_shape = (640,640)
img, ratio = preprocess(oriImg, input_shape)
ort_inputs = {session.get_inputs()[0].name: img[None, :, :, :]}
output = session.run(None, ort_inputs)
predictions = demo_postprocess(output[0], input_shape)[0]
boxes = predictions[:, :4]
scores = predictions[:, 4:5] * predictions[:, 5:]
boxes_xyxy = np.ones_like(boxes)
boxes_xyxy[:, 0] = boxes[:, 0] - boxes[:, 2]/2.
boxes_xyxy[:, 1] = boxes[:, 1] - boxes[:, 3]/2.
boxes_xyxy[:, 2] = boxes[:, 0] + boxes[:, 2]/2.
boxes_xyxy[:, 3] = boxes[:, 1] + boxes[:, 3]/2.
boxes_xyxy /= ratio
dets = multiclass_nms(boxes_xyxy, scores, nms_thr=0.45, score_thr=0.1)
if dets is not None:
final_boxes, final_scores, final_cls_inds = dets[:, :4], dets[:, 4], dets[:, 5]
isscore = final_scores>0.3
iscat = final_cls_inds == 0
isbbox = [ i and j for (i, j) in zip(isscore, iscat)]
final_boxes = final_boxes[isbbox]
else:
final_boxes = np.array([])
return final_boxes
+360
View File
@@ -0,0 +1,360 @@
from typing import List, Tuple
import cv2
import numpy as np
import onnxruntime as ort
def preprocess(
img: np.ndarray, out_bbox, input_size: Tuple[int, int] = (192, 256)
) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
"""Do preprocessing for RTMPose model inference.
Args:
img (np.ndarray): Input image in shape.
input_size (tuple): Input image size in shape (w, h).
Returns:
tuple:
- resized_img (np.ndarray): Preprocessed image.
- center (np.ndarray): Center of image.
- scale (np.ndarray): Scale of image.
"""
# get shape of image
img_shape = img.shape[:2]
out_img, out_center, out_scale = [], [], []
if len(out_bbox) == 0:
out_bbox = [[0, 0, img_shape[1], img_shape[0]]]
for i in range(len(out_bbox)):
x0 = out_bbox[i][0]
y0 = out_bbox[i][1]
x1 = out_bbox[i][2]
y1 = out_bbox[i][3]
bbox = np.array([x0, y0, x1, y1])
# get center and scale
center, scale = bbox_xyxy2cs(bbox, padding=1.25)
# do affine transformation
resized_img, scale = top_down_affine(input_size, scale, center, img)
# normalize image
mean = np.array([123.675, 116.28, 103.53])
std = np.array([58.395, 57.12, 57.375])
resized_img = (resized_img - mean) / std
out_img.append(resized_img)
out_center.append(center)
out_scale.append(scale)
return out_img, out_center, out_scale
def inference(sess: ort.InferenceSession, img: np.ndarray) -> np.ndarray:
"""Inference RTMPose model.
Args:
sess (ort.InferenceSession): ONNXRuntime session.
img (np.ndarray): Input image in shape.
Returns:
outputs (np.ndarray): Output of RTMPose model.
"""
all_out = []
# build input
for i in range(len(img)):
input = [img[i].transpose(2, 0, 1)]
# build output
sess_input = {sess.get_inputs()[0].name: input}
sess_output = []
for out in sess.get_outputs():
sess_output.append(out.name)
# run model
outputs = sess.run(sess_output, sess_input)
all_out.append(outputs)
return all_out
def postprocess(outputs: List[np.ndarray],
model_input_size: Tuple[int, int],
center: Tuple[int, int],
scale: Tuple[int, int],
simcc_split_ratio: float = 2.0
) -> Tuple[np.ndarray, np.ndarray]:
"""Postprocess for RTMPose model output.
Args:
outputs (np.ndarray): Output of RTMPose model.
model_input_size (tuple): RTMPose model Input image size.
center (tuple): Center of bbox in shape (x, y).
scale (tuple): Scale of bbox in shape (w, h).
simcc_split_ratio (float): Split ratio of simcc.
Returns:
tuple:
- keypoints (np.ndarray): Rescaled keypoints.
- scores (np.ndarray): Model predict scores.
"""
all_key = []
all_score = []
for i in range(len(outputs)):
# use simcc to decode
simcc_x, simcc_y = outputs[i]
keypoints, scores = decode(simcc_x, simcc_y, simcc_split_ratio)
# rescale keypoints
keypoints = keypoints / model_input_size * scale[i] + center[i] - scale[i] / 2
all_key.append(keypoints[0])
all_score.append(scores[0])
return np.array(all_key), np.array(all_score)
def bbox_xyxy2cs(bbox: np.ndarray,
padding: float = 1.) -> Tuple[np.ndarray, np.ndarray]:
"""Transform the bbox format from (x,y,w,h) into (center, scale)
Args:
bbox (ndarray): Bounding box(es) in shape (4,) or (n, 4), formatted
as (left, top, right, bottom)
padding (float): BBox padding factor that will be multilied to scale.
Default: 1.0
Returns:
tuple: A tuple containing center and scale.
- np.ndarray[float32]: Center (x, y) of the bbox in shape (2,) or
(n, 2)
- np.ndarray[float32]: Scale (w, h) of the bbox in shape (2,) or
(n, 2)
"""
# convert single bbox from (4, ) to (1, 4)
dim = bbox.ndim
if dim == 1:
bbox = bbox[None, :]
# get bbox center and scale
x1, y1, x2, y2 = np.hsplit(bbox, [1, 2, 3])
center = np.hstack([x1 + x2, y1 + y2]) * 0.5
scale = np.hstack([x2 - x1, y2 - y1]) * padding
if dim == 1:
center = center[0]
scale = scale[0]
return center, scale
def _fix_aspect_ratio(bbox_scale: np.ndarray,
aspect_ratio: float) -> np.ndarray:
"""Extend the scale to match the given aspect ratio.
Args:
scale (np.ndarray): The image scale (w, h) in shape (2, )
aspect_ratio (float): The ratio of ``w/h``
Returns:
np.ndarray: The reshaped image scale in (2, )
"""
w, h = np.hsplit(bbox_scale, [1])
bbox_scale = np.where(w > h * aspect_ratio,
np.hstack([w, w / aspect_ratio]),
np.hstack([h * aspect_ratio, h]))
return bbox_scale
def _rotate_point(pt: np.ndarray, angle_rad: float) -> np.ndarray:
"""Rotate a point by an angle.
Args:
pt (np.ndarray): 2D point coordinates (x, y) in shape (2, )
angle_rad (float): rotation angle in radian
Returns:
np.ndarray: Rotated point in shape (2, )
"""
sn, cs = np.sin(angle_rad), np.cos(angle_rad)
rot_mat = np.array([[cs, -sn], [sn, cs]])
return rot_mat @ pt
def _get_3rd_point(a: np.ndarray, b: np.ndarray) -> np.ndarray:
"""To calculate the affine matrix, three pairs of points are required. This
function is used to get the 3rd point, given 2D points a & b.
The 3rd point is defined by rotating vector `a - b` by 90 degrees
anticlockwise, using b as the rotation center.
Args:
a (np.ndarray): The 1st point (x,y) in shape (2, )
b (np.ndarray): The 2nd point (x,y) in shape (2, )
Returns:
np.ndarray: The 3rd point.
"""
direction = a - b
c = b + np.r_[-direction[1], direction[0]]
return c
def get_warp_matrix(center: np.ndarray,
scale: np.ndarray,
rot: float,
output_size: Tuple[int, int],
shift: Tuple[float, float] = (0., 0.),
inv: bool = False) -> np.ndarray:
"""Calculate the affine transformation matrix that can warp the bbox area
in the input image to the output size.
Args:
center (np.ndarray[2, ]): Center of the bounding box (x, y).
scale (np.ndarray[2, ]): Scale of the bounding box
wrt [width, height].
rot (float): Rotation angle (degree).
output_size (np.ndarray[2, ] | list(2,)): Size of the
destination heatmaps.
shift (0-100%): Shift translation ratio wrt the width/height.
Default (0., 0.).
inv (bool): Option to inverse the affine transform direction.
(inv=False: src->dst or inv=True: dst->src)
Returns:
np.ndarray: A 2x3 transformation matrix
"""
shift = np.array(shift)
src_w = scale[0]
dst_w = output_size[0]
dst_h = output_size[1]
# compute transformation matrix
rot_rad = np.deg2rad(rot)
src_dir = _rotate_point(np.array([0., src_w * -0.5]), rot_rad)
dst_dir = np.array([0., dst_w * -0.5])
# get four corners of the src rectangle in the original image
src = np.zeros((3, 2), dtype=np.float32)
src[0, :] = center + scale * shift
src[1, :] = center + src_dir + scale * shift
src[2, :] = _get_3rd_point(src[0, :], src[1, :])
# get four corners of the dst rectangle in the input image
dst = np.zeros((3, 2), dtype=np.float32)
dst[0, :] = [dst_w * 0.5, dst_h * 0.5]
dst[1, :] = np.array([dst_w * 0.5, dst_h * 0.5]) + dst_dir
dst[2, :] = _get_3rd_point(dst[0, :], dst[1, :])
if inv:
warp_mat = cv2.getAffineTransform(np.float32(dst), np.float32(src))
else:
warp_mat = cv2.getAffineTransform(np.float32(src), np.float32(dst))
return warp_mat
def top_down_affine(input_size: dict, bbox_scale: dict, bbox_center: dict,
img: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
"""Get the bbox image as the model input by affine transform.
Args:
input_size (dict): The input size of the model.
bbox_scale (dict): The bbox scale of the img.
bbox_center (dict): The bbox center of the img.
img (np.ndarray): The original image.
Returns:
tuple: A tuple containing center and scale.
- np.ndarray[float32]: img after affine transform.
- np.ndarray[float32]: bbox scale after affine transform.
"""
w, h = input_size
warp_size = (int(w), int(h))
# reshape bbox to fixed aspect ratio
bbox_scale = _fix_aspect_ratio(bbox_scale, aspect_ratio=w / h)
# get the affine matrix
center = bbox_center
scale = bbox_scale
rot = 0
warp_mat = get_warp_matrix(center, scale, rot, output_size=(w, h))
# do affine transform
img = cv2.warpAffine(img, warp_mat, warp_size, flags=cv2.INTER_LINEAR)
return img, bbox_scale
def get_simcc_maximum(simcc_x: np.ndarray,
simcc_y: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
"""Get maximum response location and value from simcc representations.
Note:
instance number: N
num_keypoints: K
heatmap height: H
heatmap width: W
Args:
simcc_x (np.ndarray): x-axis SimCC in shape (K, Wx) or (N, K, Wx)
simcc_y (np.ndarray): y-axis SimCC in shape (K, Wy) or (N, K, Wy)
Returns:
tuple:
- locs (np.ndarray): locations of maximum heatmap responses in shape
(K, 2) or (N, K, 2)
- vals (np.ndarray): values of maximum heatmap responses in shape
(K,) or (N, K)
"""
N, K, Wx = simcc_x.shape
simcc_x = simcc_x.reshape(N * K, -1)
simcc_y = simcc_y.reshape(N * K, -1)
# get maximum value locations
x_locs = np.argmax(simcc_x, axis=1)
y_locs = np.argmax(simcc_y, axis=1)
locs = np.stack((x_locs, y_locs), axis=-1).astype(np.float32)
max_val_x = np.amax(simcc_x, axis=1)
max_val_y = np.amax(simcc_y, axis=1)
# get maximum value across x and y axis
mask = max_val_x > max_val_y
max_val_x[mask] = max_val_y[mask]
vals = max_val_x
locs[vals <= 0.] = -1
# reshape
locs = locs.reshape(N, K, 2)
vals = vals.reshape(N, K)
return locs, vals
def decode(simcc_x: np.ndarray, simcc_y: np.ndarray,
simcc_split_ratio) -> Tuple[np.ndarray, np.ndarray]:
"""Modulate simcc distribution with Gaussian.
Args:
simcc_x (np.ndarray[K, Wx]): model predicted simcc in x.
simcc_y (np.ndarray[K, Wy]): model predicted simcc in y.
simcc_split_ratio (int): The split ratio of simcc.
Returns:
tuple: A tuple containing center and scale.
- np.ndarray[float32]: keypoints in shape (K, 2) or (n, K, 2)
- np.ndarray[float32]: scores in shape (K,) or (n, K)
"""
keypoints, scores = get_simcc_maximum(simcc_x, simcc_y)
keypoints /= simcc_split_ratio
return keypoints, scores
def inference_pose(session, out_bbox, oriImg):
h, w = session.get_inputs()[0].shape[2:]
model_input_size = (w, h)
resized_img, center, scale = preprocess(oriImg, out_bbox, model_input_size)
outputs = inference(session, resized_img)
keypoints, scores = postprocess(outputs, model_input_size, center, scale)
return keypoints, scores
+336
View File
@@ -0,0 +1,336 @@
import math
import numpy as np
import matplotlib
import cv2
eps = 0.01
def smart_resize(x, s):
Ht, Wt = s
if x.ndim == 2:
Ho, Wo = x.shape
Co = 1
else:
Ho, Wo, Co = x.shape
if Co == 3 or Co == 1:
k = float(Ht + Wt) / float(Ho + Wo)
return cv2.resize(x, (int(Wt), int(Ht)), interpolation=cv2.INTER_AREA if k < 1 else cv2.INTER_LANCZOS4)
else:
return np.stack([smart_resize(x[:, :, i], s) for i in range(Co)], axis=2)
def smart_resize_k(x, fx, fy):
if x.ndim == 2:
Ho, Wo = x.shape
Co = 1
else:
Ho, Wo, Co = x.shape
Ht, Wt = Ho * fy, Wo * fx
if Co == 3 or Co == 1:
k = float(Ht + Wt) / float(Ho + Wo)
return cv2.resize(x, (int(Wt), int(Ht)), interpolation=cv2.INTER_AREA if k < 1 else cv2.INTER_LANCZOS4)
else:
return np.stack([smart_resize_k(x[:, :, i], fx, fy) for i in range(Co)], axis=2)
def padRightDownCorner(img, stride, padValue):
h = img.shape[0]
w = img.shape[1]
pad = 4 * [None]
pad[0] = 0 # up
pad[1] = 0 # left
pad[2] = 0 if (h % stride == 0) else stride - (h % stride) # down
pad[3] = 0 if (w % stride == 0) else stride - (w % stride) # right
img_padded = img
pad_up = np.tile(img_padded[0:1, :, :]*0 + padValue, (pad[0], 1, 1))
img_padded = np.concatenate((pad_up, img_padded), axis=0)
pad_left = np.tile(img_padded[:, 0:1, :]*0 + padValue, (1, pad[1], 1))
img_padded = np.concatenate((pad_left, img_padded), axis=1)
pad_down = np.tile(img_padded[-2:-1, :, :]*0 + padValue, (pad[2], 1, 1))
img_padded = np.concatenate((img_padded, pad_down), axis=0)
pad_right = np.tile(img_padded[:, -2:-1, :]*0 + padValue, (1, pad[3], 1))
img_padded = np.concatenate((img_padded, pad_right), axis=1)
return img_padded, pad
def transfer(model, model_weights):
transfered_model_weights = {}
for weights_name in model.state_dict().keys():
transfered_model_weights[weights_name] = model_weights['.'.join(weights_name.split('.')[1:])]
return transfered_model_weights
def draw_bodypose(canvas, candidate, subset):
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], [3, 17], [6, 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]]
for i in range(17):
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])
canvas = (canvas * 0.6).astype(np.uint8)
for i in range(18):
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_body_and_foot(canvas, candidate, subset):
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]]
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])
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)
return canvas
def draw_handpose(canvas, all_hand_peaks):
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], \
[10, 11], [11, 12], [0, 13], [13, 14], [14, 15], [15, 16], [0, 17], [17, 18], [18, 19], [19, 20]]
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)
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)
return canvas
def draw_facepose(canvas, all_lmks):
H, W, C = canvas.shape
for lmks in all_lmks:
lmks = np.array(lmks)
for lmk in lmks:
x, y = lmk
x = int(x * W)
y = int(y * H)
if x > eps and y > eps:
cv2.circle(canvas, (x, y), 3, (255, 255, 255), thickness=-1)
return canvas
# detect hand according to body pose keypoints
# please refer to https://github.com/CMU-Perceptual-Computing-Lab/openpose/blob/master/src/openpose/hand/handDetector.cpp
def handDetect(candidate, subset, oriImg):
# right hand: wrist 4, elbow 3, shoulder 2
# left hand: wrist 7, elbow 6, shoulder 5
ratioWristElbow = 0.33
detect_result = []
image_height, image_width = oriImg.shape[0:2]
for person in subset.astype(int):
# if any of three not detected
has_left = np.sum(person[[5, 6, 7]] == -1) == 0
has_right = np.sum(person[[2, 3, 4]] == -1) == 0
if not (has_left or has_right):
continue
hands = []
#left hand
if has_left:
left_shoulder_index, left_elbow_index, left_wrist_index = person[[5, 6, 7]]
x1, y1 = candidate[left_shoulder_index][:2]
x2, y2 = candidate[left_elbow_index][:2]
x3, y3 = candidate[left_wrist_index][:2]
hands.append([x1, y1, x2, y2, x3, y3, True])
# right hand
if has_right:
right_shoulder_index, right_elbow_index, right_wrist_index = person[[2, 3, 4]]
x1, y1 = candidate[right_shoulder_index][:2]
x2, y2 = candidate[right_elbow_index][:2]
x3, y3 = candidate[right_wrist_index][:2]
hands.append([x1, y1, x2, y2, x3, y3, False])
for x1, y1, x2, y2, x3, y3, is_left in hands:
x = x3 + ratioWristElbow * (x3 - x2)
y = y3 + ratioWristElbow * (y3 - y2)
distanceWristElbow = math.sqrt((x3 - x2) ** 2 + (y3 - y2) ** 2)
distanceElbowShoulder = math.sqrt((x2 - x1) ** 2 + (y2 - y1) ** 2)
width = 1.5 * max(distanceWristElbow, 0.9 * distanceElbowShoulder)
# x-y refers to the center --> offset to topLeft point
# handRectangle.x -= handRectangle.width / 2.f;
# handRectangle.y -= handRectangle.height / 2.f;
x -= width / 2
y -= width / 2 # width = height
# overflow the image
if x < 0: x = 0
if y < 0: y = 0
width1 = width
width2 = width
if x + width > image_width: width1 = image_width - x
if y + width > image_height: width2 = image_height - y
width = min(width1, width2)
# the max hand box value is 20 pixels
if width >= 20:
detect_result.append([int(x), int(y), int(width), is_left])
'''
return value: [[x, y, w, True if left hand else False]].
width=height since the network require squared input.
x, y is the coordinate of top left
'''
return detect_result
# Written by Lvmin
def faceDetect(candidate, subset, oriImg):
# left right eye ear 14 15 16 17
detect_result = []
image_height, image_width = oriImg.shape[0:2]
for person in subset.astype(int):
has_head = person[0] > -1
if not has_head:
continue
has_left_eye = person[14] > -1
has_right_eye = person[15] > -1
has_left_ear = person[16] > -1
has_right_ear = person[17] > -1
if not (has_left_eye or has_right_eye or has_left_ear or has_right_ear):
continue
head, left_eye, right_eye, left_ear, right_ear = person[[0, 14, 15, 16, 17]]
width = 0.0
x0, y0 = candidate[head][:2]
if has_left_eye:
x1, y1 = candidate[left_eye][:2]
d = max(abs(x0 - x1), abs(y0 - y1))
width = max(width, d * 3.0)
if has_right_eye:
x1, y1 = candidate[right_eye][:2]
d = max(abs(x0 - x1), abs(y0 - y1))
width = max(width, d * 3.0)
if has_left_ear:
x1, y1 = candidate[left_ear][:2]
d = max(abs(x0 - x1), abs(y0 - y1))
width = max(width, d * 1.5)
if has_right_ear:
x1, y1 = candidate[right_ear][:2]
d = max(abs(x0 - x1), abs(y0 - y1))
width = max(width, d * 1.5)
x, y = x0, y0
x -= width
y -= width
if x < 0:
x = 0
if y < 0:
y = 0
width1 = width * 2
width2 = width * 2
if x + width > image_width:
width1 = image_width - x
if y + width > image_height:
width2 = image_height - y
width = min(width1, width2)
if width >= 20:
detect_result.append([int(x), int(y), int(width)])
return detect_result
# get max index of 2d array
def npmax(array):
arrayindex = array.argmax(1)
arrayvalue = array.max(1)
i = arrayvalue.argmax()
j = arrayindex[i]
return i, j
+42
View File
@@ -0,0 +1,42 @@
import numpy as np
from .jit_det import inference_detector as inference_jit_yolox
from .jit_pose import inference_pose as inference_jit_pose
import os
class Wholebody:
def __init__(self, model_det, model_pose):
self.model_det = model_det
self.model_pose = model_pose
def __call__(self, oriImg):
det_result = inference_jit_yolox(self.model_det, oriImg, detect_classes=[0])
keypoints, scores = inference_jit_pose(self.model_pose, det_result, oriImg)
keypoints_info = np.concatenate(
(keypoints, scores[..., None]), axis=-1)
# compute neck joint
neck = np.mean(keypoints_info[:, [5, 6]], axis=1)
# neck score when visualizing pred
neck[:, 2:4] = np.logical_and(
keypoints_info[:, 5, 2:4] > 0.3,
keypoints_info[:, 6, 2:4] > 0.3).astype(int)
new_keypoints_info = np.insert(
keypoints_info, 17, neck, axis=1)
mmpose_idx = [
17, 6, 8, 10, 7, 9, 12, 14, 16, 13, 15, 2, 1, 4, 3
]
openpose_idx = [
1, 2, 3, 4, 6, 7, 8, 9, 10, 12, 13, 14, 15, 16, 17
]
new_keypoints_info[:, openpose_idx] = \
new_keypoints_info[:, mmpose_idx]
keypoints_info = new_keypoints_info
keypoints, scores = keypoints_info[
..., :2], keypoints_info[..., 2]
return keypoints, scores
+784
View File
@@ -0,0 +1,784 @@
import torch.nn as nn
from ..utils import log
import comfy.model_management as mm
from comfy.utils import ProgressBar
from tqdm import tqdm
def update_transformer(transformer, state_dict):
concat_dim = 4
transformer.dwpose_embedding = nn.Sequential(
nn.Conv3d(3, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)),
nn.SiLU(),
nn.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)),
nn.SiLU(),
nn.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,1,1), padding=(1,1,1)),
nn.SiLU(),
nn.Conv3d(concat_dim * 4, concat_dim * 4, (3,3,3), stride=(1,2,2), padding=(1,1,1)),
nn.SiLU(),
nn.Conv3d(concat_dim * 4, concat_dim * 4, 3, stride=(2,2,2), padding=1),
nn.SiLU(),
nn.Conv3d(concat_dim * 4, concat_dim * 4, 3, stride=(2,2,2), padding=1),
nn.SiLU(),
nn.Conv3d(concat_dim * 4, 5120, (1,2,2), stride=(1,2,2), padding=0))
randomref_dim = 20
transformer.randomref_embedding_pose = nn.Sequential(
nn.Conv2d(3, concat_dim * 4, 3, stride=1, padding=1),
nn.SiLU(),
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=1, padding=1),
nn.SiLU(),
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=1, padding=1),
nn.SiLU(),
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=2, padding=1),
nn.SiLU(),
nn.Conv2d(concat_dim * 4, concat_dim * 4, 3, stride=2, padding=1),
nn.SiLU(),
nn.Conv2d(concat_dim * 4, randomref_dim, 3, stride=2, padding=1),
)
state_dict_new = {}
for key in list(state_dict.keys()):
if "dwpose_embedding" in key:
state_dict_new[key.split("dwpose_embedding.")[1]] = state_dict.pop(key)
transformer.dwpose_embedding.load_state_dict(state_dict_new, strict=True)
state_dict_new = {}
for key in list(state_dict.keys()):
if "randomref_embedding_pose" in key:
state_dict_new[key.split("randomref_embedding_pose.")[1]] = state_dict.pop(key)
transformer.randomref_embedding_pose.load_state_dict(state_dict_new,strict=True)
return transformer
# Openpose
# Original from CMU https://github.com/CMU-Perceptual-Computing-Lab/openpose
# 2nd Edited by https://github.com/Hzzone/pytorch-openpose
# 3rd Edited by ControlNet
# 4th Edited by ControlNet (added face and correct hands)
import os
import torch
import numpy as np
import copy
import torch
import numpy as np
import math
from .dwpose.util import draw_body_and_foot, draw_handpose, draw_facepose
from .dwpose.wholebody import Wholebody
def smoothing_factor(t_e, cutoff):
r = 2 * math.pi * cutoff * t_e
return r / (r + 1)
def exponential_smoothing(a, x, x_prev):
return a * x + (1 - a) * x_prev
class OneEuroFilter:
def __init__(self, t0, x0, dx0=0.0, min_cutoff=1.0, beta=0.0,
d_cutoff=1.0):
"""Initialize the one euro filter."""
# The parameters.
self.min_cutoff = float(min_cutoff)
self.beta = float(beta)
self.d_cutoff = float(d_cutoff)
# Previous values.
self.x_prev = x0
self.dx_prev = float(dx0)
self.t_prev = float(t0)
def __call__(self, t, x):
"""Compute the filtered signal."""
t_e = t - self.t_prev
# The filtered derivative of the signal.
a_d = smoothing_factor(t_e, self.d_cutoff)
dx = (x - self.x_prev) / t_e
dx_hat = exponential_smoothing(a_d, dx, self.dx_prev)
# The filtered signal.
cutoff = self.min_cutoff + self.beta * abs(dx_hat)
a = smoothing_factor(t_e, cutoff)
x_hat = exponential_smoothing(a, x, self.x_prev)
# Memorize the previous values.
self.x_prev = x_hat
self.dx_prev = dx_hat
self.t_prev = t
return x_hat
class DWposeDetector:
def __init__(self, model_det, model_pose):
self.pose_estimation = Wholebody(model_det, model_pose)
def __call__(self, oriImg):
oriImg = oriImg.copy()
H, W, C = oriImg.shape
with torch.no_grad():
candidate, subset = self.pose_estimation(oriImg)
candidate = candidate[0][np.newaxis, :, :]
subset = subset[0][np.newaxis, :]
nums, keys, locs = candidate.shape
candidate[..., 0] /= float(W)
candidate[..., 1] /= float(H)
body = candidate[:,:18].copy()
body = body.reshape(nums*18, locs)
score = subset[:,:18].copy()
for i in range(len(score)):
for j in range(len(score[i])):
if score[i][j] > 0.3:
score[i][j] = int(18*i+j)
else:
score[i][j] = -1
un_visible = subset<0.3
candidate[un_visible] = -1
bodyfoot_score = subset[:,:24].copy()
for i in range(len(bodyfoot_score)):
for j in range(len(bodyfoot_score[i])):
if bodyfoot_score[i][j] > 0.3:
bodyfoot_score[i][j] = int(18*i+j)
else:
bodyfoot_score[i][j] = -1
if -1 not in bodyfoot_score[:,18] and -1 not in bodyfoot_score[:,19]:
bodyfoot_score[:,18] = np.array([18.])
else:
bodyfoot_score[:,18] = np.array([-1.])
if -1 not in bodyfoot_score[:,21] and -1 not in bodyfoot_score[:,22]:
bodyfoot_score[:,19] = np.array([19.])
else:
bodyfoot_score[:,19] = np.array([-1.])
bodyfoot_score = bodyfoot_score[:, :20]
bodyfoot = candidate[:,:24].copy()
for i in range(nums):
if -1 not in bodyfoot[i][18] and -1 not in bodyfoot[i][19]:
bodyfoot[i][18] = (bodyfoot[i][18]+bodyfoot[i][19])/2
else:
bodyfoot[i][18] = np.array([-1., -1.])
if -1 not in bodyfoot[i][21] and -1 not in bodyfoot[i][22]:
bodyfoot[i][19] = (bodyfoot[i][21]+bodyfoot[i][22])/2
else:
bodyfoot[i][19] = np.array([-1., -1.])
bodyfoot = bodyfoot[:,:20,:]
bodyfoot = bodyfoot.reshape(nums*20, locs)
foot = candidate[:,18:24]
faces = candidate[:,24:92]
hands = candidate[:,92:113]
hands = np.vstack([hands, candidate[:,113:]])
# bodies = dict(candidate=body, subset=score)
bodies = dict(candidate=bodyfoot, subset=bodyfoot_score)
pose = dict(bodies=bodies, hands=hands, faces=faces)
# return draw_pose(pose, H, W)
return pose
def draw_pose(pose, H, W):
bodies = pose['bodies']
faces = pose['faces']
hands = pose['hands']
candidate = bodies['candidate']
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_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):
results_vis = []
comfy_pbar = ProgressBar(len(pose_images))
try:
pose_ref = dwpose_model(ref_image.squeeze(0))
except:
raise ValueError("No pose detected in reference image")
for img in tqdm(pose_images, desc="Pose Extraction", unit="image", total=len(pose_images)):
try:
pose = dwpose_model(img)
except:
pose = torch.zeros_like(img)
results_vis.append(pose)
comfy_pbar.update(1)
dwpose_woface, dwpose_wface = draw_pose(pose_ref, H=height, W=width)
bodies = results_vis[0]['bodies']
faces = results_vis[0]['faces']
hands = results_vis[0]['hands']
candidate = bodies['candidate']
ref_bodies = pose_ref['bodies']
ref_faces = pose_ref['faces']
ref_hands = pose_ref['hands']
ref_candidate = ref_bodies['candidate']
ref_2_x = ref_candidate[2][0]
ref_2_y = ref_candidate[2][1]
ref_5_x = ref_candidate[5][0]
ref_5_y = ref_candidate[5][1]
ref_8_x = ref_candidate[8][0]
ref_8_y = ref_candidate[8][1]
ref_11_x = ref_candidate[11][0]
ref_11_y = ref_candidate[11][1]
ref_center1 = 0.5*(ref_candidate[2]+ref_candidate[5])
ref_center2 = 0.5*(ref_candidate[8]+ref_candidate[11])
zero_2_x = candidate[2][0]
zero_2_y = candidate[2][1]
zero_5_x = candidate[5][0]
zero_5_y = candidate[5][1]
zero_8_x = candidate[8][0]
zero_8_y = candidate[8][1]
zero_11_x = candidate[11][0]
zero_11_y = candidate[11][1]
zero_center1 = 0.5*(candidate[2]+candidate[5])
zero_center2 = 0.5*(candidate[8]+candidate[11])
x_ratio = (ref_5_x-ref_2_x)/(zero_5_x-zero_2_x)
y_ratio = (ref_center2[1]-ref_center1[1])/(zero_center2[1]-zero_center1[1])
results_vis[0]['bodies']['candidate'][:,0] *= x_ratio
results_vis[0]['bodies']['candidate'][:,1] *= y_ratio
results_vis[0]['faces'][:,:,0] *= x_ratio
results_vis[0]['faces'][:,:,1] *= y_ratio
results_vis[0]['hands'][:,:,0] *= x_ratio
results_vis[0]['hands'][:,:,1] *= y_ratio
########neck########
l_neck_ref = ((ref_candidate[0][0] - ref_candidate[1][0]) ** 2 + (ref_candidate[0][1] - ref_candidate[1][1]) ** 2) ** 0.5
l_neck_0 = ((candidate[0][0] - candidate[1][0]) ** 2 + (candidate[0][1] - candidate[1][1]) ** 2) ** 0.5
neck_ratio = l_neck_ref / l_neck_0
x_offset_neck = (candidate[1][0]-candidate[0][0])*(1.-neck_ratio)
y_offset_neck = (candidate[1][1]-candidate[0][1])*(1.-neck_ratio)
results_vis[0]['bodies']['candidate'][0,0] += x_offset_neck
results_vis[0]['bodies']['candidate'][0,1] += y_offset_neck
results_vis[0]['bodies']['candidate'][14,0] += x_offset_neck
results_vis[0]['bodies']['candidate'][14,1] += y_offset_neck
results_vis[0]['bodies']['candidate'][15,0] += x_offset_neck
results_vis[0]['bodies']['candidate'][15,1] += y_offset_neck
results_vis[0]['bodies']['candidate'][16,0] += x_offset_neck
results_vis[0]['bodies']['candidate'][16,1] += y_offset_neck
results_vis[0]['bodies']['candidate'][17,0] += x_offset_neck
results_vis[0]['bodies']['candidate'][17,1] += y_offset_neck
########shoulder2########
l_shoulder2_ref = ((ref_candidate[2][0] - ref_candidate[1][0]) ** 2 + (ref_candidate[2][1] - ref_candidate[1][1]) ** 2) ** 0.5
l_shoulder2_0 = ((candidate[2][0] - candidate[1][0]) ** 2 + (candidate[2][1] - candidate[1][1]) ** 2) ** 0.5
shoulder2_ratio = l_shoulder2_ref / l_shoulder2_0
x_offset_shoulder2 = (candidate[1][0]-candidate[2][0])*(1.-shoulder2_ratio)
y_offset_shoulder2 = (candidate[1][1]-candidate[2][1])*(1.-shoulder2_ratio)
results_vis[0]['bodies']['candidate'][2,0] += x_offset_shoulder2
results_vis[0]['bodies']['candidate'][2,1] += y_offset_shoulder2
results_vis[0]['bodies']['candidate'][3,0] += x_offset_shoulder2
results_vis[0]['bodies']['candidate'][3,1] += y_offset_shoulder2
results_vis[0]['bodies']['candidate'][4,0] += x_offset_shoulder2
results_vis[0]['bodies']['candidate'][4,1] += y_offset_shoulder2
results_vis[0]['hands'][1,:,0] += x_offset_shoulder2
results_vis[0]['hands'][1,:,1] += y_offset_shoulder2
########shoulder5########
l_shoulder5_ref = ((ref_candidate[5][0] - ref_candidate[1][0]) ** 2 + (ref_candidate[5][1] - ref_candidate[1][1]) ** 2) ** 0.5
l_shoulder5_0 = ((candidate[5][0] - candidate[1][0]) ** 2 + (candidate[5][1] - candidate[1][1]) ** 2) ** 0.5
shoulder5_ratio = l_shoulder5_ref / l_shoulder5_0
x_offset_shoulder5 = (candidate[1][0]-candidate[5][0])*(1.-shoulder5_ratio)
y_offset_shoulder5 = (candidate[1][1]-candidate[5][1])*(1.-shoulder5_ratio)
results_vis[0]['bodies']['candidate'][5,0] += x_offset_shoulder5
results_vis[0]['bodies']['candidate'][5,1] += y_offset_shoulder5
results_vis[0]['bodies']['candidate'][6,0] += x_offset_shoulder5
results_vis[0]['bodies']['candidate'][6,1] += y_offset_shoulder5
results_vis[0]['bodies']['candidate'][7,0] += x_offset_shoulder5
results_vis[0]['bodies']['candidate'][7,1] += y_offset_shoulder5
results_vis[0]['hands'][0,:,0] += x_offset_shoulder5
results_vis[0]['hands'][0,:,1] += y_offset_shoulder5
########arm3########
l_arm3_ref = ((ref_candidate[3][0] - ref_candidate[2][0]) ** 2 + (ref_candidate[3][1] - ref_candidate[2][1]) ** 2) ** 0.5
l_arm3_0 = ((candidate[3][0] - candidate[2][0]) ** 2 + (candidate[3][1] - candidate[2][1]) ** 2) ** 0.5
arm3_ratio = l_arm3_ref / l_arm3_0
x_offset_arm3 = (candidate[2][0]-candidate[3][0])*(1.-arm3_ratio)
y_offset_arm3 = (candidate[2][1]-candidate[3][1])*(1.-arm3_ratio)
results_vis[0]['bodies']['candidate'][3,0] += x_offset_arm3
results_vis[0]['bodies']['candidate'][3,1] += y_offset_arm3
results_vis[0]['bodies']['candidate'][4,0] += x_offset_arm3
results_vis[0]['bodies']['candidate'][4,1] += y_offset_arm3
results_vis[0]['hands'][1,:,0] += x_offset_arm3
results_vis[0]['hands'][1,:,1] += y_offset_arm3
########arm4########
l_arm4_ref = ((ref_candidate[4][0] - ref_candidate[3][0]) ** 2 + (ref_candidate[4][1] - ref_candidate[3][1]) ** 2) ** 0.5
l_arm4_0 = ((candidate[4][0] - candidate[3][0]) ** 2 + (candidate[4][1] - candidate[3][1]) ** 2) ** 0.5
arm4_ratio = l_arm4_ref / l_arm4_0
x_offset_arm4 = (candidate[3][0]-candidate[4][0])*(1.-arm4_ratio)
y_offset_arm4 = (candidate[3][1]-candidate[4][1])*(1.-arm4_ratio)
results_vis[0]['bodies']['candidate'][4,0] += x_offset_arm4
results_vis[0]['bodies']['candidate'][4,1] += y_offset_arm4
results_vis[0]['hands'][1,:,0] += x_offset_arm4
results_vis[0]['hands'][1,:,1] += y_offset_arm4
########arm6########
l_arm6_ref = ((ref_candidate[6][0] - ref_candidate[5][0]) ** 2 + (ref_candidate[6][1] - ref_candidate[5][1]) ** 2) ** 0.5
l_arm6_0 = ((candidate[6][0] - candidate[5][0]) ** 2 + (candidate[6][1] - candidate[5][1]) ** 2) ** 0.5
arm6_ratio = l_arm6_ref / l_arm6_0
x_offset_arm6 = (candidate[5][0]-candidate[6][0])*(1.-arm6_ratio)
y_offset_arm6 = (candidate[5][1]-candidate[6][1])*(1.-arm6_ratio)
results_vis[0]['bodies']['candidate'][6,0] += x_offset_arm6
results_vis[0]['bodies']['candidate'][6,1] += y_offset_arm6
results_vis[0]['bodies']['candidate'][7,0] += x_offset_arm6
results_vis[0]['bodies']['candidate'][7,1] += y_offset_arm6
results_vis[0]['hands'][0,:,0] += x_offset_arm6
results_vis[0]['hands'][0,:,1] += y_offset_arm6
########arm7########
l_arm7_ref = ((ref_candidate[7][0] - ref_candidate[6][0]) ** 2 + (ref_candidate[7][1] - ref_candidate[6][1]) ** 2) ** 0.5
l_arm7_0 = ((candidate[7][0] - candidate[6][0]) ** 2 + (candidate[7][1] - candidate[6][1]) ** 2) ** 0.5
arm7_ratio = l_arm7_ref / l_arm7_0
x_offset_arm7 = (candidate[6][0]-candidate[7][0])*(1.-arm7_ratio)
y_offset_arm7 = (candidate[6][1]-candidate[7][1])*(1.-arm7_ratio)
results_vis[0]['bodies']['candidate'][7,0] += x_offset_arm7
results_vis[0]['bodies']['candidate'][7,1] += y_offset_arm7
results_vis[0]['hands'][0,:,0] += x_offset_arm7
results_vis[0]['hands'][0,:,1] += y_offset_arm7
########head14########
l_head14_ref = ((ref_candidate[14][0] - ref_candidate[0][0]) ** 2 + (ref_candidate[14][1] - ref_candidate[0][1]) ** 2) ** 0.5
l_head14_0 = ((candidate[14][0] - candidate[0][0]) ** 2 + (candidate[14][1] - candidate[0][1]) ** 2) ** 0.5
head14_ratio = l_head14_ref / l_head14_0
x_offset_head14 = (candidate[0][0]-candidate[14][0])*(1.-head14_ratio)
y_offset_head14 = (candidate[0][1]-candidate[14][1])*(1.-head14_ratio)
results_vis[0]['bodies']['candidate'][14,0] += x_offset_head14
results_vis[0]['bodies']['candidate'][14,1] += y_offset_head14
results_vis[0]['bodies']['candidate'][16,0] += x_offset_head14
results_vis[0]['bodies']['candidate'][16,1] += y_offset_head14
########head15########
l_head15_ref = ((ref_candidate[15][0] - ref_candidate[0][0]) ** 2 + (ref_candidate[15][1] - ref_candidate[0][1]) ** 2) ** 0.5
l_head15_0 = ((candidate[15][0] - candidate[0][0]) ** 2 + (candidate[15][1] - candidate[0][1]) ** 2) ** 0.5
head15_ratio = l_head15_ref / l_head15_0
x_offset_head15 = (candidate[0][0]-candidate[15][0])*(1.-head15_ratio)
y_offset_head15 = (candidate[0][1]-candidate[15][1])*(1.-head15_ratio)
results_vis[0]['bodies']['candidate'][15,0] += x_offset_head15
results_vis[0]['bodies']['candidate'][15,1] += y_offset_head15
results_vis[0]['bodies']['candidate'][17,0] += x_offset_head15
results_vis[0]['bodies']['candidate'][17,1] += y_offset_head15
########head16########
l_head16_ref = ((ref_candidate[16][0] - ref_candidate[14][0]) ** 2 + (ref_candidate[16][1] - ref_candidate[14][1]) ** 2) ** 0.5
l_head16_0 = ((candidate[16][0] - candidate[14][0]) ** 2 + (candidate[16][1] - candidate[14][1]) ** 2) ** 0.5
head16_ratio = l_head16_ref / l_head16_0
x_offset_head16 = (candidate[14][0]-candidate[16][0])*(1.-head16_ratio)
y_offset_head16 = (candidate[14][1]-candidate[16][1])*(1.-head16_ratio)
results_vis[0]['bodies']['candidate'][16,0] += x_offset_head16
results_vis[0]['bodies']['candidate'][16,1] += y_offset_head16
########head17########
l_head17_ref = ((ref_candidate[17][0] - ref_candidate[15][0]) ** 2 + (ref_candidate[17][1] - ref_candidate[15][1]) ** 2) ** 0.5
l_head17_0 = ((candidate[17][0] - candidate[15][0]) ** 2 + (candidate[17][1] - candidate[15][1]) ** 2) ** 0.5
head17_ratio = l_head17_ref / l_head17_0
x_offset_head17 = (candidate[15][0]-candidate[17][0])*(1.-head17_ratio)
y_offset_head17 = (candidate[15][1]-candidate[17][1])*(1.-head17_ratio)
results_vis[0]['bodies']['candidate'][17,0] += x_offset_head17
results_vis[0]['bodies']['candidate'][17,1] += y_offset_head17
########MovingAverage########
########left leg########
l_ll1_ref = ((ref_candidate[8][0] - ref_candidate[9][0]) ** 2 + (ref_candidate[8][1] - ref_candidate[9][1]) ** 2) ** 0.5
l_ll1_0 = ((candidate[8][0] - candidate[9][0]) ** 2 + (candidate[8][1] - candidate[9][1]) ** 2) ** 0.5
ll1_ratio = l_ll1_ref / l_ll1_0
x_offset_ll1 = (candidate[9][0]-candidate[8][0])*(ll1_ratio-1.)
y_offset_ll1 = (candidate[9][1]-candidate[8][1])*(ll1_ratio-1.)
results_vis[0]['bodies']['candidate'][9,0] += x_offset_ll1
results_vis[0]['bodies']['candidate'][9,1] += y_offset_ll1
results_vis[0]['bodies']['candidate'][10,0] += x_offset_ll1
results_vis[0]['bodies']['candidate'][10,1] += y_offset_ll1
results_vis[0]['bodies']['candidate'][19,0] += x_offset_ll1
results_vis[0]['bodies']['candidate'][19,1] += y_offset_ll1
l_ll2_ref = ((ref_candidate[9][0] - ref_candidate[10][0]) ** 2 + (ref_candidate[9][1] - ref_candidate[10][1]) ** 2) ** 0.5
l_ll2_0 = ((candidate[9][0] - candidate[10][0]) ** 2 + (candidate[9][1] - candidate[10][1]) ** 2) ** 0.5
ll2_ratio = l_ll2_ref / l_ll2_0
x_offset_ll2 = (candidate[10][0]-candidate[9][0])*(ll2_ratio-1.)
y_offset_ll2 = (candidate[10][1]-candidate[9][1])*(ll2_ratio-1.)
results_vis[0]['bodies']['candidate'][10,0] += x_offset_ll2
results_vis[0]['bodies']['candidate'][10,1] += y_offset_ll2
results_vis[0]['bodies']['candidate'][19,0] += x_offset_ll2
results_vis[0]['bodies']['candidate'][19,1] += y_offset_ll2
########right leg########
l_rl1_ref = ((ref_candidate[11][0] - ref_candidate[12][0]) ** 2 + (ref_candidate[11][1] - ref_candidate[12][1]) ** 2) ** 0.5
l_rl1_0 = ((candidate[11][0] - candidate[12][0]) ** 2 + (candidate[11][1] - candidate[12][1]) ** 2) ** 0.5
rl1_ratio = l_rl1_ref / l_rl1_0
x_offset_rl1 = (candidate[12][0]-candidate[11][0])*(rl1_ratio-1.)
y_offset_rl1 = (candidate[12][1]-candidate[11][1])*(rl1_ratio-1.)
results_vis[0]['bodies']['candidate'][12,0] += x_offset_rl1
results_vis[0]['bodies']['candidate'][12,1] += y_offset_rl1
results_vis[0]['bodies']['candidate'][13,0] += x_offset_rl1
results_vis[0]['bodies']['candidate'][13,1] += y_offset_rl1
results_vis[0]['bodies']['candidate'][18,0] += x_offset_rl1
results_vis[0]['bodies']['candidate'][18,1] += y_offset_rl1
l_rl2_ref = ((ref_candidate[12][0] - ref_candidate[13][0]) ** 2 + (ref_candidate[12][1] - ref_candidate[13][1]) ** 2) ** 0.5
l_rl2_0 = ((candidate[12][0] - candidate[13][0]) ** 2 + (candidate[12][1] - candidate[13][1]) ** 2) ** 0.5
rl2_ratio = l_rl2_ref / l_rl2_0
x_offset_rl2 = (candidate[13][0]-candidate[12][0])*(rl2_ratio-1.)
y_offset_rl2 = (candidate[13][1]-candidate[12][1])*(rl2_ratio-1.)
results_vis[0]['bodies']['candidate'][13,0] += x_offset_rl2
results_vis[0]['bodies']['candidate'][13,1] += y_offset_rl2
results_vis[0]['bodies']['candidate'][18,0] += x_offset_rl2
results_vis[0]['bodies']['candidate'][18,1] += y_offset_rl2
offset = ref_candidate[1] - results_vis[0]['bodies']['candidate'][1]
results_vis[0]['bodies']['candidate'] += offset[np.newaxis, :]
results_vis[0]['faces'] += offset[np.newaxis, np.newaxis, :]
results_vis[0]['hands'] += offset[np.newaxis, np.newaxis, :]
for i in range(1, len(results_vis)):
results_vis[i]['bodies']['candidate'][:,0] *= x_ratio
results_vis[i]['bodies']['candidate'][:,1] *= y_ratio
results_vis[i]['faces'][:,:,0] *= x_ratio
results_vis[i]['faces'][:,:,1] *= y_ratio
results_vis[i]['hands'][:,:,0] *= x_ratio
results_vis[i]['hands'][:,:,1] *= y_ratio
########neck########
x_offset_neck = (results_vis[i]['bodies']['candidate'][1][0]-results_vis[i]['bodies']['candidate'][0][0])*(1.-neck_ratio)
y_offset_neck = (results_vis[i]['bodies']['candidate'][1][1]-results_vis[i]['bodies']['candidate'][0][1])*(1.-neck_ratio)
results_vis[i]['bodies']['candidate'][0,0] += x_offset_neck
results_vis[i]['bodies']['candidate'][0,1] += y_offset_neck
results_vis[i]['bodies']['candidate'][14,0] += x_offset_neck
results_vis[i]['bodies']['candidate'][14,1] += y_offset_neck
results_vis[i]['bodies']['candidate'][15,0] += x_offset_neck
results_vis[i]['bodies']['candidate'][15,1] += y_offset_neck
results_vis[i]['bodies']['candidate'][16,0] += x_offset_neck
results_vis[i]['bodies']['candidate'][16,1] += y_offset_neck
results_vis[i]['bodies']['candidate'][17,0] += x_offset_neck
results_vis[i]['bodies']['candidate'][17,1] += y_offset_neck
########shoulder2########
x_offset_shoulder2 = (results_vis[i]['bodies']['candidate'][1][0]-results_vis[i]['bodies']['candidate'][2][0])*(1.-shoulder2_ratio)
y_offset_shoulder2 = (results_vis[i]['bodies']['candidate'][1][1]-results_vis[i]['bodies']['candidate'][2][1])*(1.-shoulder2_ratio)
results_vis[i]['bodies']['candidate'][2,0] += x_offset_shoulder2
results_vis[i]['bodies']['candidate'][2,1] += y_offset_shoulder2
results_vis[i]['bodies']['candidate'][3,0] += x_offset_shoulder2
results_vis[i]['bodies']['candidate'][3,1] += y_offset_shoulder2
results_vis[i]['bodies']['candidate'][4,0] += x_offset_shoulder2
results_vis[i]['bodies']['candidate'][4,1] += y_offset_shoulder2
results_vis[i]['hands'][1,:,0] += x_offset_shoulder2
results_vis[i]['hands'][1,:,1] += y_offset_shoulder2
########shoulder5########
x_offset_shoulder5 = (results_vis[i]['bodies']['candidate'][1][0]-results_vis[i]['bodies']['candidate'][5][0])*(1.-shoulder5_ratio)
y_offset_shoulder5 = (results_vis[i]['bodies']['candidate'][1][1]-results_vis[i]['bodies']['candidate'][5][1])*(1.-shoulder5_ratio)
results_vis[i]['bodies']['candidate'][5,0] += x_offset_shoulder5
results_vis[i]['bodies']['candidate'][5,1] += y_offset_shoulder5
results_vis[i]['bodies']['candidate'][6,0] += x_offset_shoulder5
results_vis[i]['bodies']['candidate'][6,1] += y_offset_shoulder5
results_vis[i]['bodies']['candidate'][7,0] += x_offset_shoulder5
results_vis[i]['bodies']['candidate'][7,1] += y_offset_shoulder5
results_vis[i]['hands'][0,:,0] += x_offset_shoulder5
results_vis[i]['hands'][0,:,1] += y_offset_shoulder5
########arm3########
x_offset_arm3 = (results_vis[i]['bodies']['candidate'][2][0]-results_vis[i]['bodies']['candidate'][3][0])*(1.-arm3_ratio)
y_offset_arm3 = (results_vis[i]['bodies']['candidate'][2][1]-results_vis[i]['bodies']['candidate'][3][1])*(1.-arm3_ratio)
results_vis[i]['bodies']['candidate'][3,0] += x_offset_arm3
results_vis[i]['bodies']['candidate'][3,1] += y_offset_arm3
results_vis[i]['bodies']['candidate'][4,0] += x_offset_arm3
results_vis[i]['bodies']['candidate'][4,1] += y_offset_arm3
results_vis[i]['hands'][1,:,0] += x_offset_arm3
results_vis[i]['hands'][1,:,1] += y_offset_arm3
########arm4########
x_offset_arm4 = (results_vis[i]['bodies']['candidate'][3][0]-results_vis[i]['bodies']['candidate'][4][0])*(1.-arm4_ratio)
y_offset_arm4 = (results_vis[i]['bodies']['candidate'][3][1]-results_vis[i]['bodies']['candidate'][4][1])*(1.-arm4_ratio)
results_vis[i]['bodies']['candidate'][4,0] += x_offset_arm4
results_vis[i]['bodies']['candidate'][4,1] += y_offset_arm4
results_vis[i]['hands'][1,:,0] += x_offset_arm4
results_vis[i]['hands'][1,:,1] += y_offset_arm4
########arm6########
x_offset_arm6 = (results_vis[i]['bodies']['candidate'][5][0]-results_vis[i]['bodies']['candidate'][6][0])*(1.-arm6_ratio)
y_offset_arm6 = (results_vis[i]['bodies']['candidate'][5][1]-results_vis[i]['bodies']['candidate'][6][1])*(1.-arm6_ratio)
results_vis[i]['bodies']['candidate'][6,0] += x_offset_arm6
results_vis[i]['bodies']['candidate'][6,1] += y_offset_arm6
results_vis[i]['bodies']['candidate'][7,0] += x_offset_arm6
results_vis[i]['bodies']['candidate'][7,1] += y_offset_arm6
results_vis[i]['hands'][0,:,0] += x_offset_arm6
results_vis[i]['hands'][0,:,1] += y_offset_arm6
########arm7########
x_offset_arm7 = (results_vis[i]['bodies']['candidate'][6][0]-results_vis[i]['bodies']['candidate'][7][0])*(1.-arm7_ratio)
y_offset_arm7 = (results_vis[i]['bodies']['candidate'][6][1]-results_vis[i]['bodies']['candidate'][7][1])*(1.-arm7_ratio)
results_vis[i]['bodies']['candidate'][7,0] += x_offset_arm7
results_vis[i]['bodies']['candidate'][7,1] += y_offset_arm7
results_vis[i]['hands'][0,:,0] += x_offset_arm7
results_vis[i]['hands'][0,:,1] += y_offset_arm7
########head14########
x_offset_head14 = (results_vis[i]['bodies']['candidate'][0][0]-results_vis[i]['bodies']['candidate'][14][0])*(1.-head14_ratio)
y_offset_head14 = (results_vis[i]['bodies']['candidate'][0][1]-results_vis[i]['bodies']['candidate'][14][1])*(1.-head14_ratio)
results_vis[i]['bodies']['candidate'][14,0] += x_offset_head14
results_vis[i]['bodies']['candidate'][14,1] += y_offset_head14
results_vis[i]['bodies']['candidate'][16,0] += x_offset_head14
results_vis[i]['bodies']['candidate'][16,1] += y_offset_head14
########head15########
x_offset_head15 = (results_vis[i]['bodies']['candidate'][0][0]-results_vis[i]['bodies']['candidate'][15][0])*(1.-head15_ratio)
y_offset_head15 = (results_vis[i]['bodies']['candidate'][0][1]-results_vis[i]['bodies']['candidate'][15][1])*(1.-head15_ratio)
results_vis[i]['bodies']['candidate'][15,0] += x_offset_head15
results_vis[i]['bodies']['candidate'][15,1] += y_offset_head15
results_vis[i]['bodies']['candidate'][17,0] += x_offset_head15
results_vis[i]['bodies']['candidate'][17,1] += y_offset_head15
########head16########
x_offset_head16 = (results_vis[i]['bodies']['candidate'][14][0]-results_vis[i]['bodies']['candidate'][16][0])*(1.-head16_ratio)
y_offset_head16 = (results_vis[i]['bodies']['candidate'][14][1]-results_vis[i]['bodies']['candidate'][16][1])*(1.-head16_ratio)
results_vis[i]['bodies']['candidate'][16,0] += x_offset_head16
results_vis[i]['bodies']['candidate'][16,1] += y_offset_head16
########head17########
x_offset_head17 = (results_vis[i]['bodies']['candidate'][15][0]-results_vis[i]['bodies']['candidate'][17][0])*(1.-head17_ratio)
y_offset_head17 = (results_vis[i]['bodies']['candidate'][15][1]-results_vis[i]['bodies']['candidate'][17][1])*(1.-head17_ratio)
results_vis[i]['bodies']['candidate'][17,0] += x_offset_head17
results_vis[i]['bodies']['candidate'][17,1] += y_offset_head17
# ########MovingAverage########
########left leg########
x_offset_ll1 = (results_vis[i]['bodies']['candidate'][9][0]-results_vis[i]['bodies']['candidate'][8][0])*(ll1_ratio-1.)
y_offset_ll1 = (results_vis[i]['bodies']['candidate'][9][1]-results_vis[i]['bodies']['candidate'][8][1])*(ll1_ratio-1.)
results_vis[i]['bodies']['candidate'][9,0] += x_offset_ll1
results_vis[i]['bodies']['candidate'][9,1] += y_offset_ll1
results_vis[i]['bodies']['candidate'][10,0] += x_offset_ll1
results_vis[i]['bodies']['candidate'][10,1] += y_offset_ll1
results_vis[i]['bodies']['candidate'][19,0] += x_offset_ll1
results_vis[i]['bodies']['candidate'][19,1] += y_offset_ll1
x_offset_ll2 = (results_vis[i]['bodies']['candidate'][10][0]-results_vis[i]['bodies']['candidate'][9][0])*(ll2_ratio-1.)
y_offset_ll2 = (results_vis[i]['bodies']['candidate'][10][1]-results_vis[i]['bodies']['candidate'][9][1])*(ll2_ratio-1.)
results_vis[i]['bodies']['candidate'][10,0] += x_offset_ll2
results_vis[i]['bodies']['candidate'][10,1] += y_offset_ll2
results_vis[i]['bodies']['candidate'][19,0] += x_offset_ll2
results_vis[i]['bodies']['candidate'][19,1] += y_offset_ll2
########right leg########
x_offset_rl1 = (results_vis[i]['bodies']['candidate'][12][0]-results_vis[i]['bodies']['candidate'][11][0])*(rl1_ratio-1.)
y_offset_rl1 = (results_vis[i]['bodies']['candidate'][12][1]-results_vis[i]['bodies']['candidate'][11][1])*(rl1_ratio-1.)
results_vis[i]['bodies']['candidate'][12,0] += x_offset_rl1
results_vis[i]['bodies']['candidate'][12,1] += y_offset_rl1
results_vis[i]['bodies']['candidate'][13,0] += x_offset_rl1
results_vis[i]['bodies']['candidate'][13,1] += y_offset_rl1
results_vis[i]['bodies']['candidate'][18,0] += x_offset_rl1
results_vis[i]['bodies']['candidate'][18,1] += y_offset_rl1
x_offset_rl2 = (results_vis[i]['bodies']['candidate'][13][0]-results_vis[i]['bodies']['candidate'][12][0])*(rl2_ratio-1.)
y_offset_rl2 = (results_vis[i]['bodies']['candidate'][13][1]-results_vis[i]['bodies']['candidate'][12][1])*(rl2_ratio-1.)
results_vis[i]['bodies']['candidate'][13,0] += x_offset_rl2
results_vis[i]['bodies']['candidate'][13,1] += y_offset_rl2
results_vis[i]['bodies']['candidate'][18,0] += x_offset_rl2
results_vis[i]['bodies']['candidate'][18,1] += y_offset_rl2
results_vis[i]['bodies']['candidate'] += offset[np.newaxis, :]
results_vis[i]['faces'] += offset[np.newaxis, np.newaxis, :]
results_vis[i]['hands'] += offset[np.newaxis, np.newaxis, :]
dwpose_woface_list = []
for i in range(len(results_vis)):
dwpose_woface, dwpose_wface = draw_pose(results_vis[i], H=height, W=width)
dwpose_woface_list.append(torch.from_numpy(dwpose_woface))
dwpose_woface_tensor = torch.stack(dwpose_woface_list, dim=0)
dwpose_woface_ref, dwpose_wface_ref = draw_pose(pose_ref, H=height, W=width)
dwpose_woface_ref_tensor = torch.from_numpy(dwpose_woface_ref)
return dwpose_woface_tensor, dwpose_woface_ref_tensor
class WanVideoUniAnimateDWPoseDetector:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"pose_images": ("IMAGE", {"tooltip": "Pose images"}),
"reference_pose_image": ("IMAGE", {"tooltip": "Reference pose image"}),
},
}
RETURN_TYPES = ("IMAGE", "IMAGE", )
RETURN_NAMES = ("poses", "reference_pose",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, pose_images, reference_pose_image):
device = mm.get_torch_device()
#model loading
dw_pose_model = "dw-ll_ucoco_384_bs5.torchscript.pt"
yolo_model = "yolox_l.torchscript.pt"
script_directory = os.path.dirname(os.path.abspath(__file__))
model_base_path = os.path.join(script_directory, "models", "DWPose")
model_det=os.path.join(model_base_path, yolo_model)
model_pose=os.path.join(model_base_path, dw_pose_model)
if not os.path.exists(model_det):
log.info(f"Downloading yolo model to: {model_base_path}")
from huggingface_hub import snapshot_download
snapshot_download(repo_id="hr16/yolox-onnx",
allow_patterns=[f"*{yolo_model}*"],
local_dir=model_base_path,
local_dir_use_symlinks=False)
if not os.path.exists(model_pose):
log.info(f"Downloading dwpose model to: {model_base_path}")
from huggingface_hub import snapshot_download
snapshot_download(repo_id="hr16/DWPose-TorchScript-BatchSize5",
allow_patterns=[f"*{dw_pose_model}*"],
local_dir=model_base_path,
local_dir_use_symlinks=False)
if not hasattr(self, "det") or not hasattr(self, "pose"):
self.det = torch.jit.load(model_det, map_location=device)
self.pose = torch.jit.load(model_pose, map_location=device)
self.dwpose_detector = DWposeDetector(self.det, self.pose)
#model inference
height, width = pose_images.shape[1:3]
pose_np = pose_images.cpu().numpy() * 255
ref = reference_pose_image
ref_np = ref.cpu().numpy() * 255
poses, reference_pose = pose_extract(pose_np, ref_np, self.dwpose_detector, height, width)
poses = poses / 255.0
reference_pose = reference_pose.unsqueeze(0) / 255.0
return (poses, reference_pose, )
class WanVideoUniAnimatePoseInput:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"pose_images": ("IMAGE", {"tooltip": "Pose images"}),
"reference_pose_image": ("IMAGE", {"tooltip": "Reference pose image"}),
},
}
RETURN_TYPES = ("UNIANIMATE_POSE", )
RETURN_NAMES = ("unianimate_poses",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, pose_images, reference_pose_image):
pose = pose_images.permute(3, 0, 1, 2).unsqueeze(0).contiguous()
ref = reference_pose_image.permute(0, 3, 1, 2).contiguous()
unianim_poses = {
"pose": pose,
"ref": ref,
}
return (unianim_poses,)
NODE_CLASS_MAPPINGS = {
"WanVideoUniAnimatePoseInput": WanVideoUniAnimatePoseInput,
"WanVideoUniAnimateDWPoseDetector": WanVideoUniAnimateDWPoseDetector,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoUniAnimatePoseInput": "WanVideo UniAnimate Pose Input",
"WanVideoUniAnimateDWPoseDetector": "WanVideo UniAnimate DWPose Detector",
}
+54 -1
View File
@@ -164,4 +164,57 @@ def encode_image_(clip_vision, image):
pixel_values = clip_preprocess(image, size=224, crop=True).float()
out = clip_vision.visual(pixel_values)
return out
return out
# Code based on https://github.com/WikiChao/FreSca (MIT License)
import torch
import torch.fft as fft
def fourier_filter(x, scale_low=1.0, scale_high=1.5, freq_cutoff=20):
"""
Apply frequency-dependent scaling to an image tensor using Fourier transforms.
Parameters:
x: Input tensor of shape (B, C, H, W)
scale_low: Scaling factor for low-frequency components (default: 1.0)
scale_high: Scaling factor for high-frequency components (default: 1.5)
freq_cutoff: Number of frequency indices around center to consider as low-frequency (default: 20)
Returns:
x_filtered: Filtered version of x in spatial domain with frequency-specific scaling applied.
"""
# Preserve input dtype and device
dtype, device = x.dtype, x.device
# Convert to float32 for FFT computations
x = x.to(torch.float32)
# 1) Apply FFT and shift low frequencies to center
x_freq = fft.fftn(x, dim=(-2, -1))
x_freq = fft.fftshift(x_freq, dim=(-2, -1))
# 2) Create a mask to scale frequencies differently
C, B, H, W = x_freq.shape
crow, ccol = H // 2, W // 2
# Initialize mask with high-frequency scaling factor
mask = torch.ones((C, B, H, W), device=device) * scale_high
# Apply low-frequency scaling factor to center region
mask[
...,
crow - freq_cutoff : crow + freq_cutoff,
ccol - freq_cutoff : ccol + freq_cutoff,
] = scale_low
# 3) Apply frequency-specific scaling
x_freq = x_freq * mask
# 4) Convert back to spatial domain
x_freq = fft.ifftshift(x_freq, dim=(-2, -1))
x_filtered = fft.ifftn(x_freq, dim=(-2, -1)).real
# 5) Restore original dtype
x_filtered = x_filtered.to(dtype)
return x_filtered
+5 -1
View File
@@ -966,7 +966,8 @@ class WanModel(ModelMixin, ConfigMixin):
pred_id=None,
control_lora_enabled=False,
vace_data = None,
camera_embed = None
camera_embed = None,
unianim_data = None
):
r"""
Forward pass through the diffusion model
@@ -1112,6 +1113,9 @@ class WanModel(ModelMixin, ConfigMixin):
if self.enable_teacache:
original_x = x.clone().to(self.teacache_cache_device, non_blocking=self.use_non_blocking)
if hasattr(self, "dwpose_embedding") and unianim_data is not None:
x += rearrange(unianim_data['dwpose'], 'b c f h w -> b (f h w) c').contiguous()
# arguments
kwargs = dict(
e=e0,