Download nanosaur2_support/vae.py from levzalt/Nanosaur2-Inpaint-ControlNet: direct link, hf CLI and curl.
- Browser
- Download file 4.1 kB
-
https://huggingface.co/levzalt/Nanosaur2-Inpaint-ControlNet/resolve/main/nanosaur2_support/vae.py
- Command line
-
hf download hf://levzalt/Nanosaur2-Inpaint-ControlNet/nanosaur2_support/vae.py
-
curl -L -o vae.py https://huggingface.co/levzalt/Nanosaur2-Inpaint-ControlNet/resolve/main/nanosaur2_support/vae.py
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) | |