From 538771b836f07f862c0a6bd4050d9e724369ad95 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 11 Oct 2024 10:55:46 +0300 Subject: [PATCH] round gamma to avoid covariance_matrix error on Windows --- pyramid_dit/pyramid_dit_for_video_gen_pipeline.py | 12 +----------- 1 file changed, 1 insertion(+), 11 deletions(-) diff --git a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py index 3b0ceb8..123239a 100644 --- a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py +++ b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py @@ -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(