File size: 4,102 Bytes
943b415
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
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)