From 617ce55938bab23e579febce61d770a7ce96fbce Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 3 Sep 2025 19:24:47 +0300 Subject: [PATCH] Fix up pusa noise multipliers indexing --- nodes.py | 29 +++++++++++++++------------ wanvideo/schedulers/flowmatch_pusa.py | 2 +- 2 files changed, 17 insertions(+), 14 deletions(-) diff --git a/nodes.py b/nodes.py index 3492c3f..1fca961 100644 --- a/nodes.py +++ b/nodes.py @@ -2334,11 +2334,16 @@ class WanVideoSampler: ).repeat(1, noise.shape[0], 1, 1, 1) # extra latents (Pusa) and 5b - latents_to_insert = add_index = None + latents_to_insert = add_index = noise_multipliers = None extra_latents = image_embeds.get("extra_latents", None) all_indices = [] + noise_multiplier_list = image_embeds.get("pusa_noise_multipliers", None) + if noise_multiplier_list is not None: + noise_multipliers = torch.zeros(latent_video_length) if extra_latents is not None and transformer.multitalk_model_type.lower() != "infinitetalk": - for entry in extra_latents: + if noise_multiplier_list is not None: + noise_multiplier_list = list(noise_multiplier_list) + [1.0] * (len(all_indices) - len(noise_multiplier_list)) + for i, entry in enumerate(extra_latents): add_index = entry["index"] num_extra_frames = entry["samples"].shape[2] # Handle negative indices @@ -2349,6 +2354,10 @@ class WanVideoSampler: noise[:, add_index:add_index+num_extra_frames] = entry["samples"].to(noise) log.info(f"Adding extra samples to latent indices {add_index} to {add_index+num_extra_frames-1}") all_indices.extend(range(add_index, add_index+num_extra_frames)) + if noise_multipliers is not None: + for i, idx in enumerate(all_indices): + noise_multipliers[idx] = noise_multiplier_list[i] + log.info(f"Using Pusa noise multipliers: {noise_multipliers}") latent = noise.to(device) @@ -3007,15 +3016,9 @@ class WanVideoSampler: latent = current_latent if is_pusa and all_indices: - noise_multipliers = image_embeds.get("pusa_noise_multipliers", None) pusa_noisy_steps = image_embeds.get("pusa_noisy_steps", -1) if pusa_noisy_steps == -1: pusa_noisy_steps = len(timesteps) - if noise_multipliers is not None: - if len(noise_multipliers) < len(all_indices): - noise_multipliers = list(noise_multipliers) + [1.0] * (len(all_indices) - len(noise_multipliers)) - #pusa_noise_flag = torch.zeros(latent_video_length) - try: pbar = ProgressBar(len(timesteps)) #region main loop start @@ -3050,10 +3053,10 @@ class WanVideoSampler: if is_pusa: scheduler_step_args["cond_frame_latent_indices"] = all_indices scheduler_step_args["noise_multipliers"] = noise_multipliers - for latent_idx, multiplier in zip(all_indices, noise_multipliers): - timestep[:, latent_idx] = timestep[:, latent_idx] * multiplier + for latent_idx in all_indices: + timestep[:, latent_idx] = timestep[:, latent_idx] * noise_multipliers[latent_idx] # add noise for conditioning frames if multiplier > 0 - if idx < pusa_noisy_steps and multiplier > 0: + if idx < pusa_noisy_steps and noise_multipliers[latent_idx] > 0: latent_size = (1, latent.shape[0], latent.shape[1], latent.shape[2], latent.shape[3]) noise_for_cond = torch.randn(latent_size, generator=seed_g, device=torch.device("cpu")) timestep_cond = torch.ones_like(timestep) * timestep.max() @@ -3062,10 +3065,10 @@ class WanVideoSampler: latent[:, latent_idx:latent_idx+1].to(device), noise_for_cond[:, :, latent_idx:latent_idx+1].to(device), timestep_cond[:, latent_idx:latent_idx+1].to(device), - noise_multiplier=multiplier) + noise_multiplier=noise_multipliers[latent_idx]) else: timestep[:, all_indices] = 0 - print("timestep: ", timestep) + #print("timestep: ", timestep) ### latent shift if latent_shift_loop: diff --git a/wanvideo/schedulers/flowmatch_pusa.py b/wanvideo/schedulers/flowmatch_pusa.py index 20a2678..8c708cf 100644 --- a/wanvideo/schedulers/flowmatch_pusa.py +++ b/wanvideo/schedulers/flowmatch_pusa.py @@ -76,7 +76,7 @@ class FlowMatchSchedulerPusa(): sigma[:, :, latent_idx] = 0 sigma_[:, :, latent_idx] = 0 continue - multiplier = noise_multipliers.get(latent_idx, 1.0) if isinstance(noise_multipliers, dict) else noise_multipliers[latent_idx] + multiplier = noise_multipliers[latent_idx] sigma[:, :, latent_idx] = sigma[:, :, latent_idx] * multiplier sigma_[:, :, latent_idx] = sigma_[:, :, latent_idx] * multiplier