Make UniAnimate ref pose optional
This commit is contained in:
@@ -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
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user