Update nodes.py

This commit is contained in:
kijai
2025-08-21 18:26:06 +03:00
parent bd5f76a3c8
commit bd0a634ec0
+42 -28
View File
@@ -1305,10 +1305,6 @@ 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,)
@@ -2051,25 +2047,25 @@ class WanVideoSampler:
else:
noise = input_samples
mask = samples.get("noise_mask", None)
if mask is not None:
log.info(f"Latent mask shape: {mask.shape}")
noise_mask = samples.get("noise_mask", None)
if noise_mask is not None:
log.info(f"Latent noise_mask shape: {noise_mask.shape}")
original_image = input_samples.to(device)
if len(mask.shape) == 4:
mask = mask.squeeze(1)
if len(noise_mask.shape) == 4:
noise_mask = noise_mask.squeeze(1)
mask = torch.nn.functional.interpolate(
mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
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]),
mode='trilinear',
align_corners=False
).squeeze(0) # Remove batch dim, keep channel dim
# Add batch & channel dims for final output
mask = mask.unsqueeze(0).repeat(1, noise.shape[0], 1, 1, 1)
noise_mask = noise_mask.unsqueeze(0).repeat(1, noise.shape[0], 1, 1, 1)
if mask.shape[2] != noise.shape[1]:
mask = torch.cat([torch.zeros(1, noise.shape[0], noise.shape[1] - mask.shape[2], noise.shape[2], noise.shape[3]), mask], dim=2)
if noise_mask.shape[2] != noise.shape[1]:
noise_mask = torch.cat([torch.zeros(1, noise.shape[0], noise.shape[1] - noise_mask.shape[2], noise.shape[2], noise.shape[3]), noise_mask], dim=2)
# extra latents (Pusa) and 5b
latents_to_insert = add_index = None
@@ -2628,11 +2624,11 @@ class WanVideoSampler:
# diff diff prep
masks = None
if not multitalk_sampling and samples is not None and mask is not None:
mask = 1 - mask
if not multitalk_sampling and samples is not None and noise_mask is not None:
noise_mask = 1 - noise_mask
thresholds = torch.arange(len(timesteps), dtype=original_image.dtype) / len(timesteps)
thresholds = thresholds.unsqueeze(1).unsqueeze(1).unsqueeze(1).unsqueeze(1).to(device)
masks = mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)
masks = noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)
masks = masks > thresholds
latent_shift_loop = False
@@ -3137,6 +3133,15 @@ class WanVideoSampler:
else:
noise = input_samples
# diff diff prep
masks = None
if noise_mask is not None:
noise_mask = 1 - noise_mask
thresholds = torch.arange(len(timesteps), dtype=original_image.dtype) / len(timesteps)
thresholds = thresholds.unsqueeze(1).unsqueeze(1).unsqueeze(1).unsqueeze(1).to(device)
masks = noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)
masks = masks > thresholds
window_vace_data = None
if vace_data is not None:
window_vace_data = []
@@ -3153,7 +3158,7 @@ class WanVideoSampler:
"seq_len": vace_entry["seq_len"]
})
# get mask
# get image cond mask
msk = torch.ones(1, frame_num, lat_h, lat_w, device=device)
if mode == "multitalk":
msk[:, cur_motion_frames_num:] = 0
@@ -3308,13 +3313,22 @@ class WanVideoSampler:
latent = latent + noise_pred * dt[:, None, None, None]
else:
latent = latent.to(intermediate_device)
temp_x0 = sample_scheduler.step(
noise_pred.unsqueeze(0),
timestep,
latent.unsqueeze(0),
**scheduler_step_args)[0]
latent = temp_x0.squeeze(0)
# 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)
# injecting motion frames
if not is_first_clip and mode == "multitalk":
@@ -3479,6 +3493,15 @@ class WanVideoSampler:
latent_backwards = torch.flip(latent_backwards, dims=[1])
latent = latent * 0.5 + latent_backwards * 0.5
#InfiniteTalk first frame handling
if (extra_latents is not None
and not multitalk_sampling
and transformer.multitalk_model_type=="InfiniteTalk"):
for entry in extra_latents:
add_index = entry["index"]
num_extra_frames = entry["samples"].shape[2]
latent[:, add_index:add_index+num_extra_frames] = entry["samples"].to(latent)
# differential diffusion inpaint
if masks is not None:
if idx < len(timesteps) - 1:
@@ -3488,15 +3511,6 @@ class WanVideoSampler:
)
mask = masks[idx].to(latent)
latent = image_latent * mask + latent * (1-mask)
#InfiniteTalk first frame handling
if (extra_latents is not None
and not multitalk_sampling
and transformer.multitalk_model_type=="InfiniteTalk"):
for entry in extra_latents:
add_index = entry["index"]
num_extra_frames = entry["samples"].shape[2]
latent[:, add_index:add_index+num_extra_frames] = entry["samples"].to(latent)
if freeinit_args is not None:
current_latent = latent.clone()