diff --git a/nodes.py b/nodes.py index a6081cb..d477846 100644 --- a/nodes.py +++ b/nodes.py @@ -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, diff --git a/unianimate/nodes.py b/unianimate/nodes.py index a656956..be1e2ff 100644 --- a/unianimate/nodes.py +++ b/unianimate/nodes.py @@ -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, diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 40f68ed..ea032df 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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