Instructions to use nebulette/nano-transformers with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use nebulette/nano-transformers with Transformers:
# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("nebulette/nano-transformers", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download vae.py from nebulette/nano-transformers: direct link, hf CLI and curl.
- Browser
- Download file 13.5 kB
-
https://huggingface.co/nebulette/nano-transformers/resolve/main/vae.py
- Command line
-
hf download hf://nebulette/nano-transformers/vae.py
-
curl -L -o vae.py https://huggingface.co/nebulette/nano-transformers/resolve/main/vae.py
13.5 kB
| 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 | |
| ) | |
| 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) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return self.decode(self.encode(x)) | |
| 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) | |