Strength controls for Unianimate

This commit is contained in:
kijai
2025-04-18 23:06:49 +03:00
parent 19044adc78
commit 421e375f13
3 changed files with 20 additions and 4 deletions
+6 -2
View File
@@ -20,6 +20,7 @@ from .taehv import TAEHV
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
from einops import rearrange
import folder_paths
import comfy.model_management as mm
@@ -2380,6 +2381,7 @@ class WanVideoSampler:
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"])
dwpose_data = rearrange(dwpose_data, 'b c f h w -> b (f h w) c').contiguous()
random_ref_dwpose_data = None
if image_cond is not None:
@@ -2387,11 +2389,13 @@ class WanVideoSampler:
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
"random_ref": random_ref_dwpose_data.squeeze(0) if random_ref_dwpose_data is not None else None,
"strength": unianimate_poses["strength"],
"start_percent": unianimate_poses["start_percent"],
"end_percent": unianimate_poses["end_percent"]
}
latent_video_length = noise.shape[1]
+7 -1
View File
@@ -764,6 +764,9 @@ class WanVideoUniAnimatePoseInput:
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"}),
},
}
@@ -772,7 +775,7 @@ class WanVideoUniAnimatePoseInput:
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, pose_images, reference_pose_image):
def process(self, pose_images, reference_pose_image, strength, start_percent, end_percent):
pose = pose_images.permute(3, 0, 1, 2).unsqueeze(0).contiguous()
ref = reference_pose_image.permute(0, 3, 1, 2).contiguous()
@@ -780,6 +783,9 @@ class WanVideoUniAnimatePoseInput:
unianim_poses = {
"pose": pose,
"ref": ref,
"strength": strength,
"start_percent": start_percent,
"end_percent": end_percent
}
return (unianim_poses,)
+7 -1
View File
@@ -998,6 +998,10 @@ class WanModel(ModelMixin, ConfigMixin):
_, F, H, W = x[0].shape
if y is not None:
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"]
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)]
# embeddings
@@ -1114,7 +1118,9 @@ class WanModel(ModelMixin, ConfigMixin):
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()
if unianim_data['start_percent'] <= current_step_percentage <= unianim_data['end_percent']:
dwpose_emb = unianim_data['dwpose']
x += dwpose_emb * unianim_data['strength']
# arguments
kwargs = dict(