Update nodes.py
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user