From bea3fcad8f94d2e63d5db84c612d4102e56fdec4 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 29 Sep 2025 19:53:35 +0300 Subject: [PATCH] Fix diff diff masking with InfiniteTalk --- nodes_sampler.py | 9 ++++----- 1 file changed, 4 insertions(+), 5 deletions(-) diff --git a/nodes_sampler.py b/nodes_sampler.py index 9cd82ef..ab3e6e0 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -1948,6 +1948,7 @@ class WanVideoSampler: latent_end_idx = latent_start_idx + noise.shape[1] if samples is not None: + noise_mask = samples.get("noise_mask", None) 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]: @@ -1968,14 +1969,12 @@ class WanVideoSampler: noise = input_samples # diff diff prep - noise_mask = samples.get("noise_mask", None) if noise_mask is not None: if len(noise_mask.shape) == 4: noise_mask = noise_mask.squeeze(1) - if noise_mask.shape[0] < noise.shape[1]: - noise_mask = noise_mask.repeat(noise.shape[1] // noise_mask.shape[0], 1, 1) - else: - noise_mask = noise_mask[latent_start_idx:latent_end_idx] + if audio_end_idx > noise_mask.shape[0]: + noise_mask = noise_mask.repeat(audio_end_idx // noise_mask.shape[0], 1, 1) + noise_mask = noise_mask[audio_start_idx:audio_end_idx] noise_mask = torch.nn.functional.interpolate( noise_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W] size=(noise.shape[1], noise.shape[2], noise.shape[3]),