flowmatch_pusa split sampling fixes

This commit is contained in:
kijai
2025-09-01 13:34:49 +03:00
parent d4d942c5ee
commit 43d8a2c8a5
3 changed files with 27 additions and 33 deletions
+3 -3
View File
@@ -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))
+4 -4
View File
@@ -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)
+20 -26
View File
@@ -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