import math import numpy as np from safetensors.torch import load_file import timm import torch import torch.nn as nn import torch.nn.functional as F from transformers import Dinov2Model from typing import Optional, Tuple from timm.models.vision_transformer import VisionTransformer DINO_MODEL_NAME = "vit_base_patch14_dinov2.lvd142m" DINO_MEAN = (0.485, 0.456, 0.406) DINO_STD = (0.229, 0.224, 0.225) DINO_PATCH_SIZE = 14 DINO_GRID = 37 LATENT_DOWNSAMPLE_FACTOR = 16 ENCODER_LAYERS = 6 class ResnetBlock(nn.Module): def __init__(self, in_ch: int, out_ch: int) -> None: super().__init__() self.norm1 = nn.GroupNorm(32, in_ch, eps=1e-6, affine=True) self.conv1 = nn.Conv2d(in_ch, out_ch, 3, padding=1) self.norm2 = nn.GroupNorm(32, out_ch, eps=1e-6, affine=True) self.conv2 = nn.Conv2d(out_ch, out_ch, 3, padding=1) self.nin_shortcut = ( nn.Conv2d(in_ch, out_ch, 1) if in_ch != out_ch else nn.Identity() ) def forward(self, x: torch.Tensor) -> torch.Tensor: h = self.conv1(F.silu(self.norm1(x))) h = self.conv2(F.silu(self.norm2(h))) return self.nin_shortcut(x) + h class AttnBlock(nn.Module): def __init__(self, channels: int) -> None: super().__init__() self.norm = nn.GroupNorm(32, channels, eps=1e-6, affine=True) self.q = nn.Conv2d(channels, channels, kernel_size=1) self.k = nn.Conv2d(channels, channels, kernel_size=1) self.v = nn.Conv2d(channels, channels, kernel_size=1) self.proj_out = nn.Conv2d(channels, channels, 1) def forward(self, x: torch.Tensor) -> torch.Tensor: batch, channels, height, width = x.shape normalized = self.norm(x) # (B, C, H, W) -> (B, H*W, C) q = self.q(normalized).flatten(2).transpose(1, 2) k = self.k(normalized).flatten(2).transpose(1, 2) v = self.v(normalized).flatten(2).transpose(1, 2) h = F.scaled_dot_product_attention(q, k, v) h = h.transpose(1, 2).reshape(batch, channels, height, width) return x + self.proj_out(h) class Upsample(nn.Module): def __init__(self, channels: int, scale_factor: float) -> None: super().__init__() self.conv = nn.Conv2d(channels, channels, 3, padding=1) self.scale_factor = scale_factor def forward(self, x: torch.Tensor) -> torch.Tensor: return self.conv(F.interpolate(x, scale_factor=self.scale_factor, mode="nearest")) class Decoder(nn.Module): # ldm/modules/diffusionmodules/model.py def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), num_res_blocks, attn_resolutions, dropout=0.0, resamp_with_conv=True, in_channels, resolution, z_channels, tanh_out=False, use_linear_attn=False, attn_op=AttnBlock, conv3d=False, time_compress=None, **ignorekwargs): super().__init__() self.ch = ch self.temb_ch = 0 self.num_resolutions = len(ch_mult) self.num_res_blocks = num_res_blocks self.resolution = resolution self.in_channels = in_channels self.tanh_out = tanh_out self.carried = False # compute block_in and curr_res at lowest res block_in = ch*ch_mult[self.num_resolutions-1] curr_res = resolution // 2**(self.num_resolutions-1) self.z_shape = (1,z_channels,curr_res,curr_res) # z to block_in self.conv_in = nn.Conv2d( z_channels, block_in, kernel_size=3, stride=1, padding=1) # middle self.mid = MidBlock(block_in) # upsampling self.up = nn.ModuleList() for i_level in reversed(range(self.num_resolutions)): block = nn.ModuleList() attn = nn.ModuleList() block_out = ch*ch_mult[i_level] for i_block in range(self.num_res_blocks+1): block.append(ResnetBlock(block_in, block_out)) block_in = block_out if curr_res in attn_resolutions: attn.append(AttnBlock(block_in)) up = nn.Module() up.block = block up.attn = attn if i_level != 0: scale_factor = 2.0 if time_compress is not None: if i_level > math.log2(time_compress): scale_factor = (1.0, 2.0, 2.0) up.upsample = Upsample(block_in, scale_factor=scale_factor) curr_res = curr_res * 2 self.up.insert(0, up) # prepend to get consistent order self.norm_out = nn.GroupNorm(num_groups=32, num_channels=block_in, eps=1e-6, affine=True) self.conv_out = nn.Conv2d( block_in, out_ch, kernel_size=3, stride=1, padding=1) def forward(self, z, **kwargs) -> list[torch.Tensor]: # timestep embedding temb = None h = self.conv_in(z) h = self.mid(h) if self.carried: h = torch.split(h, 2, dim=2) else: h = [ h ] out = [] conv_carry_in = None for i, h1 in enumerate(h): conv_carry_out = [] if i == len(h) - 1: conv_carry_out = None for i_level in reversed(range(self.num_resolutions)): for i_block in range(self.num_res_blocks+1): h1 = self.up[i_level].block[i_block](h1) if len(self.up[i_level].attn) > 0: assert i == 0 #carried should not happen if attn exists h1 = self.up[i_level].attn[i_block](h1) if i_level != 0: h1 = self.up[i_level].upsample(h1) h1 = self.norm_out(h1) h1 = F.silu(h1) h1 = self.conv_out(h1) if self.tanh_out: h1 = torch.tanh(h1) out.append(h1) conv_carry_in = conv_carry_out return out class MidBlock(nn.Module): def __init__(self, block_in): super().__init__() self.block_1 = ResnetBlock(block_in, block_in) self.attn_1 = AttnBlock(block_in) self.block_2 = ResnetBlock(block_in, block_in) def forward(self, x): x = self.block_1(x) x = self.attn_1(x) x = self.block_2(x) return x class VAE(nn.Module): def __init__(self): super().__init__() self.latent_channels = 64 self._dino = [timm.create_model( DINO_MODEL_NAME, pretrained=True, num_classes=0, dynamic_img_size=True, dynamic_img_pad=True )] # The call_dino() works without changing the output_fmt. # self.dino.patch_embed.output_fmt = 'NCHW' self.dino.eval() # self._dino[0].patch_embed.strict_img_size = False self.embed_dim = self.dino.embed_dim # The original model uses the final six DINO transformer blocks. self.feature_norms = nn.ModuleList([ nn.LayerNorm(self.embed_dim, eps=1e-6) for _ in range(ENCODER_LAYERS) ]) self.encoder_projection = nn.Conv2d( self.embed_dim * ENCODER_LAYERS, self.latent_channels // 2, kernel_size=1, ) self.semantic_projection = nn.Conv2d( self.embed_dim, self.latent_channels // 2, kernel_size=1, ) self.semantic_patch_embeddings = Dinov2PatchEmbeddings(self.embed_dim) # Replace this with a PyTorch implementation. self.decoder = Decoder( ch=128, in_channels=3, out_ch=3, ch_mult=(1, 1, 2, 2, 4), num_res_blocks=2, attn_resolutions=(16,), resolution=256, z_channels=self.latent_channels, tanh_out=True, ) self.register_buffer( "latent_mean", torch.empty(1, self.latent_channels, 1, 1), ) self.register_buffer( "latent_std", torch.empty(1, self.latent_channels, 1, 1), ) self.register_buffer( "dino_mean", torch.tensor(DINO_MEAN).view(1, 3, 1, 1), persistent=False ) self.register_buffer( "dino_std", torch.tensor(DINO_STD).view(1, 3, 1, 1), persistent=False ) @property def dino(self) -> VisionTransformer: return self._dino[0] def call_dino(self, pixel_values, patch_embeddings=None): # This was probably from Dinov2Embeddings.interpolate_pos_encoding() function. if patch_embeddings is None: x = self.dino.patch_embed(pixel_values) else: x = patch_embeddings(pixel_values) if x.shape[-1] == self.dino.embed_dim: x = x.flatten(1, 2) else: x = x.flatten(2).transpose(1, 2) # NCHW -> NLC pos = self.dino.pos_embed.to(x.dtype) patch_pos = pos[:, 1:].reshape(1, DINO_GRID, DINO_GRID, -1).permute(0, 3, 1, 2) grid = (pixel_values.shape[-2] // DINO_PATCH_SIZE, pixel_values.shape[-1] // DINO_PATCH_SIZE) patch_pos = F.interpolate(patch_pos, size=grid, mode="bicubic", antialias=True).flatten(2).transpose(1, 2) pos = torch.cat([pos[:, :1], patch_pos], dim=1).to(x.dtype) B, _, _ = x.shape cls = self.dino.cls_token.expand(B, 1, -1).to(x.dtype) x = torch.cat([cls, x], dim=1) + pos hidden_states = [] for layer in self.dino.blocks: x = layer(x) hidden_states.append(x) return hidden_states def encode(self, x: torch.Tensor) -> torch.Tensor: """ Input: x: [-1, 1], shape [B, 3, H, W] Output: Normalized latent tensor. """ batch_size, _, height, width = x.shape grid_h = max(1, (height + LATENT_DOWNSAMPLE_FACTOR // 2) // LATENT_DOWNSAMPLE_FACTOR) grid_w = max(1, (width + LATENT_DOWNSAMPLE_FACTOR // 2) // LATENT_DOWNSAMPLE_FACTOR) target_height = grid_h * DINO_PATCH_SIZE target_width = grid_w * DINO_PATCH_SIZE x = F.interpolate( x, size=(target_height, target_width), mode="bicubic", align_corners=False, antialias=True, ) # Convert [-1, 1] to [0, 1], then apply ImageNet normalization. # x = (x + 1.0) * 0.5 x = (x - self.dino_mean) / self.dino_std print(self.dino.patch_embed.output_fmt) hidden_states = self.call_dino(x) # Match: # dino_hidden_states(...)[-ENCODER_LAYERS:] features = hidden_states[-ENCODER_LAYERS:] normalized_features = [] for feature, norm in zip(features, self.feature_norms): # Remove CLS token. feature = feature[:, 1:] feature = norm(feature) normalized_features.append(feature) features = torch.cat(normalized_features, dim=-1) # (B, patches, C * ENCODER_LAYERS) features = ( features.transpose(1, 2) .reshape(batch_size, -1, grid_h, grid_w) ) # (B, C * ENCODER_LAYERS, grid_h, grid_w) # Semantic branch uses the final transformer block followed by # the model's final LayerNorm. assert not torch.equal(self.dino.patch_embed.proj.weight, self.semantic_patch_embeddings.projection.weight) semantic = self.call_dino(x, self.semantic_patch_embeddings)[-1][:, 1:] semantic = self.dino.norm(semantic) semantic = ( semantic.transpose(1, 2) .reshape(batch_size, -1, grid_h, grid_w,) ) z = torch.cat([self.encoder_projection(features), self.semantic_projection(semantic)], dim=1) return (z - self.latent_mean) / self.latent_std def decode(self, z: torch.Tensor) -> torch.Tensor: z = z * self.latent_std + self.latent_mean return self.decoder(z) @torch.no_grad() def forward(self, x: torch.Tensor) -> torch.Tensor: return self.decode(self.encode(x)) @staticmethod def from_safetensors(path, device=None): model = VAE() sd = load_file(path) for key in list(sd.keys()): if key.startswith('encoder.'): del sd[key] model.load_state_dict(sd) del sd return model.to(device) def to_numpy_array(self, x: torch.Tensor): # Clip the tanh output. x = (x.detach().clamp(-1, 1) + 1) / 2 * 255.0 if len(x.shape) == 4: x = x.squeeze(0) x = x.cpu().movedim(0, 2) # (CHW) -> (HWC) return x.numpy().astype(np.uint8) def to_tensor(self, image): pixels = torch.from_numpy(np.array(image)).float() pixels = pixels.permute(2, 0, 1) # Dino requires [0, 1] values. pixels = (pixels / 255.0) # * 2.0 - 1.0 return pixels.unsqueeze(0) class Dinov2PatchEmbeddings(nn.Module): def __init__(self, hidden_dim: int): super().__init__() self.projection = nn.Conv2d(3, hidden_dim, kernel_size=14, stride=14) def forward(self, x): # See transformers.models.dinov2.modeling_dinov2 for implementation. return self.projection(x)