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