only round the gamma once

This commit is contained in:
kijai
2024-10-11 11:08:34 +03:00
parent 538771b836
commit cfc0a6a36b
@@ -56,6 +56,8 @@ class PyramidDiTForVideoGeneration:
self.sample_ratios = sample_ratios
self.corrupt_ratio = corrupt_ratio
dit_path = os.path.join(model_path, model_variant)
@@ -115,6 +117,9 @@ class PyramidDiTForVideoGeneration:
shift=timestep_shift, stages=len(self.stages),
stage_range=stage_range, gamma=scheduler_gamma,
)
#round the gamma as 1/3 seems to have issues on some systems
self.gamma = round(self.scheduler.config.gamma, 5)
print(f"The start sigmas and end sigmas of each stage is Start: {self.scheduler.start_sigmas}, End: {self.scheduler.end_sigmas}, Ori_start: {self.scheduler.ori_start_sigmas}")
self.cfg_rate = 0.1
@@ -191,8 +196,7 @@ class PyramidDiTForVideoGeneration:
return latents
def sample_block_noise(self, bs, ch, temp, height, width):
gamma = round(self.scheduler.config.gamma, 5)
dist = torch.distributions.multivariate_normal.MultivariateNormal(torch.zeros(4), torch.eye(4) * (1 + gamma) - torch.ones(4, 4) * gamma)
dist = torch.distributions.multivariate_normal.MultivariateNormal(torch.zeros(4), torch.eye(4) * (1 + self.gamma) - torch.ones(4, 4) * self.gamma)
block_number = bs * ch * temp * (height // 2) * (width // 2)
noise = torch.stack([dist.sample() for _ in range(block_number)]) # [block number, 4]
noise = rearrange(noise, '(b c t h w) (p q) -> b c t (h p) (w q)',b=bs,c=ch,t=temp,h=height//2,w=width//2,p=2,q=2)
@@ -230,9 +234,8 @@ class PyramidDiTForVideoGeneration:
latents = rearrange(latents, '(b t) c h w -> b c t h w', t=temp)
# Fix the stage
ori_sigma = 1 - self.scheduler.ori_start_sigmas[i_s] # the original coeff of signal
gamma = self.scheduler.config.gamma
alpha = 1 / (math.sqrt(1 + (1 / gamma)) * (1 - ori_sigma) + ori_sigma)
beta = alpha * (1 - ori_sigma) / math.sqrt(gamma)
alpha = 1 / (math.sqrt(1 + (1 / self.gamma)) * (1 - ori_sigma) + ori_sigma)
beta = alpha * (1 - ori_sigma) / math.sqrt(self.gamma)
bs, ch, temp, height, width = latents.shape
noise = self.sample_block_noise(bs, ch, temp, height, width)