Fix differential diffusion, VACE tweaks

This commit is contained in:
kijai
2025-08-21 12:26:29 +03:00
parent bcdd6dc664
commit bd5f76a3c8
2 changed files with 56 additions and 36 deletions
+54 -34
View File
@@ -1232,8 +1232,6 @@ class WanVideoVACEEncode:
CATEGORY = "WanVideoWrapper"
def process(self, vae, width, height, num_frames, strength, vace_start_percent, vace_end_percent, input_frames=None, ref_images=None, input_masks=None, prev_vace_embeds=None, tiled_vae=False):
vae = vae.to(device)
width = (width // 16) * 16
height = (height // 16) * 16
@@ -1244,8 +1242,8 @@ class WanVideoVACEEncode:
if input_frames is None:
input_frames = torch.zeros((1, 3, num_frames, height, width), device=device, dtype=vae.dtype)
else:
input_frames = input_frames[:num_frames]
input_frames = common_upscale(input_frames.clone().movedim(-1, 1), width, height, "lanczos", "disabled").movedim(1, -1)
input_frames = input_frames.clone()[:num_frames, :, :, :3]
input_frames = common_upscale(input_frames.movedim(-1, 1), width, height, "lanczos", "disabled").movedim(1, -1)
input_frames = input_frames.to(vae.dtype).to(device).unsqueeze(0).permute(0, 4, 1, 2, 3) # B, C, T, H, W
input_frames = input_frames * 2 - 1
if input_masks is None:
@@ -1258,10 +1256,11 @@ class WanVideoVACEEncode:
input_masks = input_masks.unsqueeze(-1).unsqueeze(0).permute(0, 4, 1, 2, 3).repeat(1, 3, 1, 1, 1) # B, C, T, H, W
if ref_images is not None:
ref_images = ref_images.clone()[..., :3]
# Create padded image
if ref_images.shape[0] > 1:
ref_images = torch.cat([ref_images[i] for i in range(ref_images.shape[0])], dim=1).unsqueeze(0)
B, H, W, C = ref_images.shape
current_aspect = W / H
target_aspect = width / height
@@ -1283,12 +1282,12 @@ class WanVideoVACEEncode:
ref_images = ref_images.to(vae.dtype).to(device).unsqueeze(0).permute(0, 4, 1, 2, 3).unsqueeze(0)
ref_images = ref_images * 2 - 1
vae = vae.to(device)
z0 = self.vace_encode_frames(vae, input_frames, ref_images, masks=input_masks, tiled_vae=tiled_vae)
vae.model.clear_cache()
m0 = self.vace_encode_masks(input_masks, ref_images)
z = self.vace_latent(z0, m0)
vae.to(offload_device)
vace_input = {
@@ -1306,9 +1305,14 @@ class WanVideoVACEEncode:
if prev_vace_embeds is not None:
if "additional_vace_inputs" in prev_vace_embeds and prev_vace_embeds["additional_vace_inputs"]:
vace_input["additional_vace_inputs"] = prev_vace_embeds["additional_vace_inputs"].copy()
else:
new_entry = prev_vace_embeds.copy()
new_entry.update(vace_input)
return (new_entry,)
vace_input["additional_vace_inputs"].append(prev_vace_embeds)
return (vace_input,)
def vace_encode_frames(self, vae, frames, ref_images, masks=None, tiled_vae=False):
if ref_images is None:
ref_images = [None] * len(frames)
@@ -2710,18 +2714,6 @@ class WanVideoSampler:
if idx < skip_steps:
continue
# diff diff
if masks is not None:
if idx < len(timesteps) - 1:
noise_timestep = timesteps[idx+1]
image_latent = sample_scheduler.scale_noise(
original_image, torch.tensor([noise_timestep]), noise.to(device)
)
mask = masks[idx]
mask = mask.to(latent)
latent = image_latent * mask + latent * (1-mask)
# end diff diff
if bidirectional_sampling:
latent_flipped = torch.flip(latent, dims=[1])
latent_model_input_flipped = latent_flipped.to(device)
@@ -3015,7 +3007,6 @@ class WanVideoSampler:
if mode == "auto":
mode = transformer.multitalk_model_type.lower()
log.info(f"Multitalk mode: {mode}")
original_images = cond_image = image_embeds.get("multitalk_start_image", None)
cond_frame = None
offload = image_embeds.get("force_offload", False)
tiled_vae = image_embeds.get("tiled_vae", False)
@@ -3028,6 +3019,9 @@ class WanVideoSampler:
motion_frame = image_embeds.get("motion_frame", 25)
target_w = image_embeds.get("target_w", None)
target_h = image_embeds.get("target_h", None)
original_images = cond_image = image_embeds.get("multitalk_start_image", None)
if original_images is None:
original_images = torch.zeros([noise.shape[0], 1, target_h, target_w], device=device)
if len(multitalk_embeds['audio_features'])==2 and (multitalk_embeds['ref_target_masks'] is None):
face_scale = 0.1
@@ -3115,18 +3109,18 @@ class WanVideoSampler:
16, (frame_num - 1) // 4 + 1,
lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
# Calculate the correct latent slice based on current iteration
if is_first_clip:
latent_start_idx = 0
latent_end_idx = noise.shape[1]
else:
new_frames_per_iteration = frame_num - motion_frame
new_latent_frames_per_iteration = ((new_frames_per_iteration - 1) // 4 + 1)
latent_start_idx = iteration_count * new_latent_frames_per_iteration
latent_end_idx = latent_start_idx + noise.shape[1]
if samples is not None:
input_samples = samples["samples"].squeeze(0).to(noise)
# Calculate the correct slice based on current iteration
if is_first_clip:
latent_start_idx = 0
latent_end_idx = noise.shape[1]
else:
new_frames_per_iteration = frame_num - motion_frame
new_latent_frames_per_iteration = ((new_frames_per_iteration - 1) // 4 + 1)
latent_start_idx = iteration_count * new_latent_frames_per_iteration
latent_end_idx = latent_start_idx + noise.shape[1]
input_samples = samples["samples"].squeeze(0).to(noise)
# Check if we have enough frames in input_samples
if latent_end_idx > input_samples.shape[1]:
# We need more frames than available - pad the input_samples at the end
@@ -3143,6 +3137,22 @@ class WanVideoSampler:
else:
noise = input_samples
window_vace_data = None
if vace_data is not None:
window_vace_data = []
for vace_entry in vace_data:
partial_context = vace_entry["context"][0][:, latent_start_idx:latent_end_idx]
if has_ref:
partial_context[:, 0] = vace_entry["context"][0][:, 0]
window_vace_data.append({
"context": [partial_context],
"scale": vace_entry["scale"],
"start": vace_entry["start"],
"end": vace_entry["end"],
"seq_len": vace_entry["seq_len"]
})
# get mask
msk = torch.ones(1, frame_num, lat_h, lat_w, device=device)
if mode == "multitalk":
@@ -3158,8 +3168,8 @@ class WanVideoSampler:
mm.soft_empty_cache()
# zero padding and vae encode
if cond_image is not None or cond_frame is not None:
video_frames = torch.zeros(1, cond_image.shape[1], frame_num-cond_image.shape[2], target_h, target_w, device=device, dtype=vae.dtype)
if cond_image is not None:
video_frames = torch.zeros(1, 3, frame_num-cond_image.shape[2], target_h, target_w, device=device, dtype=vae.dtype)
padding_frames_pixels_values = torch.concat([cond_image.to(device, vae.dtype), video_frames], dim=2)
# encode
@@ -3279,7 +3289,7 @@ class WanVideoSampler:
cfg[idx],
positive,
text_embeds["negative_prompt_embeds"],
timestep, idx, y, clip_embeds, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
timestep, idx, y, clip_embeds, control_latents, window_vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs)
sampling_pbar.update(1)
@@ -3468,6 +3478,16 @@ class WanVideoSampler:
**scheduler_step_args)[0].squeeze(0)
latent_backwards = torch.flip(latent_backwards, dims=[1])
latent = latent * 0.5 + latent_backwards * 0.5
# differential diffusion inpaint
if masks is not None:
if idx < len(timesteps) - 1:
noise_timestep = timesteps[idx+1]
image_latent = sample_scheduler.scale_noise(
original_image, torch.tensor([noise_timestep]), noise.to(device)
)
mask = masks[idx].to(latent)
latent = image_latent * mask + latent * (1-mask)
#InfiniteTalk first frame handling
if (extra_latents is not None
+2 -2
View File
@@ -1406,7 +1406,7 @@ class VideoVAE38_(VideoVAE_):
attn_scales, self.temperal_upsample, dropout)
def encode(self, x, pbar=False):
def encode(self, x, pbar=True):
self.clear_cache()
x = patchify(x, patch_size=2)
t = x.shape[2]
@@ -1434,7 +1434,7 @@ class VideoVAE38_(VideoVAE_):
return mu
def decode(self, z, pbar=False):
def decode(self, z, pbar=True):
self.clear_cache()
z = z / self.inv_std.to(z) + self.mean.to(z)