Fix decoder init

This commit is contained in:
City
2023-11-22 20:46:34 +01:00
parent 06fdd1f2d6
commit 4c711dddc4
2 changed files with 4 additions and 6 deletions
+1 -1
View File
@@ -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:
+3 -5
View File
@@ -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