Make UniAnimate ref pose optional

This commit is contained in:
kijai
2025-04-19 17:41:04 +03:00
parent 90c9415d31
commit 2cfae9fa8f
3 changed files with 15 additions and 8 deletions
+5 -4
View File
@@ -2400,10 +2400,11 @@ class WanVideoSampler:
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]
random_ref_dwpose = unianimate_poses.get("ref", None)
if random_ref_dwpose is not None:
random_ref_dwpose_data = transformer.randomref_embedding_pose(
random_ref_dwpose.to(device)
).unsqueeze(2).to(model["dtype"]) # [1, 20, 104, 60]
unianim_data = {
"dwpose": dwpose_data,
+8 -3
View File
@@ -763,11 +763,13 @@ class WanVideoUniAnimatePoseInput:
def INPUT_TYPES(s):
return {"required": {
"pose_images": ("IMAGE", {"tooltip": "Pose images"}),
"reference_pose_image": ("IMAGE", {"tooltip": "Reference pose image"}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Strength of the pose control"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage for the pose control"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage for the pose control"}),
},
"optional": {
"reference_pose_image": ("IMAGE", {"tooltip": "Reference pose image"}),
},
}
RETURN_TYPES = ("UNIANIMATE_POSE", )
@@ -775,10 +777,13 @@ class WanVideoUniAnimatePoseInput:
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, pose_images, reference_pose_image, strength, start_percent, end_percent):
def process(self, pose_images, strength, start_percent, end_percent, reference_pose_image=None):
pose = pose_images.permute(3, 0, 1, 2).unsqueeze(0).contiguous()
ref = reference_pose_image.permute(0, 3, 1, 2).contiguous()
ref = None
if reference_pose_image is not None:
ref = reference_pose_image.permute(0, 3, 1, 2).contiguous()
unianim_poses = {
"pose": pose,
+2 -1
View File
@@ -1001,7 +1001,8 @@ class WanModel(ModelMixin, ConfigMixin):
if hasattr(self, "randomref_embedding_pose") and unianim_data is not None:
if unianim_data['start_percent'] <= current_step_percentage <= unianim_data['end_percent']:
random_ref_emb = unianim_data["random_ref"]
y[0] = y[0] + random_ref_emb * unianim_data["strength"]
if random_ref_emb is not None:
y[0] = y[0] + random_ref_emb * unianim_data["strength"]
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)]
# embeddings