Fix differential diffusion, VACE tweaks
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user