round gamma to avoid covariance_matrix error on Windows

This commit is contained in:
kijai
2024-10-11 10:55:46 +03:00
parent 265f46b54b
commit 538771b836
@@ -191,22 +191,12 @@ class PyramidDiTForVideoGeneration:
return latents
def sample_block_noise(self, bs, ch, temp, height, width):
gamma = self.scheduler.config.gamma
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)
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)
return noise
# def sample_block_noise(self, bs, ch, temp, height, width):
# gamma = self.scheduler.config.gamma
# epsilon = 1e-5 # Small value to ensure positive definiteness
# covariance_matrix = torch.eye(4) * (1 + gamma) - torch.ones(4, 4) * gamma + torch.eye(4) * epsilon
# dist = torch.distributions.multivariate_normal.MultivariateNormal(torch.zeros(4), covariance_matrix)
# 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)
# return noise
@torch.no_grad()
def generate_one_unit(