Strength controls for Unianimate
This commit is contained in:
@@ -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
@@ -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,)
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user