From 4c711dddc4bf04d07025a6f33b707fa5829d6add Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Wed, 22 Nov 2023 20:46:34 +0100 Subject: [PATCH] Fix decoder init --- README.md | 2 +- VAE/models/temporal_ae.py | 8 +++----- 2 files changed, 4 insertions(+), 6 deletions(-) diff --git a/README.md b/README.md index 350adbf..151ca3e 100644 --- a/README.md +++ b/README.md @@ -133,7 +133,7 @@ This now works thanks to the work of @mrsteyk and @madebyollin - [Gist with more This is the VAE that comes baked into the [Stable Video Diffusion](https://stability.ai/news/stable-video-diffusion-open-ai-video-model) model. -It doesn't seem particularly good as a normal VAE (color issues, pretty bad with finer details). Parts of it also seem to be missing weights and/or are never called, so I'm not even sure it does anything special with batch sizes greater than 1 (i.e. the deflickering part). +It doesn't seem particularly good as a normal VAE (color issues, pretty bad with finer details). Still for completeness sake the code to run it is mostly implemented. To obtain the weights just extract them from the sdv model: diff --git a/VAE/models/temporal_ae.py b/VAE/models/temporal_ae.py index 5ddb8b6..690ca94 100644 --- a/VAE/models/temporal_ae.py +++ b/VAE/models/temporal_ae.py @@ -9,7 +9,7 @@ from comfy import model_management from .kl import ( Encoder, Decoder, Upsample, Normalize, AttnBlock, ResnetBlock, #MemoryEfficientAttnBlock, - DiagonalGaussianDistribution, nonlinearity + DiagonalGaussianDistribution, nonlinearity, make_attn ) class AutoencoderKL(nn.Module): @@ -130,11 +130,9 @@ class VideoDecoder(nn.Module): alpha=self.alpha, merge_strategy=self.merge_strategy, ) - self.mid.attn_1 = make_time_attn( + self.mid.attn_1 = make_attn( block_in, attn_type=attn_type, - alpha=self.alpha, - merge_strategy=self.merge_strategy, ) self.mid.block_2 = VideoResBlock( in_channels=block_in, @@ -211,7 +209,7 @@ class VideoDecoder(nn.Module): # middle h = self.mid.block_1(h, temb, **kwargs) - h = self.mid.attn_1(h, **kwargs) + h = self.mid.attn_1(h) h = self.mid.block_2(h, temb, **kwargs) # upsampling