Allow unianimate with DF

This commit is contained in:
kijai
2025-04-22 01:01:22 +03:00
parent 3ae0bc1fec
commit f93e54d21b
2 changed files with 40 additions and 3 deletions
+39 -2
View File
@@ -10,6 +10,7 @@ from ..wanvideo.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from ..wanvideo.utils.scheduling_flow_match_lcm import FlowMatchLCMScheduler
from ..nodes import optimized_scale
from einops import rearrange
import comfy.model_management as mm
from comfy.utils import load_torch_file, ProgressBar, common_upscale
@@ -122,6 +123,7 @@ class WanVideoDiffusionForcingSampler:
"slg_args": ("SLGARGS", ),
"rope_function": (["default", "comfy"], {"default": "comfy", "tooltip": "Comfy's RoPE implementation doesn't use complex numbers and can thus be compiled, that should be a lot faster when using torch.compile"}),
"experimental_args": ("EXPERIMENTALARGS", ),
"unianimate_poses": ("UNIANIMATE_POSE", ),
}
}
@@ -131,7 +133,8 @@ class WanVideoDiffusionForcingSampler:
CATEGORY = "WanVideoWrapper"
def process(self, model, text_embeds, image_embeds, shift, fps, steps, addnoise_condition, cfg, seed, scheduler,
force_offload=True, samples=None, prefix_samples=None, denoise_strength=1.0, slg_args=None, rope_function="default", teacache_args=None, flowedit_args=None, batched_cfg=False, experimental_args=None):
force_offload=True, samples=None, prefix_samples=None, denoise_strength=1.0, slg_args=None, rope_function="default", teacache_args=None,
experimental_args=None, unianimate_poses=None):
#assert not (context_options and teacache_args), "Context options cannot currently be used together with teacache."
patcher = model
model = model.model
@@ -267,6 +270,39 @@ class WanVideoDiffusionForcingSampler:
sample_schedulers.append(sample_scheduler)
sample_schedulers_counter = [0] * latent_video_length
if unianimate_poses is not None:
transformer.dwpose_embedding.to(device)
transformer.randomref_embedding_pose.to(device)
dwpose_data = unianimate_poses["pose"]
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"])
log.info(f"UniAnimate pose embed shape: {dwpose_data.shape}")
if dwpose_data.shape[2] > latent_video_length:
log.warning(f"UniAnimate pose embed length {dwpose_data.shape[2]} is longer than the video length {latent_video_length}, truncating")
dwpose_data = dwpose_data[:,:, :latent_video_length]
elif dwpose_data.shape[2] < latent_video_length:
log.warning(f"UniAnimate pose embed length {dwpose_data.shape[2]} is shorter than the video length {latent_video_length}, padding with last pose")
pad_len = latent_video_length - dwpose_data.shape[2]
pad = dwpose_data[:,:,:1].repeat(1,1,pad_len,1,1)
dwpose_data = torch.cat([dwpose_data, pad], dim=2)
dwpose_data_flat = 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:
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_flat,
"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"]
}
freqs = None
@@ -392,6 +428,7 @@ class WanVideoDiffusionForcingSampler:
'current_step': idx,
'control_lora_enabled': control_lora_enabled,
'vace_data': vace_data,
'unianim_data': unianim_data,
'fps_embeds': fps_embeds,
}
@@ -491,7 +528,7 @@ class WanVideoDiffusionForcingSampler:
cfg[i],
text_embeds["prompt_embeds"],
text_embeds["negative_prompt_embeds"],
timestep, i, image_cond, clip_fea, vace_data=vace_data,
timestep, i, image_cond, clip_fea, unianim_data=unianim_data, vace_data=vace_data,
teacache_state=self.teacache_state)
for idx in range(valid_interval_start, valid_interval_end):
+1 -1
View File
@@ -1131,7 +1131,7 @@ class WanModel(ModelMixin, ConfigMixin):
if self.enable_teacache and self.teacache_start_step <= current_step <= self.teacache_end_step:
if pred_id is None:
pred_id = self.teacache_state.new_prediction(cache_device=self.teacache_cache_device)
log.info(current_step)
#log.info(current_step)
#log.info(f"TeaCache: Initializing TeaCache variables for model pred: {pred_id}")
should_calc = True
else: