Fix res_multistep steps
This commit is contained in:
@@ -35,9 +35,7 @@ class FlowMatchSchedulerResMultistep():
|
||||
self.sigmas = torch.FloatTensor(sigmas)
|
||||
self.sigmas = self.shift * self.sigmas / \
|
||||
(1 + (self.shift - 1) * self.sigmas)
|
||||
self.timesteps = self.sigmas * self.num_train_timesteps
|
||||
#print(f"Timesteps: {self.timesteps}, Sigmas: {self.sigmas}")
|
||||
|
||||
self.timesteps = self.sigmas[:-1] * self.num_train_timesteps
|
||||
|
||||
def step(self, model_output, timestep, sample):
|
||||
if timestep.ndim == 2:
|
||||
|
||||
Reference in New Issue
Block a user