Merge branch 'main' into longvie2

This commit is contained in:
kijai
2025-12-29 15:39:43 +02:00
2 changed files with 5 additions and 7 deletions
+1 -1
View File
@@ -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:
@@ -48,14 +46,14 @@ class FlowMatchSchedulerResMultistep():
timestep_id = torch.argmin((self.timesteps - timestep).abs(), dim=0) timestep_id = torch.argmin((self.timesteps - timestep).abs(), dim=0)
else: else:
timestep_id = torch.argmin((self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1) timestep_id = torch.argmin((self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1) 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 sigma_prev = self.sigmas[timestep_id - 1].reshape(-1, 1, 1, 1) if timestep_id > 0 else sigma
if (timestep_id + 1 >= len(self.sigmas)).any(): if (timestep_id + 1 >= len(self.sigmas)).any():
sigma_next = torch.tensor(0) sigma_next = torch.tensor(0)
else: else:
sigma_next = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1) sigma_next = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
x0_pred = (sample - sigma * model_output) x0_pred = (sample - sigma * model_output)
if sigma_next == 0 or self.prev_model_output is None: if sigma_next == 0 or self.prev_model_output is None:
@@ -73,7 +71,7 @@ class FlowMatchSchedulerResMultistep():
self.old_sigma_next = sigma_next self.old_sigma_next = sigma_next
self.prev_model_output = x0_pred self.prev_model_output = x0_pred
return x return x
def add_noise(self, original_samples, noise, timestep): def add_noise(self, original_samples, noise, timestep):
""" """