Transformers
nebulette's picture
Upload 4 files
069dce3 verified
Raw History Blame
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
)
@property
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)
@torch.no_grad()
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.decode(self.encode(x))
@staticmethod
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)