Fix res_multistep.step to set sigma_next correctly when ending early

Use the next sigma from the schedule when sampler ends early, only assume `sigma_next = 0` when on the true final step.
This commit is contained in:
FB
2025-08-11 14:17:08 +01:00
committed by GitHub
parent 7c93aea182
commit 1e091d181b
@@ -51,7 +51,7 @@ class FlowMatchSchedulerResMultistep():
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
sigma_prev = self.sigmas[timestep_id - 1].reshape(-1, 1, 1, 1) if timestep_id > 0 else sigma
if (timestep_id + 1 >= len(self.timesteps)).any():
if (timestep_id + 1 >= len(self.sigmas)).any():
sigma_next = torch.tensor(0)
else:
sigma_next = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)