From d9be7b9c263905b6079630543b12db574b298f83 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 1 Aug 2025 21:21:55 +0300 Subject: [PATCH] Allow other schedulers for 5B --- nodes.py | 32 +++++++++++++++++++++----------- 1 file changed, 21 insertions(+), 11 deletions(-) diff --git a/nodes.py b/nodes.py index fc8dc11..d2dd9fe 100644 --- a/nodes.py +++ b/nodes.py @@ -1802,17 +1802,18 @@ class WanVideoSampler: mask = torch.cat([torch.zeros(1, noise.shape[0], noise.shape[1] - mask.shape[2], noise.shape[2], noise.shape[3]), mask], dim=2) # extra latents (Pusa) and 5b - encoded_image_latents = None + latents_to_insert = None if (extra_latents := image_embeds.get("extra_latents", None)) is not None: - encoded_image_latents = extra_latents["samples"].squeeze(0).to(noise) + latents_to_insert = extra_latents["samples"].squeeze(0).to(noise) + num_latents_to_insert = latents_to_insert.shape[1] if (empty_latent_indices := extra_latents.get("empty_latent_indices", None)) is not None and len(empty_latent_indices) > 0: - noise_out = encoded_image_latents.clone() + noise_out = latents_to_insert.clone() for idx in empty_latent_indices: #print(f"Adding noise to Empty latent index: {idx}") noise_out[:, idx] = noise[:, idx] noise = noise_out else: - noise[:,0:encoded_image_latents.shape[1]] = encoded_image_latents + noise[:,0:num_latents_to_insert] = latents_to_insert latent = noise.to(device) @@ -2415,7 +2416,7 @@ class WanVideoSampler: latent_model_input = latent.to(device) timestep = torch.tensor([t]).to(device) - if scheduler == "flowmatch_pusa" or (is_5b and encoded_image_latents is not None): + if scheduler == "flowmatch_pusa" or (is_5b and latents_to_insert is not None): timestep = timestep.unsqueeze(1).repeat(1, latent_video_length) if extra_latents is not None: if empty_latent_indices is not None and len(empty_latent_indices) > 0: @@ -2423,7 +2424,7 @@ class WanVideoSampler: non_noise_indices = [i for i in range(timestep.shape[1]) if i not in empty_latent_indices] timestep[:, non_noise_indices] = 0 else: - timestep[:,0:encoded_image_latents.shape[1]] = 0 + timestep[:,0:num_latents_to_insert] = 0 #print(f"timestep: {timestep}") current_step_percentage = idx / len(timesteps) @@ -2946,11 +2947,20 @@ class WanVideoSampler: if flowedit_args is None: latent = latent.to(intermediate_device) - latent = sample_scheduler.step( - noise_pred[:, :orig_noise_len].unsqueeze(0) if recammaster is not None else noise_pred.unsqueeze(0), - timestep, - latent[:, :orig_noise_len].unsqueeze(0) if recammaster is not None else latent.unsqueeze(0), - **scheduler_step_args)[0].squeeze(0) + + if len(timestep.shape) != 1 and scheduler != "flowmatch_pusa": #pusa and 5b + latent_slice = sample_scheduler.step( + noise_pred[:, num_latents_to_insert:].unsqueeze(0), + timestep.flatten()[-1], + latent[:, num_latents_to_insert:].unsqueeze(0), + **scheduler_step_args)[0].squeeze(0) + latent = torch.cat([latent[:, :num_latents_to_insert], latent_slice], dim=1) + else: + latent = sample_scheduler.step( + noise_pred[:, :orig_noise_len].unsqueeze(0) if recammaster is not None else noise_pred.unsqueeze(0), + timestep, + latent[:, :orig_noise_len].unsqueeze(0) if recammaster is not None else latent.unsqueeze(0), + **scheduler_step_args)[0].squeeze(0) if freeinit_args is not None: current_latent = latent.clone()