Merge branch 'main' into longvie2
This commit is contained in:
@@ -34,7 +34,7 @@ class ERSDEScheduler():
|
|||||||
sigmas.append(0.0)
|
sigmas.append(0.0)
|
||||||
self.sigmas = torch.FloatTensor(sigmas)
|
self.sigmas = torch.FloatTensor(sigmas)
|
||||||
self.sigmas = self.shift * self.sigmas / (1 + (self.shift - 1) * self.sigmas)
|
self.sigmas = self.shift * 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
|
||||||
self.step_index = 0
|
self.step_index = 0
|
||||||
self.old_denoised = None
|
self.old_denoised = None
|
||||||
self.old_denoised_d = None
|
self.old_denoised_d = None
|
||||||
|
|||||||
@@ -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