diff --git a/nodes.py b/nodes.py index d5497ee..05edbe4 100644 --- a/nodes.py +++ b/nodes.py @@ -1845,7 +1845,6 @@ class WanVideoSampler: else: timesteps = torch.tensor([1000, 750, 500, 250], device=device) - log.info(f"timesteps: {timesteps}") total_steps = steps steps = len(timesteps) @@ -2315,8 +2314,9 @@ class WanVideoSampler: for entry in extra_latents: add_index = entry["index"] num_extra_frames = entry["samples"].shape[2] - 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}") + if start_step == 0: + 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)) diff --git a/wanvideo/schedulers/__init__.py b/wanvideo/schedulers/__init__.py index 8f67870..414266b 100644 --- a/wanvideo/schedulers/__init__.py +++ b/wanvideo/schedulers/__init__.py @@ -89,10 +89,10 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo sample_scheduler.timesteps = denoising_step_list[:steps].clone().detach().to(device) sample_scheduler.sigmas = torch.cat([sample_scheduler.timesteps / 1000, torch.tensor([0.0], device=device)]) elif 'flowmatch_pusa' in scheduler: - sample_scheduler = FlowMatchSchedulerPusa( - shift=shift, sigma_min=0.0, extra_one_step=True - ) - sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength, shift=shift, sigmas=sigmas[:-1].tolist() if sigmas is not None else None) + sample_scheduler = FlowMatchSchedulerPusa(shift=shift, sigma_min=0.0, extra_one_step=True) + sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength, shift=shift, + sigmas=sigmas[:-1].tolist() if sigmas is not None else None, + finalize=(end_step == -1 or end_step >= steps)) elif scheduler == 'res_multistep': sample_scheduler = FlowMatchSchedulerResMultistep(shift=shift) sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength, sigmas=sigmas[:-1].tolist() if sigmas is not None else None) diff --git a/wanvideo/schedulers/flowmatch_pusa.py b/wanvideo/schedulers/flowmatch_pusa.py index 94dd9d9..ab3133e 100644 --- a/wanvideo/schedulers/flowmatch_pusa.py +++ b/wanvideo/schedulers/flowmatch_pusa.py @@ -15,15 +15,16 @@ class FlowMatchSchedulerPusa(): self.set_timesteps(num_inference_steps) - def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, training=False, shift=None, sigmas=None): + def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, training=False, shift=None, sigmas=None, finalize=False): if shift is not None: self.shift = shift sigma_start = self.sigma_min + (self.sigma_max - self.sigma_min) * denoising_strength if sigmas is None: + steps = num_inference_steps + 1 if not finalize else num_inference_steps if self.extra_one_step: - self.sigmas = torch.linspace(sigma_start, self.sigma_min, num_inference_steps + 1)[:-1] + self.sigmas = torch.linspace(sigma_start, self.sigma_min, steps)[:-1] else: - self.sigmas = torch.linspace(sigma_start, self.sigma_min, num_inference_steps) + self.sigmas = torch.linspace(sigma_start, self.sigma_min, steps) if self.inverse_timesteps: self.sigmas = torch.flip(self.sigmas, dims=[0]) else: @@ -42,7 +43,6 @@ class FlowMatchSchedulerPusa(): def step(self, model_output, timestep, sample, to_final=False, **kwargs): if isinstance(timestep, torch.Tensor): - # timestep = timestep.cpu() self.timesteps = self.timesteps.to(timestep.device) self.sigmas = self.sigmas.to(timestep.device) model_output = model_output.to(timestep.device) @@ -50,35 +50,29 @@ class FlowMatchSchedulerPusa(): if len(timestep.shape) == 1: timestep_id = torch.argmin((self.timesteps - timestep).abs()) sigma = self.sigmas[timestep_id] - if to_final or timestep_id + 1 >= len(self.timesteps): - sigma_ = 1 if (self.inverse_timesteps or self.reverse_sigmas) else 0 - else: + if timestep_id + 1 < len(self.sigmas): sigma_ = self.sigmas[timestep_id + 1] + else: + sigma_ = 0.0 # Only zero at the true end + # Zero sigma and sigma_ for indices where timestep == 0 + if torch.any(timestep == 0): + sigma = torch.where(timestep == 0, torch.zeros_like(sigma), sigma) + sigma_ = torch.where(timestep == 0, torch.zeros_like(sigma_), sigma_) prev_sample = sample + model_output * (sigma_ - sigma) else: timestep_id = torch.argmin((self.timesteps.unsqueeze(1) - timestep).abs(), dim=0) sigma = self.sigmas[timestep_id].unsqueeze(0).unsqueeze(1).unsqueeze(3).unsqueeze(4).to(sample.device) - # Handle sigma_ calculation for each timestep_id element - if to_final or torch.any(timestep_id + 1 >= len(self.timesteps)): - default_value = 1.0 if (self.inverse_timesteps or self.reverse_sigmas) else 0.0 - # Create sigma_ with the same dtype as self.sigmas - sigma_ = torch.ones_like(timestep_id, dtype=self.sigmas.dtype, device=sample.device) * default_value - valid_indices = timestep_id + 1 < len(self.timesteps) - if torch.any(valid_indices): - # Convert indices to the appropriate type for indexing - valid_timestep_ids = timestep_id[valid_indices] - sigma_[valid_indices] = self.sigmas[(valid_timestep_ids + 1).to(torch.long)] - else: - sigma_ = self.sigmas[(timestep_id + 1).to(torch.long)] - - - # Reshape sigma_ to match sigma's dimensions for the operation + sigma_ = torch.zeros_like(timestep_id, dtype=self.sigmas.dtype, device=sample.device) + last_step = (timestep_id == len(self.sigmas) - 1) + not_last = ~last_step + sigma_[not_last] = self.sigmas[(timestep_id[not_last] + 1).to(torch.long)] sigma_ = sigma_.unsqueeze(0).unsqueeze(1).unsqueeze(3).unsqueeze(4).to(sample.device) - if torch.any(timestep == 0): - zero_indices = torch.where(timestep == 0)[1].to(torch.long) - sigma[:,:,zero_indices] = 0 + # Zero sigma and sigma_ for batch indices where timestep == 0 + zero_indices = (timestep[0] == 0).nonzero(as_tuple=True)[0] + if zero_indices.numel() > 0: + sigma[:,:,zero_indices,:,:] = 0 + sigma_[:,:,zero_indices,:,:] = 0 #print("sigma", sigma[0,0,:,0,0], '\n', "sigma_", sigma_[0,0,:,0,0]) - prev_sample = sample + model_output * (sigma_ - sigma) return prev_sample