import torch import torch.nn as nn import torch.nn.functional as F import comfy.model_management import comfy.ops from comfy.image_encoders.dino2 import Dino2PatchEmbeddings, Dinov2Model from comfy.ldm.modules.attention import optimized_attention_for_device from comfy.ldm.modules.diffusionmodules.model import Decoder ops = comfy.ops.disable_weight_init DINO_CONFIG = {"num_hidden_layers": 12, "hidden_size": 768, "num_attention_heads": 12, "layer_norm_eps": 1e-6, "use_swiglu_ffn": False, "use_mask_token": False} 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 def dino_hidden_states(dino, pixels, patch_embeddings): embeddings = dino.embeddings x = patch_embeddings(pixels) pos = comfy.model_management.cast_to(embeddings.position_embeddings, dtype=torch.float32, device=x.device) patch_pos = pos[:, 1:].reshape(1, DINO_GRID, DINO_GRID, -1).permute(0, 3, 1, 2) grid = (pixels.shape[-2] // DINO_PATCH_SIZE, pixels.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) cls = comfy.model_management.cast_to(embeddings.cls_token, dtype=x.dtype, device=x.device) x = torch.cat([cls.expand(x.shape[0], -1, -1), x], dim=1) + pos optimized_attention = optimized_attention_for_device(x.device, False, small_input=True) hidden_states = [] for layer in dino.encoder.layer: x = layer(x, optimized_attention) hidden_states.append(x) return hidden_states class Nanosaur2VAE(nn.Module): def __init__(self, latent_channels=64): super().__init__() embed_dim = DINO_CONFIG["hidden_size"] self.encoder = Dinov2Model(DINO_CONFIG, None, None, ops) self.semantic_patch_embeddings = Dino2PatchEmbeddings(embed_dim, patch_size=DINO_PATCH_SIZE, operations=ops) self.feature_norms = nn.ModuleList([ops.LayerNorm(embed_dim) for _ in range(ENCODER_LAYERS)]) self.encoder_projection = ops.Conv2d(embed_dim * ENCODER_LAYERS, latent_channels // 2, 1) self.semantic_projection = ops.Conv2d(embed_dim, latent_channels // 2, 1) self.decoder = Decoder(ch=128, out_ch=3, ch_mult=(1, 1, 2, 2, 4), num_res_blocks=2, attn_resolutions=(16,), in_channels=3, resolution=256, z_channels=latent_channels, tanh_out=True) self.register_buffer("latent_mean", torch.empty(1, latent_channels, 1, 1)) self.register_buffer("latent_std", torch.empty(1, latent_channels, 1, 1)) def encode(self, x): b, _, h, w = x.shape grid_h = max(1, (h + LATENT_DOWNSAMPLE_FACTOR // 2) // LATENT_DOWNSAMPLE_FACTOR) grid_w = max(1, (w + LATENT_DOWNSAMPLE_FACTOR // 2) // LATENT_DOWNSAMPLE_FACTOR) x = F.interpolate(x, size=(grid_h * DINO_PATCH_SIZE, grid_w * DINO_PATCH_SIZE), mode="bicubic", align_corners=False, antialias=True) mean = torch.tensor(DINO_MEAN, device=x.device, dtype=x.dtype).view(1, 3, 1, 1) std = torch.tensor(DINO_STD, device=x.device, dtype=x.dtype).view(1, 3, 1, 1) x = (x.add(1.0).mul(0.5) - mean) / std features = dino_hidden_states(self.encoder, x, self.encoder.embeddings.patch_embeddings)[-ENCODER_LAYERS:] features = [norm(f[:, 1:]) for f, norm in zip(features, self.feature_norms)] features = torch.cat(features, dim=-1).transpose(1, 2).reshape(b, -1, grid_h, grid_w) semantic = self.encoder.layernorm(dino_hidden_states(self.encoder, x, self.semantic_patch_embeddings)[-1][:, 1:]) semantic = semantic.transpose(1, 2).reshape(b, -1, grid_h, grid_w) z = torch.cat([self.encoder_projection(features), self.semantic_projection(semantic)], dim=1) return (z - comfy.ops.cast_to_input(self.latent_mean, z)) / comfy.ops.cast_to_input(self.latent_std, z) def decode(self, z): z = z * comfy.ops.cast_to_input(self.latent_std, z) + comfy.ops.cast_to_input(self.latent_mean, z) return self.decoder(z)