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