levzalt's picture
Release step-1000 Nanosaur2 inpainting adapter and ComfyUI support
943b415 verified
Raw History Blame Contribute Delete
4.1 kB
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)