diff --git a/nodes.py b/nodes.py index f9bc82e..4682a66 100644 --- a/nodes.py +++ b/nodes.py @@ -2929,7 +2929,7 @@ class WanVideoSampler: #region model pred def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, - control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None, teacache_state=None): + control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None, add_cond=None, teacache_state=None): z = z.to(dtype) with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=("fp8" in model["quantization"])): @@ -2984,6 +2984,7 @@ class WanVideoSampler: if recammaster is not None: z = torch.cat([z, recam_latents.to(z)], dim=1) + use_phantom = False if phantom_latents is not None: if (phantom_start_percent <= current_step_percentage <= phantom_end_percent) or \ @@ -3400,12 +3401,16 @@ class WanVideoSampler: "start_percent": unianimate_poses["start_percent"], "end_percent": unianimate_poses["end_percent"] } + + if add_cond is not None: + partial_add_cond = add_cond[:, :, c].to(device, dtype) noise_pred_context, new_teacache = predict_with_cfg( partial_latent_model_input, cfg[idx], positive, text_embeds["negative_prompt_embeds"], - timestep, idx, partial_img_emb, clip_fea, partial_control_latents, partial_vace_context, partial_unianim_data,partial_audio_proj,partial_control_camera_latents, + timestep, idx, partial_img_emb, clip_fea, partial_control_latents, partial_vace_context, partial_unianim_data,partial_audio_proj, + partial_control_camera_latents, partial_add_cond, current_teacache) if teacache_args is not None: @@ -3422,7 +3427,7 @@ class WanVideoSampler: cfg[idx], text_embeds["prompt_embeds"], text_embeds["negative_prompt_embeds"], - timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, + timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond, teacache_state=self.teacache_state) if latent_shift_loop: