hugging-apps/echo-memory
0
1import torch2from .sd_unet import ResnetBlock, DownSampler3from .sd_vae_encoder import VAEAttentionBlock, SDVAEEncoderStateDictConverter4from .tiler import TileWorker5from einops import rearrange6 7 8class SD3VAEEncoder(torch.nn.Module):9 def __init__(self):10 super().__init__()11 self.scaling_factor = 1.5305 # Different from SD 1.x12 self.shift_factor = 0.0609 # Different from SD 1.x13 self.conv_in = torch.nn.Conv2d(3, 128, kernel_size=3, padding=1)14 15 self.blocks = torch.nn.ModuleList([16 # DownEncoderBlock2D17 ResnetBlock(128, 128, eps=1e-6),18 ResnetBlock(128, 128, eps=1e-6),19 DownSampler(128, padding=0, extra_padding=True),20 # DownEncoderBlock2D21 ResnetBlock(128, 256, eps=1e-6),22 ResnetBlock(256, 256, eps=1e-6),23 DownSampler(256, padding=0, extra_padding=True),24 # DownEncoderBlock2D25 ResnetBlock(256, 512, eps=1e-6),26 ResnetBlock(512, 512, eps=1e-6),27 DownSampler(512, padding=0, extra_padding=True),28 # DownEncoderBlock2D29 ResnetBlock(512, 512, eps=1e-6),30 ResnetBlock(512, 512, eps=1e-6),31 # UNetMidBlock2D32 ResnetBlock(512, 512, eps=1e-6),33 VAEAttentionBlock(1, 512, 512, 1, eps=1e-6),34 ResnetBlock(512, 512, eps=1e-6),35 ])36 37 self.conv_norm_out = torch.nn.GroupNorm(num_channels=512, num_groups=32, eps=1e-6)38 self.conv_act = torch.nn.SiLU()39 self.conv_out = torch.nn.Conv2d(512, 32, kernel_size=3, padding=1)40 41 def tiled_forward(self, sample, tile_size=64, tile_stride=32):42 hidden_states = TileWorker().tiled_forward(43 lambda x: self.forward(x),44 sample,45 tile_size,46 tile_stride,47 tile_device=sample.device,48 tile_dtype=sample.dtype49 )50 return hidden_states51 52 def forward(self, sample, tiled=False, tile_size=64, tile_stride=32, **kwargs):53 # For VAE Decoder, we do not need to apply the tiler on each layer.54 if tiled:55 return self.tiled_forward(sample, tile_size=tile_size, tile_stride=tile_stride)56 57 # 1. pre-process58 hidden_states = self.conv_in(sample)59 time_emb = None60 text_emb = None61 res_stack = None62 63 # 2. blocks64 for i, block in enumerate(self.blocks):65 hidden_states, time_emb, text_emb, res_stack = block(hidden_states, time_emb, text_emb, res_stack)66 67 # 3. output68 hidden_states = self.conv_norm_out(hidden_states)69 hidden_states = self.conv_act(hidden_states)70 hidden_states = self.conv_out(hidden_states)71 hidden_states = hidden_states[:, :16]72 hidden_states = (hidden_states - self.shift_factor) * self.scaling_factor73 74 return hidden_states75 76 def encode_video(self, sample, batch_size=8):77 B = sample.shape[0]78 hidden_states = []79 80 for i in range(0, sample.shape[2], batch_size):81 82 j = min(i + batch_size, sample.shape[2])83 sample_batch = rearrange(sample[:,:,i:j], "B C T H W -> (B T) C H W")84 85 hidden_states_batch = self(sample_batch)86 hidden_states_batch = rearrange(hidden_states_batch, "(B T) C H W -> B C T H W", B=B)87 88 hidden_states.append(hidden_states_batch)89 90 hidden_states = torch.concat(hidden_states, dim=2)91 return hidden_states92 93 @staticmethod94 def state_dict_converter():95 return SDVAEEncoderStateDictConverter()96 