v1.0.3 update
This commit is contained in:
@@ -1,4 +1,4 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.model.backbone import (autoencoder, image, unet, utils,
|
||||
video)
|
||||
from scepter.modules.model.backbone import (autoencoder, image, mmdit, pixart,
|
||||
unet, utils, video)
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
from .sd3 import MMDiT
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,2 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
from .pixart_alpha import PixArt
|
||||
@@ -0,0 +1,503 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
# All rights reserved.
|
||||
# This file contains code that is adapted from
|
||||
# timm: https://github.com/huggingface/pytorch-image-models
|
||||
# pixart: https://github.com/PixArt-alpha/PixArt-alpha
|
||||
|
||||
# This source code is also licensed under the license found in the
|
||||
# LICENSE file in the root directory of this source tree.
|
||||
# --------------------------------------------------------
|
||||
|
||||
from collections import OrderedDict
|
||||
# MAE: https://github.com/facebookresearch/mae/blob/main/models_mae.py
|
||||
# --------------------------------------------------------
|
||||
from typing import Iterable
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
# References:
|
||||
# GLIDE: https://github.com/openai/glide-text2im
|
||||
from torch.utils.checkpoint import checkpoint, checkpoint_sequential
|
||||
|
||||
from scepter.modules.model.backbone.transformer.attention import \
|
||||
MultiHeadAttention
|
||||
from scepter.modules.model.backbone.transformer.layers import (
|
||||
CaptionEmbedder, DropPath, LabelEmbedder, Mlp, SizeEmbedder,
|
||||
TimestepEmbedder, modulate)
|
||||
from scepter.modules.model.backbone.transformer.patchify import (PatchEmbed,
|
||||
unpatchify)
|
||||
from scepter.modules.model.backbone.transformer.pos_embed import \
|
||||
get_2d_sincos_pos_embed
|
||||
from scepter.modules.model.base_model import BaseModel
|
||||
from scepter.modules.model.registry import BACKBONES
|
||||
from scepter.modules.utils.config import dict_to_yaml
|
||||
from scepter.modules.utils.file_system import FS
|
||||
|
||||
|
||||
def auto_grad_checkpoint(module, *args, use_grad_checkpoint=False, **kwargs):
|
||||
if use_grad_checkpoint:
|
||||
if not isinstance(module, Iterable):
|
||||
return checkpoint(module, *args, use_reentrant=False, **kwargs)
|
||||
gc_step = module[0].grad_checkpointing_step
|
||||
return checkpoint_sequential(module,
|
||||
gc_step,
|
||||
*args,
|
||||
use_reentrant=False,
|
||||
**kwargs)
|
||||
return module(*args, **kwargs)
|
||||
|
||||
|
||||
class FinalLayer(nn.Module):
|
||||
"""
|
||||
The final layer of PixArt.
|
||||
"""
|
||||
def __init__(self, hidden_size, patch_size, out_channels):
|
||||
super().__init__()
|
||||
self.norm_final = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6)
|
||||
self.linear = nn.Linear(hidden_size,
|
||||
patch_size * patch_size * out_channels,
|
||||
bias=True)
|
||||
self.adaLN_modulation = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True))
|
||||
|
||||
def forward(self, x, c):
|
||||
shift, scale = self.adaLN_modulation(c).chunk(2, dim=1)
|
||||
# modulate x
|
||||
x = modulate(self.norm_final(x), shift, scale, unsqueeze=True)
|
||||
x = self.linear(x)
|
||||
return x
|
||||
|
||||
|
||||
class T2IFinalLayer(nn.Module):
|
||||
"""
|
||||
The final layer of PixArt.
|
||||
"""
|
||||
def __init__(self, hidden_size, patch_size, out_channels):
|
||||
super().__init__()
|
||||
self.norm_final = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6)
|
||||
self.linear = nn.Linear(hidden_size,
|
||||
patch_size * patch_size * out_channels,
|
||||
bias=True)
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
torch.randn(2, hidden_size) / hidden_size**0.5)
|
||||
self.out_channels = out_channels
|
||||
|
||||
def forward(self, x, t):
|
||||
shift, scale = (self.scale_shift_table[None] + t[:, None]).chunk(2,
|
||||
dim=1)
|
||||
x = modulate(self.norm_final(x), shift, scale)
|
||||
x = self.linear(x)
|
||||
return x
|
||||
|
||||
|
||||
class DitFinalLayer(nn.Module):
|
||||
"""
|
||||
The final layer of PixArt.
|
||||
"""
|
||||
def __init__(self, hidden_size, patch_size, out_channels):
|
||||
super().__init__()
|
||||
self.norm_final = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6)
|
||||
self.linear = nn.Linear(hidden_size,
|
||||
patch_size * patch_size * out_channels,
|
||||
bias=True)
|
||||
self.adaLN_modulation = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True))
|
||||
self.out_channels = out_channels
|
||||
|
||||
def forward(self, x, t):
|
||||
shift, scale = self.adaLN_modulation(t).chunk(2, dim=1)
|
||||
x = modulate(self.norm_final(x), shift, scale, unsqueeze=True)
|
||||
x = self.linear(x)
|
||||
return x
|
||||
|
||||
|
||||
class PixArtBlock(nn.Module):
|
||||
"""
|
||||
A PixArt block with adaptive layer norm zero (adaLN-Zero) conditioning.
|
||||
"""
|
||||
def __init__(self,
|
||||
hidden_size,
|
||||
num_heads,
|
||||
mlp_ratio=4.0,
|
||||
drop_path=0.,
|
||||
window_size=0,
|
||||
use_rel_pos=False,
|
||||
backend=None,
|
||||
use_condition=True,
|
||||
**block_kwargs):
|
||||
super().__init__()
|
||||
self.hidden_size = hidden_size
|
||||
self.use_condition = use_condition
|
||||
self.norm1 = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6)
|
||||
self.attn = MultiHeadAttention(hidden_size,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=True,
|
||||
backend=backend,
|
||||
**block_kwargs)
|
||||
if self.use_condition:
|
||||
self.cross_attn = MultiHeadAttention(hidden_size,
|
||||
context_dim=hidden_size,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=True,
|
||||
backend=backend,
|
||||
**block_kwargs)
|
||||
self.norm2 = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6)
|
||||
# to be compatible with lower version pytorch
|
||||
approx_gelu = lambda: nn.GELU(approximate='tanh')
|
||||
self.mlp = Mlp(in_features=hidden_size,
|
||||
hidden_features=int(hidden_size * mlp_ratio),
|
||||
act_layer=approx_gelu,
|
||||
drop=0)
|
||||
self.drop_path = DropPath(
|
||||
drop_path) if drop_path > 0. else nn.Identity()
|
||||
self.window_size = window_size
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
torch.randn(6, hidden_size) / hidden_size**0.5)
|
||||
|
||||
def forward(self, x, y, t, mask=None, **kwargs):
|
||||
B, N, C = x.shape
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
|
||||
self.scale_shift_table[None] + t.reshape(B, 6, -1)).chunk(6, dim=1)
|
||||
x = x + self.drop_path(gate_msa * self.attn(
|
||||
modulate(self.norm1(x), shift_msa, scale_msa, unsqueeze=False)))
|
||||
if self.use_condition:
|
||||
x = x + self.cross_attn(x, y, mask)
|
||||
x = x + self.drop_path(gate_mlp * self.mlp(
|
||||
modulate(self.norm2(x), shift_mlp, scale_mlp, unsqueeze=False)))
|
||||
return x
|
||||
|
||||
|
||||
@BACKBONES.register_class()
|
||||
class PixArt(BaseModel):
|
||||
"""
|
||||
Diffusion model with a Transformer backbone.
|
||||
"""
|
||||
para_dict = BaseModel.para_dict
|
||||
para_dict.update({
|
||||
'PATCH_SIZE': {
|
||||
'value': 2,
|
||||
'description': ''
|
||||
},
|
||||
'IN_CHANNELS': {
|
||||
'value': 4,
|
||||
'description': ''
|
||||
},
|
||||
'HIDDEN_SIZE': {
|
||||
'value': 1152,
|
||||
'description': ''
|
||||
},
|
||||
'DEPTH': {
|
||||
'value': 28,
|
||||
'description': ''
|
||||
},
|
||||
'NUM_HEADS': {
|
||||
'value': 16,
|
||||
'description': ''
|
||||
},
|
||||
'MLP_RATIO': {
|
||||
'value': 4.0,
|
||||
'description': ''
|
||||
},
|
||||
'CLASS_DROPOUT_PROB': {
|
||||
'value': 0.1,
|
||||
'description': ''
|
||||
},
|
||||
'PRED_SIGMA': {
|
||||
'value': True,
|
||||
'description': ''
|
||||
},
|
||||
'DROP_PATH': {
|
||||
'value': 0.,
|
||||
'description': ''
|
||||
},
|
||||
'WINDOW_DIZE': {
|
||||
'value': 0,
|
||||
'description': ''
|
||||
},
|
||||
'WINDOW_BLOCK_INDEXES': {
|
||||
'value': None,
|
||||
'description': ''
|
||||
},
|
||||
'USE_REL_POS': {
|
||||
'value': False,
|
||||
'description': ''
|
||||
},
|
||||
'CAPTION_CHANNELS': {
|
||||
'value': 4096,
|
||||
'description': ''
|
||||
},
|
||||
'USE_AR_SIZE': {
|
||||
'value': True,
|
||||
'description': ''
|
||||
},
|
||||
'DIT_FINAL_LAYER': {
|
||||
'value': False,
|
||||
'description': ''
|
||||
},
|
||||
'LEWEI_SCALE': {
|
||||
'value': 1.0,
|
||||
'description': ''
|
||||
},
|
||||
'MODEL_MAX_LENGTH': {
|
||||
'value': 120,
|
||||
'description': ''
|
||||
},
|
||||
'NUM_CLASSES': {
|
||||
'value':
|
||||
None,
|
||||
'description':
|
||||
'The class num for class guided setting, also can be set as continuous.'
|
||||
},
|
||||
'ATTENTION_BACKEND': {
|
||||
'value': None,
|
||||
'description': ''
|
||||
}
|
||||
})
|
||||
|
||||
def __init__(self, cfg, logger):
|
||||
super().__init__(cfg, logger=logger)
|
||||
self.window_block_indexes = cfg.get('WINDOW_BLOCK_INDEXES', None)
|
||||
if self.window_block_indexes is None:
|
||||
self.window_block_indexes = []
|
||||
self.pred_sigma = cfg.get('PRED_SIGMA', True)
|
||||
self.in_channels = cfg.get('IN_CHANNELS', 4)
|
||||
self.out_channels = self.in_channels * 2 if self.pred_sigma else self.in_channels
|
||||
self.patch_size = cfg.get('PATCH_SIZE', 2)
|
||||
self.num_heads = cfg.get('NUM_HEADS', 16)
|
||||
self.hidden_size = cfg.get('HIDDEN_SIZE', 1152)
|
||||
self.lewei_scale = cfg.get('LEWEI_SCALE', 1.0),
|
||||
self.caption_channels = cfg.get('CAPTION_CHANNELS', 4096)
|
||||
self.class_dropout_prob = cfg.get('CLASS_DROPOUT_PROB', 0.1)
|
||||
self.model_max_length = cfg.get('MODEL_MAX_LENGTH', 120)
|
||||
self.drop_path = cfg.get('DROP_PATH', 0.)
|
||||
self.depth = cfg.get('DEPTH', 28)
|
||||
self.mlp_ratio = cfg.get('MLP_RATIO', 4.0)
|
||||
self.num_classes = cfg.get('NUM_CLASSES', None)
|
||||
self.use_grad_checkpoint = cfg.get('USE_GRAD_CHECKPOINT', False)
|
||||
self.use_ar_size = cfg.get('USE_AR_SIZE', True)
|
||||
self.use_dit_final_layer = cfg.get('DIT_FINAL_LAYER', False)
|
||||
self.attention_backend = cfg.get('ATTENTION_BACKEND', None)
|
||||
self.ignore_keys = cfg.get('IGNORE_KEYS', [])
|
||||
|
||||
if self.num_classes is not None:
|
||||
if isinstance(self.num_classes, int):
|
||||
self.label_embedder = LabelEmbedder(
|
||||
self.num_classes,
|
||||
self.hidden_size,
|
||||
dropout_prob=self.class_dropout_prob)
|
||||
elif self.num_classes == 'continuous':
|
||||
print('setting up linear c_adm embedding layer')
|
||||
self.label_embedder = nn.Linear(1, self.hidden_size)
|
||||
else:
|
||||
raise ValueError()
|
||||
|
||||
self.x_embedder = PatchEmbed(self.patch_size,
|
||||
self.in_channels,
|
||||
self.hidden_size,
|
||||
bias=True)
|
||||
self.t_embedder = TimestepEmbedder(self.hidden_size)
|
||||
|
||||
# self.base_size = self.input_size // self.patch_size
|
||||
approx_gelu = lambda: nn.GELU(approximate='tanh')
|
||||
self.t_block = nn.Sequential(
|
||||
nn.SiLU(),
|
||||
nn.Linear(self.hidden_size, 6 * self.hidden_size, bias=True))
|
||||
if self.num_classes is None:
|
||||
self.y_embedder = CaptionEmbedder(
|
||||
in_channels=self.caption_channels,
|
||||
hidden_size=self.hidden_size,
|
||||
uncond_prob=self.class_dropout_prob,
|
||||
act_layer=approx_gelu,
|
||||
token_num=self.model_max_length)
|
||||
if self.use_ar_size:
|
||||
self.csize_embedder = SizeEmbedder(self.hidden_size //
|
||||
3) # c_size embed
|
||||
self.ar_embedder = SizeEmbedder(self.hidden_size //
|
||||
3) # aspect ratio embed
|
||||
|
||||
drop_path = [
|
||||
x.item() for x in torch.linspace(0, self.drop_path, self.depth)
|
||||
] # stochastic depth decay rule
|
||||
self.blocks = nn.ModuleList([
|
||||
PixArtBlock(self.hidden_size,
|
||||
self.num_heads,
|
||||
mlp_ratio=self.mlp_ratio,
|
||||
drop_path=drop_path[i],
|
||||
window_size=self.window_size
|
||||
if i in self.window_block_indexes else 0,
|
||||
use_rel_pos=self.use_rel_pos
|
||||
if i in self.window_block_indexes else False,
|
||||
backend=self.attention_backend,
|
||||
use_condition=self.num_classes is None)
|
||||
for i in range(self.depth)
|
||||
])
|
||||
if self.use_dit_final_layer:
|
||||
self.final_layer = DitFinalLayer(self.hidden_size, self.patch_size,
|
||||
self.out_channels)
|
||||
else:
|
||||
self.final_layer = T2IFinalLayer(self.hidden_size, self.patch_size,
|
||||
self.out_channels)
|
||||
|
||||
self.initialize_weights()
|
||||
|
||||
def load_pretrained_model(self, pretrained_model):
|
||||
if pretrained_model:
|
||||
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
||||
model = torch.load(local_path, map_location='cpu')
|
||||
if 'state_dict' in model:
|
||||
model = model['state_dict']
|
||||
new_ckpt = OrderedDict()
|
||||
for k, v in model.items():
|
||||
k = k.replace('.cross_attn.q_linear.', '.cross_attn.q.')
|
||||
k = k.replace('.cross_attn.proj.',
|
||||
'.cross_attn.o.').replace(
|
||||
'.attn.proj.', '.attn.o.')
|
||||
if '.cross_attn.kv_linear.' in k:
|
||||
k_p, v_p = torch.split(v, v.shape[0] // 2)
|
||||
new_ckpt[k.replace('.cross_attn.kv_linear.',
|
||||
'.cross_attn.k.')] = k_p
|
||||
new_ckpt[k.replace('.cross_attn.kv_linear.',
|
||||
'.cross_attn.v.')] = v_p
|
||||
elif '.attn.qkv.' in k:
|
||||
q_p, k_p, v_p = torch.split(v, v.shape[0] // 3)
|
||||
new_ckpt[k.replace('.attn.qkv.', '.attn.q.')] = q_p
|
||||
new_ckpt[k.replace('.attn.qkv.', '.attn.k.')] = k_p
|
||||
new_ckpt[k.replace('.attn.qkv.', '.attn.v.')] = v_p
|
||||
else:
|
||||
new_ckpt[k] = v
|
||||
missing, unexpected = self.load_state_dict(new_ckpt,
|
||||
strict=False)
|
||||
print(
|
||||
f'Restored from {pretrained_model} with {len(missing)} missing and {len(unexpected)} unexpected keys'
|
||||
)
|
||||
if len(missing) > 0:
|
||||
print(f'Missing Keys:\n {missing}')
|
||||
if len(unexpected) > 0:
|
||||
print(f'\nUnexpected Keys:\n {unexpected}')
|
||||
|
||||
def forward(self,
|
||||
x,
|
||||
t=None,
|
||||
cond=dict(),
|
||||
mask=None,
|
||||
data_info=None,
|
||||
**kwargs):
|
||||
"""
|
||||
Forward pass of PixArt.
|
||||
x: (N, C, H, W) tensor of spatial inputs (images or latent representations of images)
|
||||
t: (N,) tensor of diffusion timesteps
|
||||
y: (N, 1, 120, C) tensor of class labels
|
||||
"""
|
||||
label = None
|
||||
if isinstance(cond, dict):
|
||||
if 'label' in cond and cond['label'] is not None:
|
||||
label = cond['label']
|
||||
if 'concat' in cond:
|
||||
concat = cond['concat']
|
||||
x = torch.cat([x, concat], dim=1)
|
||||
context = cond.get('crossattn', None)
|
||||
else:
|
||||
context = cond
|
||||
|
||||
y = context
|
||||
h, w = x.shape[-2] // self.patch_size, x.shape[-1] // self.patch_size
|
||||
|
||||
x = self.x_embedder(x) # (N, T, D), where T = H * W / patch_size ** 2
|
||||
pos_embed = torch.from_numpy(
|
||||
get_2d_sincos_pos_embed(self.hidden_size, (h, w),
|
||||
lewei_scale=self.lewei_scale,
|
||||
base_h_size=h,
|
||||
base_w_size=w)).unsqueeze(0).float().to(
|
||||
x.device)
|
||||
x = x + pos_embed
|
||||
|
||||
t = self.t_embedder(t) # (N, D)
|
||||
if self.num_classes is not None and label is not None:
|
||||
t = t + self.label_embedder(label, self.training)
|
||||
if self.use_ar_size and data_info is not None:
|
||||
bs = x.shape[0]
|
||||
c_size, ar = data_info['img_hw'], data_info['aspect_ratio']
|
||||
csize = self.csize_embedder(c_size, bs) # (N, D)
|
||||
ar = self.ar_embedder(ar, bs) # (N, D)
|
||||
t = t + torch.cat([csize, ar], dim=1)
|
||||
t0 = self.t_block(t)
|
||||
if self.num_classes is not None:
|
||||
y = None
|
||||
else:
|
||||
y = self.y_embedder(y, self.training)
|
||||
for block in self.blocks:
|
||||
x = auto_grad_checkpoint(
|
||||
block,
|
||||
x,
|
||||
y,
|
||||
t0,
|
||||
mask,
|
||||
use_grad_checkpoint=self.use_grad_checkpoint)
|
||||
# (N, T, D) #support grad checkpoint
|
||||
x = self.final_layer(x, t) # (N, T, patch_size ** 2 * out_channels)
|
||||
x = unpatchify(x, h, w, self.out_channels, self.patch_size,
|
||||
self.patch_size) # (N, out_channels, H, W)
|
||||
if self.pred_sigma:
|
||||
return x.chunk(2, dim=1)[0]
|
||||
else:
|
||||
return x
|
||||
|
||||
def initialize_weights(self):
|
||||
# Initialize transformer layers:
|
||||
def _basic_init(module):
|
||||
if isinstance(module, nn.Linear):
|
||||
torch.nn.init.xavier_uniform_(module.weight)
|
||||
if module.bias is not None:
|
||||
nn.init.constant_(module.bias, 0)
|
||||
|
||||
self.apply(_basic_init)
|
||||
|
||||
# Initialize patch_embed like nn.Linear (instead of nn.Conv2d):
|
||||
w = self.x_embedder.proj.weight.data
|
||||
nn.init.xavier_uniform_(w.view([w.shape[0], -1]))
|
||||
# Initialize timestep embedding MLP:
|
||||
nn.init.normal_(self.t_embedder.mlp[0].weight, std=0.02)
|
||||
nn.init.normal_(self.t_embedder.mlp[2].weight, std=0.02)
|
||||
nn.init.normal_(self.t_block[1].weight, std=0.02)
|
||||
if self.use_ar_size:
|
||||
nn.init.normal_(self.csize_embedder.mlp[0].weight, std=0.02)
|
||||
nn.init.normal_(self.csize_embedder.mlp[2].weight, std=0.02)
|
||||
nn.init.normal_(self.ar_embedder.mlp[0].weight, std=0.02)
|
||||
nn.init.normal_(self.ar_embedder.mlp[2].weight, std=0.02)
|
||||
if self.num_classes is not None:
|
||||
nn.init.normal_(self.label_embedder.embedding_table.weight,
|
||||
std=0.02)
|
||||
# Initialize caption embedding MLP:
|
||||
if hasattr(self, 'y_embedder'):
|
||||
nn.init.normal_(self.y_embedder.y_proj.fc1.weight, std=0.02)
|
||||
nn.init.normal_(self.y_embedder.y_proj.fc2.weight, std=0.02)
|
||||
# Zero-out adaLN modulation layers in PixArt blocks:
|
||||
if self.num_classes is None:
|
||||
for block in self.blocks:
|
||||
nn.init.constant_(block.cross_attn.o.weight, 0)
|
||||
nn.init.constant_(block.cross_attn.o.bias, 0)
|
||||
# Zero-out output layers:
|
||||
nn.init.constant_(self.final_layer.linear.weight, 0)
|
||||
nn.init.constant_(self.final_layer.linear.bias, 0)
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
return next(self.parameters()).dtype
|
||||
|
||||
@staticmethod
|
||||
def get_config_template():
|
||||
return dict_to_yaml('BACKBONE',
|
||||
__class__.__name__,
|
||||
PixArt.para_dict,
|
||||
set_name=True)
|
||||
@@ -0,0 +1,760 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
# All rights reserved.
|
||||
# This file contains code that is adapted from
|
||||
# timm: https://github.com/huggingface/pytorch-image-models
|
||||
# pixart: https://github.com/PixArt-alpha/PixArt-alpha
|
||||
import math
|
||||
import time
|
||||
import warnings
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.cuda import amp
|
||||
from torch.nn import functional as F
|
||||
from torch.nn.utils.rnn import pad_sequence
|
||||
from tqdm import tqdm
|
||||
|
||||
from scepter.modules.model.backbone.transformer.pos_embed import apply_2d_rope
|
||||
|
||||
try:
|
||||
import xformers
|
||||
import xformers.ops
|
||||
XFORMERS_IS_AVAILABLE = True
|
||||
except Exception as e:
|
||||
XFORMERS_IS_AVAILABLE = False
|
||||
warnings.warn(f'{e}')
|
||||
try:
|
||||
from flash_attn import (flash_attn_varlen_func)
|
||||
FLASHATTN_IS_AVAILABLE = True
|
||||
except ImportError:
|
||||
FLASHATTN_IS_AVAILABLE = False
|
||||
flash_attn_varlen_func = None
|
||||
|
||||
|
||||
def drop_path(x, drop_prob: float = 0., training: bool = False):
|
||||
"""Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
|
||||
This is the same as the DropConnect impl I created for EfficientNet, etc networks, however,
|
||||
the original name is misleading as 'Drop Connect' is a different form of dropout in a separate paper...
|
||||
See discussion: https://github.com/tensorflow/tpu/issues/494#issuecomment-532968956 ... I've opted for
|
||||
changing the layer and argument names to 'drop path' rather than mix DropConnect as a layer name and use
|
||||
'survival rate' as the argument.
|
||||
"""
|
||||
if drop_prob == 0. or not training:
|
||||
return x
|
||||
keep_prob = 1 - drop_prob
|
||||
shape = (x.shape[0], ) + (1, ) * (
|
||||
x.ndim - 1) # work with diff dim tensors, not just 2D ConvNets
|
||||
random_tensor = keep_prob + torch.rand(
|
||||
shape, dtype=x.dtype, device=x.device)
|
||||
random_tensor.floor_() # binarize
|
||||
output = x.div(keep_prob) * random_tensor
|
||||
return output
|
||||
|
||||
|
||||
class MultiHeadAttention(nn.Module):
|
||||
def __init__(self,
|
||||
dim,
|
||||
context_dim=None,
|
||||
num_heads=None,
|
||||
head_dim=None,
|
||||
attn_drop=0.0,
|
||||
qkv_bias=False,
|
||||
dropout=0.0,
|
||||
backend=None,
|
||||
**block_kwargs):
|
||||
super().__init__()
|
||||
# consider head_dim first, then num_heads
|
||||
num_heads = dim // head_dim if head_dim else num_heads
|
||||
head_dim = dim // num_heads
|
||||
assert num_heads * head_dim == dim
|
||||
context_dim = context_dim or dim
|
||||
self.dim = dim
|
||||
self.context_dim = context_dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = head_dim
|
||||
self.scale = math.pow(head_dim, -0.25)
|
||||
# layers
|
||||
self.q = nn.Linear(dim, dim, bias=qkv_bias)
|
||||
self.k = nn.Linear(context_dim, dim, bias=qkv_bias)
|
||||
self.v = nn.Linear(context_dim, dim, bias=qkv_bias)
|
||||
self.o = nn.Linear(dim, dim)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.attention_op = None
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
self.backend = backend
|
||||
assert self.backend in ('flash_attn', 'xformer_attn', 'pytorch_attn',
|
||||
None)
|
||||
if FLASHATTN_IS_AVAILABLE and self.backend in ('flash_attn', None):
|
||||
self.backend = 'flash_attn'
|
||||
self.softmax_scale = block_kwargs.get('softmax_scale', None)
|
||||
self.causal = block_kwargs.get('causal', False)
|
||||
self.window_size = block_kwargs.get('window_size', (-1, -1))
|
||||
self.deterministic = block_kwargs.get('deterministic', False)
|
||||
elif XFORMERS_IS_AVAILABLE and self.backend in ('xformer_attn', None):
|
||||
self.backend = 'xformer_attn'
|
||||
else:
|
||||
self.backend = 'pytorch_attn'
|
||||
|
||||
def xformer_attn(self, x, context=None, mask=None, **kwargs):
|
||||
context = x if context is None else context
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
# compute query, key, value
|
||||
q = self.q(x).view(b, -1, n, d)
|
||||
k = self.k(context).view(b, -1, n, d)
|
||||
v = self.v(context).view(b, -1, n, d)
|
||||
|
||||
attn_bias = None
|
||||
if mask is not None:
|
||||
assert mask.ndim in [2, 3]
|
||||
mask = mask.view(b, 1, 1,
|
||||
-1) if mask.ndim == 2 else mask.unsqueeze(1)
|
||||
# To use an `attn_bias` with a sequence length that is not a multiple of 8,
|
||||
# you need to ensure memory is aligned by slicing a bigger tensor.
|
||||
# Example: use `attn_bias = torch.zeros([1, 1, 5, 8])[:,:,:,:5]`
|
||||
# instead of `torch.zeros([1, 1, 5, 5])
|
||||
q_size = math.ceil(q.size(1) / 8) * 8
|
||||
k_size = math.ceil(k.size(1) / 8) * 8
|
||||
attn_bias = x.new_zeros(b, n, q_size,
|
||||
k_size)[:, :, :q.size(1), :k.size(1)]
|
||||
attn_bias = attn_bias.masked_fill_(mask == 0,
|
||||
torch.finfo(x.dtype).min).to(
|
||||
q.dtype)
|
||||
x = xformers.ops.memory_efficient_attention(q,
|
||||
k,
|
||||
v,
|
||||
p=self.attn_drop.p,
|
||||
attn_bias=attn_bias)
|
||||
x = x.reshape(b, -1, n * d)
|
||||
x = self.o(x)
|
||||
x = self.dropout(x)
|
||||
return x
|
||||
|
||||
def flash_attn(self, x, context=None, mask=None, **kwargs):
|
||||
'''
|
||||
The implementation will be very slow when mask is not None,
|
||||
because we need rearange the x/context features according to mask.
|
||||
Args:
|
||||
x:
|
||||
context:
|
||||
mask:
|
||||
**kwargs:
|
||||
Returns: x
|
||||
'''
|
||||
context = x if context is None else context
|
||||
dtype = kwargs.get('dtype', torch.float16)
|
||||
q_lens = kwargs.get('q_lens', None)
|
||||
|
||||
# if mask is not None or q_lens is not None:
|
||||
# warnings.warn("Detected mask or q_lens is not None, "
|
||||
# "which will be very slow because of the x/context features' rearrangement,"
|
||||
# "please use FlashMultiHeadAttention instead.")
|
||||
def half(x):
|
||||
return x if x.dtype in [torch.float16, torch.bfloat16
|
||||
] else x.to(dtype)
|
||||
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
q = self.q(x).view(b, -1, n, d) # [B, Lq, Nq, C1].
|
||||
k = self.k(context).view(b, -1, n, d) # [B, Lk, Nk, C1]
|
||||
v = self.v(context).view(
|
||||
b, -1, n, d) # [B, Lk, Nk, C2] Nq must be divisible by Nk.
|
||||
|
||||
assert q.device.type == 'cuda' and q.size(-1) <= 256
|
||||
lq, lk, out_dtype = int(q.size(1)), int(k.size(1)), q.dtype
|
||||
# preprocess query
|
||||
if q_lens is None:
|
||||
q_lens = torch.tensor([lq] * b,
|
||||
dtype=torch.int32).to(q.device,
|
||||
non_blocking=True)
|
||||
# q_lens = (q.flatten(2, ).bool() + 1).sum(dim=-1).bool().sum(dim=-1)
|
||||
q = half(q.flatten(0, 1))
|
||||
else:
|
||||
q = half(torch.cat([q_v[:q_l] for q_v, q_l in zip(q, q_lens)]))
|
||||
|
||||
# preprocess key, value
|
||||
if mask is None:
|
||||
k_lens = torch.tensor([lk] * b,
|
||||
dtype=torch.int32).to(k.device,
|
||||
non_blocking=True)
|
||||
# k_lens = (k.flatten(2, ).bool() + 1).sum(dim=-1).bool().sum(dim=-1)
|
||||
k = half(k.flatten(0, 1))
|
||||
v = half(v.flatten(0, 1))
|
||||
else:
|
||||
assert mask.ndim in [1, 2, 3]
|
||||
k_lens = mask if mask.ndim == 1 else mask.flatten(start_dim=1).sum(
|
||||
dim=-1)
|
||||
k = half(torch.cat([k_v[:k_l] for k_v, k_l in zip(k, k_lens)]))
|
||||
v = half(torch.cat([v_v[:v_l] for v_v, v_l in zip(v, k_lens)]))
|
||||
|
||||
x = flash_attn_varlen_func(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
cu_seqlens_q=torch.cat([q_lens.new_zeros([1]),
|
||||
q_lens]).cumsum(0, dtype=torch.int32),
|
||||
cu_seqlens_k=torch.cat([k_lens.new_zeros([1]),
|
||||
k_lens]).cumsum(0, dtype=torch.int32),
|
||||
max_seqlen_q=int(torch.max(q_lens).cpu().numpy()),
|
||||
max_seqlen_k=int(torch.max(k_lens).cpu().numpy()),
|
||||
dropout_p=self.attn_drop.p,
|
||||
softmax_scale=self.softmax_scale,
|
||||
causal=self.causal,
|
||||
window_size=self.window_size, # -1 means infinite context window
|
||||
deterministic=self.deterministic).unflatten(0, (b, lq))
|
||||
x = x.type(out_dtype)
|
||||
x = x.flatten(2)
|
||||
# output
|
||||
x = self.o(x)
|
||||
x = self.dropout(x)
|
||||
return x
|
||||
|
||||
def pytorch_attn(self, x, context=None, mask=None, **kwargs):
|
||||
"""x: [B, L, C].
|
||||
context: [B, L', C'] or None.
|
||||
"""
|
||||
context = x if context is None else context
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.q(x).view(b, -1, n, d)
|
||||
k = self.k(context).view(b, -1, n, d)
|
||||
v = self.v(context).view(b, -1, n, d)
|
||||
# attention bias
|
||||
attn_bias = x.new_zeros(b, n, q.size(1), k.size(1))
|
||||
if mask is not None:
|
||||
assert mask.ndim in [2, 3]
|
||||
mask = mask.view(b, 1, 1,
|
||||
-1) if mask.ndim == 2 else mask.unsqueeze(1)
|
||||
attn_bias = attn_bias.masked_fill_(mask == 0,
|
||||
torch.finfo(x.dtype).min).to(
|
||||
q.dtype)
|
||||
|
||||
# compute attention (T5 does not use scaling)
|
||||
attn = torch.einsum('binc,bjnc->bnij', q * self.scale,
|
||||
k * self.scale) + attn_bias
|
||||
attn = F.softmax(attn.float(), dim=-1).type_as(attn)
|
||||
x = torch.einsum('bnij,bjnc->binc', attn, v.float())
|
||||
# output
|
||||
x = x.reshape(b, -1, n * d)
|
||||
x = self.o(x)
|
||||
x = self.dropout(x)
|
||||
return x
|
||||
|
||||
def forward(self, x, context=None, mask=None, **kwargs):
|
||||
"""x: [B, L, C].
|
||||
context: [B, L', C'] or None.
|
||||
"""
|
||||
x = getattr(self, self.backend)(x,
|
||||
context=context,
|
||||
mask=mask,
|
||||
**kwargs)
|
||||
return x
|
||||
|
||||
|
||||
def flash_preprocess(x, context=None, q_mask=None, mask=None):
|
||||
context = x if context is None else context
|
||||
b, x_l, x_hidden_size = x.shape
|
||||
x = x.flatten(0, 1)
|
||||
if q_mask is None:
|
||||
q_lens = torch.tensor([x_l] * b,
|
||||
dtype=torch.int32).to(x.device,
|
||||
non_blocking=True)
|
||||
else:
|
||||
assert q_mask.ndim in [1, 2, 3]
|
||||
q_lens = q_mask if q_mask.ndim == 1 else q_mask.flatten(
|
||||
start_dim=1).sum(dim=-1)
|
||||
|
||||
mask_b, mask_l, mask_hidden_size = context.shape
|
||||
|
||||
if mask is None:
|
||||
mask_lens = torch.tensor([mask_l] * mask_b,
|
||||
dtype=torch.int32).to(context.device,
|
||||
non_blocking=True)
|
||||
else:
|
||||
assert mask.ndim in [1, 2, 3]
|
||||
mask_lens = mask if mask.ndim == 1 else mask.flatten(start_dim=1).sum(
|
||||
dim=-1)
|
||||
|
||||
return_data = {
|
||||
'x':
|
||||
x,
|
||||
'context':
|
||||
torch.cat([u[:v] for u, v in zip(context, mask_lens)]),
|
||||
'cu_seqlens_q':
|
||||
torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(0,
|
||||
dtype=torch.int32),
|
||||
'max_seqlen_q':
|
||||
int(torch.max(q_lens).cpu().numpy()),
|
||||
'cu_seqlens_k':
|
||||
torch.cat([mask_lens.new_zeros([1]),
|
||||
mask_lens]).cumsum(0, dtype=torch.int32),
|
||||
'max_seqlen_k':
|
||||
int(torch.max(mask_lens).cpu().numpy())
|
||||
}
|
||||
return return_data
|
||||
|
||||
|
||||
class FlashMultiHeadAttention(nn.Module):
|
||||
def __init__(self,
|
||||
dim,
|
||||
context_dim=None,
|
||||
num_heads=None,
|
||||
head_dim=None,
|
||||
attn_drop=0.0,
|
||||
qkv_bias=False,
|
||||
dropout=0.0,
|
||||
softmax_scale=None,
|
||||
causal=False,
|
||||
window_size=(-1, -1),
|
||||
deterministic=False,
|
||||
**block_kwargs):
|
||||
super().__init__()
|
||||
# consider head_dim first, then num_heads
|
||||
num_heads = dim // head_dim if head_dim else num_heads
|
||||
head_dim = dim // num_heads
|
||||
assert num_heads * head_dim == dim
|
||||
context_dim = context_dim or dim
|
||||
self.dim = dim
|
||||
self.context_dim = context_dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = head_dim
|
||||
self.scale = math.pow(head_dim, -0.25)
|
||||
# layers
|
||||
self.q = nn.Linear(dim, dim, bias=qkv_bias)
|
||||
self.k = nn.Linear(context_dim, dim, bias=qkv_bias)
|
||||
self.v = nn.Linear(context_dim, dim, bias=qkv_bias)
|
||||
self.o = nn.Linear(dim, dim)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.attention_op = None
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
self.softmax_scale = softmax_scale
|
||||
self.causal = causal
|
||||
self.window_size = window_size
|
||||
self.deterministic = deterministic
|
||||
|
||||
def forward(self,
|
||||
x,
|
||||
context=None,
|
||||
cu_seqlens_q=None,
|
||||
max_seqlen_q=None,
|
||||
cu_seqlens_k=None,
|
||||
max_seqlen_k=None,
|
||||
dtype=torch.float16,
|
||||
**kwargs):
|
||||
'''
|
||||
The implementation used the rearanaged x/context according to q_lens or k_lens.
|
||||
Args:
|
||||
x: [batch_size * max_seq_len or sum(q_lens) , heads, hidden_size].
|
||||
context: [batch_size * max_seq_len or sum(q_lens) , heads, hidden_size].
|
||||
cu_seqlens_q: cumsum of seq_q to index the postion of query in the batch.
|
||||
max_seqlen_q: max length of query.
|
||||
cu_seqlens_k: cumsum of seq_k to index the postion of key/value in the batch.
|
||||
max_seqlen_k: max length of key/value.
|
||||
dtype: the dtype for attention.
|
||||
**kwargs:
|
||||
Returns: x
|
||||
'''
|
||||
context = x if context is None else context
|
||||
|
||||
def half(x):
|
||||
return x if x.dtype in [torch.float16, torch.bfloat16
|
||||
] else x.to(dtype)
|
||||
|
||||
n, d, out_dtype = self.num_heads, self.head_dim, x.dtype
|
||||
q = self.q(x).view(-1, n, d) # [B * Lq, Nq, C1].
|
||||
k = self.k(context).view(-1, n, d) # [B * Lk, Nk, C1]
|
||||
v = self.v(context).view(
|
||||
-1, n, d) # [B * Lk, Nk, C2] Nq must be divisible by Nk.
|
||||
q, k, v = half(q), half(k), half(v)
|
||||
assert q.device.type == 'cuda' and d <= 256
|
||||
x = flash_attn_varlen_func(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
cu_seqlens_q=cu_seqlens_q,
|
||||
cu_seqlens_k=cu_seqlens_k,
|
||||
max_seqlen_q=max_seqlen_q,
|
||||
max_seqlen_k=max_seqlen_k,
|
||||
dropout_p=self.attn_drop.p,
|
||||
softmax_scale=self.softmax_scale,
|
||||
causal=self.causal,
|
||||
window_size=self.window_size, # -1 means infinite context window
|
||||
deterministic=self.deterministic).unflatten(0, (x.shape[0], ))
|
||||
x = x.flatten(1).type(out_dtype)
|
||||
# output
|
||||
x = self.o(x)
|
||||
x = self.dropout(x)
|
||||
return x
|
||||
|
||||
|
||||
def multi_head_varlen_attention(q_img,
|
||||
k_img,
|
||||
v_img,
|
||||
q_txt,
|
||||
k_txt,
|
||||
v_txt,
|
||||
n,
|
||||
d,
|
||||
img_lens,
|
||||
txt_lens,
|
||||
dropout_p=0.0,
|
||||
flash_dtype=torch.bfloat16):
|
||||
'''
|
||||
q/k/v: b, s, n*d
|
||||
q_lens/k_lens: b,
|
||||
'''
|
||||
from flash_attn import flash_attn_varlen_func
|
||||
q_lens = k_lens = img_lens + txt_lens
|
||||
|
||||
cu_seqlens_q = torch.cat([q_lens.new_zeros([1]),
|
||||
q_lens]).cumsum(0, dtype=torch.int32)
|
||||
cu_seqlens_k = torch.cat([k_lens.new_zeros([1]),
|
||||
k_lens]).cumsum(0, dtype=torch.int32)
|
||||
max_seqlen_q = q_lens.max()
|
||||
max_seqlen_k = k_lens.max()
|
||||
|
||||
# concat img & txt for joint attention
|
||||
q = torch.cat([
|
||||
torch.cat([i[:i_len], t[:t_len]], dim=0)
|
||||
for i, i_len, t, t_len in zip(q_img, img_lens, q_txt, txt_lens)
|
||||
],
|
||||
dim=0).view(-1, n, d)
|
||||
|
||||
k = torch.cat([
|
||||
torch.cat([i[:i_len], t[:t_len]], dim=0)
|
||||
for i, i_len, t, t_len in zip(k_img, img_lens, k_txt, txt_lens)
|
||||
],
|
||||
dim=0).view(-1, n, d)
|
||||
|
||||
v = torch.cat([
|
||||
torch.cat([i[:i_len], t[:t_len]], dim=0)
|
||||
for i, i_len, t, t_len in zip(v_img, img_lens, v_txt, txt_lens)
|
||||
],
|
||||
dim=0).view(-1, n, d)
|
||||
|
||||
# attention
|
||||
dtype = q.dtype
|
||||
if dtype != flash_dtype:
|
||||
q = q.type(flash_dtype)
|
||||
k = k.type(flash_dtype)
|
||||
v = v.type(flash_dtype)
|
||||
|
||||
with amp.autocast():
|
||||
x = flash_attn_varlen_func(q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens_q=cu_seqlens_q,
|
||||
cu_seqlens_k=cu_seqlens_k,
|
||||
max_seqlen_q=max_seqlen_q,
|
||||
max_seqlen_k=max_seqlen_k,
|
||||
dropout_p=dropout_p).type(dtype)
|
||||
|
||||
return x, cu_seqlens_q
|
||||
|
||||
|
||||
class RMSNorm(nn.Module):
|
||||
def __init__(self, dim, eps=1e-6):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.eps = eps
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def forward(self, x):
|
||||
return self._norm(x.float()).type_as(x) * self.weight
|
||||
|
||||
def _norm(self, x):
|
||||
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
|
||||
|
||||
|
||||
class FullAttention(nn.Module):
|
||||
def __init__(self,
|
||||
dim,
|
||||
num_heads=None,
|
||||
head_dim=None,
|
||||
dropout=0.0,
|
||||
qkv_bias=False,
|
||||
qk_norm=False,
|
||||
eps=1e-6,
|
||||
flash_dtype=torch.bfloat16):
|
||||
|
||||
# consider head_dim first, then num_heads
|
||||
num_heads = dim // head_dim if head_dim else num_heads
|
||||
head_dim = dim // num_heads
|
||||
assert num_heads * head_dim == dim
|
||||
assert flash_dtype in (None, torch.float16, torch.bfloat16)
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = head_dim
|
||||
self.scale = math.pow(head_dim, -0.25)
|
||||
self.flash_dtype = flash_dtype
|
||||
# layers
|
||||
self.qkv_W = nn.Linear(dim, dim * 3, bias=qkv_bias)
|
||||
|
||||
self.out_proj = nn.Linear(dim, dim)
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
|
||||
if qk_norm:
|
||||
from apex.normalization import FusedRMSNorm
|
||||
self.q_img_norm = FusedRMSNorm(head_dim, eps=eps)
|
||||
self.k_img_norm = FusedRMSNorm(head_dim, eps=eps)
|
||||
self.q_txt_norm = FusedRMSNorm(head_dim, eps=eps)
|
||||
self.k_txt_norm = FusedRMSNorm(head_dim, eps=eps)
|
||||
else:
|
||||
self.q_img_norm = nn.Identity()
|
||||
self.k_img_norm = nn.Identity()
|
||||
self.q_txt_norm = nn.Identity()
|
||||
self.k_txt_norm = nn.Identity()
|
||||
|
||||
def forward(self,
|
||||
img,
|
||||
txt,
|
||||
img_lens=None,
|
||||
txt_lens=None,
|
||||
padded_pos_index=None):
|
||||
'''
|
||||
img: B, L, C
|
||||
txt: B, L', C
|
||||
'''
|
||||
b, img_len, c = img.shape
|
||||
txt_len, n, d = txt.shape[1], self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
img_txt = torch.cat([img, txt], dim=1)
|
||||
img_tokens, txt_tokens = torch.split(self.qkv_W(img_txt),
|
||||
[img_len, txt_len],
|
||||
dim=1)
|
||||
|
||||
q_img, k_img, v_img = img_tokens.chunk(3, dim=-1)
|
||||
q_txt, k_txt, v_txt = txt_tokens.chunk(3, dim=-1)
|
||||
|
||||
# multi-head qk norm
|
||||
q_img, k_img = q_img.view(b, -1, n, d), k_img.view(b, -1, n, d)
|
||||
q_txt, k_txt = q_txt.view(b, -1, n, d), k_txt.view(b, -1, n, d)
|
||||
q_img, q_txt = self.q_img_norm(q_img).view(
|
||||
b, -1, n * d), self.q_txt_norm(q_txt).view(b, -1, n * d)
|
||||
k_img, k_txt = self.k_img_norm(k_img).view(
|
||||
b, -1, n * d), self.k_txt_norm(k_txt).view(b, -1, n * d)
|
||||
|
||||
### add position
|
||||
q_img, k_img = apply_2d_rope(q_img, k_img, padded_pos_index, n, d)
|
||||
|
||||
# support varying length
|
||||
if img_lens is None:
|
||||
img_lens = torch.tensor([img.size(1)] * b,
|
||||
dtype=torch.int32,
|
||||
device=img.device)
|
||||
if txt_lens is None:
|
||||
txt_lens = torch.tensor([txt.size(1)] * b,
|
||||
dtype=torch.int32,
|
||||
device=txt.device)
|
||||
|
||||
# attention
|
||||
x, cu_seqlens_q = multi_head_varlen_attention(
|
||||
q_img,
|
||||
k_img,
|
||||
v_img,
|
||||
q_txt,
|
||||
k_txt,
|
||||
v_txt,
|
||||
n,
|
||||
d,
|
||||
img_lens,
|
||||
txt_lens,
|
||||
dropout_p=self.dropout.p if self.training else 0.0,
|
||||
flash_dtype=self.flash_dtype)
|
||||
|
||||
# output proj.
|
||||
x = x.reshape(-1, n * d)
|
||||
x = self.out_proj(x)
|
||||
x = self.dropout(x)
|
||||
|
||||
# split img & txt and padding to max_len
|
||||
img = pad_sequence(tuple([
|
||||
x[s:s + img_len] for s, e, img_len in zip(
|
||||
cu_seqlens_q[:-1], cu_seqlens_q[1:], img_lens)
|
||||
]),
|
||||
batch_first=True)
|
||||
txt = pad_sequence(tuple([
|
||||
x[s + img_len:e] for s, e, img_len in zip(
|
||||
cu_seqlens_q[:-1], cu_seqlens_q[1:], img_lens)
|
||||
]),
|
||||
batch_first=True)
|
||||
|
||||
return img, txt
|
||||
|
||||
|
||||
class FFNSwiGLU(nn.Module):
|
||||
def __init__(self, in_features, hidden_features):
|
||||
super().__init__()
|
||||
self.W1 = nn.Linear(in_features, hidden_features, bias=False)
|
||||
self.W2 = nn.Linear(in_features, hidden_features, bias=False)
|
||||
self.W3 = nn.Linear(hidden_features, in_features, bias=False)
|
||||
self.silu = nn.SiLU()
|
||||
|
||||
def forward(self, x):
|
||||
return self.W3(self.silu(self.W1(x)) * self.W2(x))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
# Align results for different attention implementation
|
||||
torch.manual_seed(2023)
|
||||
hidden_dim = 4096
|
||||
q_weight = torch.randn((hidden_dim, hidden_dim))
|
||||
q_bias = torch.zeros((hidden_dim))
|
||||
k_weight = torch.randn((hidden_dim, hidden_dim))
|
||||
k_bias = torch.zeros((hidden_dim))
|
||||
v_weight = torch.randn((hidden_dim, hidden_dim))
|
||||
v_bias = torch.zeros((hidden_dim))
|
||||
o_weight = torch.randn((hidden_dim, hidden_dim))
|
||||
o_bias = torch.randn((hidden_dim))
|
||||
pytorch_attn = MultiHeadAttention(hidden_dim,
|
||||
context_dim=hidden_dim,
|
||||
num_heads=32,
|
||||
head_dim=None,
|
||||
attn_drop=0.0,
|
||||
dropout=0.0,
|
||||
backend='pytorch_attn')
|
||||
|
||||
pytorch_attn.load_state_dict({
|
||||
'q.weight': q_weight,
|
||||
'k.weight': k_weight,
|
||||
'v.weight': v_weight,
|
||||
'o.weight': o_weight,
|
||||
'o.bias': o_bias
|
||||
})
|
||||
pytorch_attn.to(0)
|
||||
|
||||
xformer_attn = MultiHeadAttention(hidden_dim,
|
||||
context_dim=hidden_dim,
|
||||
num_heads=32,
|
||||
head_dim=None,
|
||||
attn_drop=0.0,
|
||||
dropout=0.0,
|
||||
backend='xformer_attn')
|
||||
|
||||
xformer_attn.load_state_dict({
|
||||
'q.weight': q_weight,
|
||||
'k.weight': k_weight,
|
||||
'v.weight': v_weight,
|
||||
'o.weight': o_weight,
|
||||
'o.bias': o_bias
|
||||
})
|
||||
xformer_attn.to(0)
|
||||
|
||||
flash_attn = MultiHeadAttention(hidden_dim,
|
||||
context_dim=hidden_dim,
|
||||
num_heads=32,
|
||||
head_dim=None,
|
||||
attn_drop=0.0,
|
||||
dropout=0.0,
|
||||
backend='flash_attn',
|
||||
dtype=torch.float16)
|
||||
flash_attn.load_state_dict({
|
||||
'q.weight': q_weight,
|
||||
'k.weight': k_weight,
|
||||
'v.weight': v_weight,
|
||||
'o.weight': o_weight,
|
||||
'o.bias': o_bias
|
||||
})
|
||||
flash_attn.to(0)
|
||||
|
||||
improved_flash_attn = FlashMultiHeadAttention(hidden_dim,
|
||||
context_dim=hidden_dim,
|
||||
num_heads=32,
|
||||
head_dim=None,
|
||||
attn_drop=0.0,
|
||||
dropout=0.0,
|
||||
backend='flash_attn',
|
||||
dtype=torch.float16)
|
||||
|
||||
improved_flash_attn.load_state_dict({
|
||||
'q.weight': q_weight,
|
||||
'k.weight': k_weight,
|
||||
'v.weight': v_weight,
|
||||
'o.weight': o_weight,
|
||||
'o.bias': o_bias
|
||||
})
|
||||
improved_flash_attn.to(0)
|
||||
|
||||
batch_size = 1
|
||||
query_length = 1024
|
||||
key_length = 1024
|
||||
# mask = None
|
||||
run_num = 10
|
||||
torch.cuda.empty_cache()
|
||||
x = torch.randn((batch_size, query_length, hidden_dim)).to(0)
|
||||
context = torch.randn((batch_size, key_length, hidden_dim)).to(0)
|
||||
# mask = torch.cat([torch.ones((batch_size, 80)), torch.zeros((batch_size, key_length - 80))], dim=1).long().to(0)
|
||||
# mask = torch.randint(1, key_length, [batch_size]).to(0)
|
||||
mask = None
|
||||
st = time.time()
|
||||
for i in tqdm(range(run_num)):
|
||||
pytorch_res = pytorch_attn(x.clone(), context.clone(),
|
||||
mask.clone() if mask is not None else mask)
|
||||
if i == run_num - 1:
|
||||
free_mem, total_mem = torch.cuda.mem_get_info(0)
|
||||
free_mem, total_mem = free_mem / (1024**3), total_mem / (1024**3)
|
||||
mem_msg = f'GPU {0}: free mem {free_mem:.3f}G, total mem {total_mem:.3f}G \n'
|
||||
pytorch_res_data = pytorch_res.clone().detach().cpu()
|
||||
print('pytorch attn ', mem_msg,
|
||||
f'Cost time per time {(time.time() - st) / run_num}s')
|
||||
#
|
||||
torch.cuda.empty_cache()
|
||||
st = time.time()
|
||||
for i in tqdm(range(run_num)):
|
||||
xformer_res = xformer_attn(x.clone(), context.clone(),
|
||||
mask.clone() if mask is not None else mask)
|
||||
if i == run_num - 1:
|
||||
free_mem, total_mem = torch.cuda.mem_get_info(0)
|
||||
free_mem, total_mem = free_mem / (1024**3), total_mem / (1024**3)
|
||||
mem_msg = f'GPU {0}: free mem {free_mem:.3f}G, total mem {total_mem:.3f}G \n'
|
||||
xformer_res_data = xformer_res.clone().detach().cpu()
|
||||
print('xformer attn ', mem_msg,
|
||||
f'Cost time per time {(time.time() - st) / run_num}s')
|
||||
#
|
||||
torch.cuda.empty_cache()
|
||||
# mask = None
|
||||
st = time.time()
|
||||
for i in tqdm(range(run_num)):
|
||||
flash_res = flash_attn(x.clone(), context.clone(),
|
||||
mask.clone() if mask is not None else mask)
|
||||
if i == run_num - 1:
|
||||
free_mem, total_mem = torch.cuda.mem_get_info(0)
|
||||
free_mem, total_mem = free_mem / (1024**3), total_mem / (1024**3)
|
||||
mem_msg = f'GPU {0}: free mem {free_mem:.3f}G, total mem {total_mem:.3f}G \n'
|
||||
flash_res_data = flash_res.clone().detach().cpu()
|
||||
print('flash attn ', mem_msg,
|
||||
f'Cost time per time {(time.time() - st) / run_num}s')
|
||||
|
||||
# recommend this style for multi blocks to save the preprocess time.
|
||||
|
||||
flash_input = flash_preprocess(
|
||||
x.clone(),
|
||||
context.clone(),
|
||||
mask=mask.clone() if mask is not None else mask)
|
||||
st = time.time()
|
||||
for i in tqdm(range(run_num)):
|
||||
improved_flash_res_v1 = improved_flash_attn(**flash_input).reshape(
|
||||
(batch_size, -1, hidden_dim))
|
||||
if i == 0:
|
||||
improved_flash_res_v1_data = improved_flash_res_v1.clone().detach(
|
||||
).cpu()
|
||||
if i == run_num - 1:
|
||||
free_mem, total_mem = torch.cuda.mem_get_info(0)
|
||||
free_mem, total_mem = free_mem / (1024**3), total_mem / (1024**3)
|
||||
mem_msg = f'GPU {0}: free mem {free_mem:.3f}G, total mem {total_mem:.3f}G \n'
|
||||
print('improved flash attn ', mem_msg,
|
||||
f'Cost time per time {(time.time() - st) / run_num}s')
|
||||
#
|
||||
print(pytorch_res_data, xformer_res_data, flash_res_data,
|
||||
improved_flash_res_v1_data)
|
||||
print(pytorch_res_data.shape, xformer_res_data.shape, flash_res_data.shape,
|
||||
improved_flash_res_v1_data.shape)
|
||||
print(
|
||||
torch.sum(pytorch_res_data) / (batch_size * query_length * hidden_dim),
|
||||
torch.sum(xformer_res_data) / (batch_size * query_length * hidden_dim),
|
||||
torch.sum(flash_res_data) / (batch_size * query_length * hidden_dim),
|
||||
torch.sum(improved_flash_res_v1_data) /
|
||||
(batch_size * query_length * hidden_dim))
|
||||
@@ -0,0 +1,303 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
# All rights reserved.
|
||||
# This file contains code that is adapted from
|
||||
# timm: https://github.com/huggingface/pytorch-image-models
|
||||
# pixart: https://github.com/PixArt-alpha/PixArt-alpha
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
|
||||
from scepter.modules.model.backbone.transformer.attention import drop_path
|
||||
|
||||
|
||||
def modulate(x, shift, scale, unsqueeze=False):
|
||||
if unsqueeze:
|
||||
return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
||||
else:
|
||||
return x * (1 + scale) + shift
|
||||
|
||||
|
||||
class DropPath(nn.Module):
|
||||
"""Drop paths (Stochastic Depth) per sample (when applied in main path of residual blocks).
|
||||
"""
|
||||
def __init__(self, drop_prob=None):
|
||||
super(DropPath, self).__init__()
|
||||
self.drop_prob = drop_prob
|
||||
|
||||
def forward(self, x):
|
||||
return drop_path(x, self.drop_prob, self.training)
|
||||
|
||||
|
||||
class MaskFinalLayer(nn.Module):
|
||||
"""
|
||||
The final layer of PixArt.
|
||||
"""
|
||||
def __init__(self, final_hidden_size, c_emb_size, patch_size,
|
||||
out_channels):
|
||||
super().__init__()
|
||||
self.norm_final = nn.LayerNorm(final_hidden_size,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6)
|
||||
self.linear = nn.Linear(final_hidden_size,
|
||||
patch_size * patch_size * out_channels,
|
||||
bias=True)
|
||||
self.adaLN_modulation = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(c_emb_size, 2 * final_hidden_size, bias=True))
|
||||
|
||||
def forward(self, x, t):
|
||||
shift, scale = self.adaLN_modulation(t).chunk(2, dim=1)
|
||||
x = modulate(self.norm_final(x), shift, scale, unsqueeze=True)
|
||||
x = self.linear(x)
|
||||
return x
|
||||
|
||||
|
||||
class DecoderLayer(nn.Module):
|
||||
"""
|
||||
The final layer of PixArt.
|
||||
"""
|
||||
def __init__(self, hidden_size, decoder_hidden_size):
|
||||
super().__init__()
|
||||
self.norm_decoder = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=False,
|
||||
eps=1e-6)
|
||||
self.linear = nn.Linear(hidden_size, decoder_hidden_size, bias=True)
|
||||
self.adaLN_modulation = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(hidden_size, 2 * hidden_size, bias=True))
|
||||
|
||||
def forward(self, x, t):
|
||||
shift, scale = self.adaLN_modulation(t).chunk(2, dim=1)
|
||||
x = modulate(self.norm_decoder(x), shift, scale, unsqueeze=True)
|
||||
x = self.linear(x)
|
||||
return x
|
||||
|
||||
|
||||
class TimestepEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds scalar timesteps into vector representations.
|
||||
"""
|
||||
def __init__(self, hidden_size, frequency_embedding_size=256):
|
||||
super().__init__()
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(frequency_embedding_size, hidden_size, bias=True),
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size, hidden_size, bias=True),
|
||||
)
|
||||
self.frequency_embedding_size = frequency_embedding_size
|
||||
|
||||
@staticmethod
|
||||
def timestep_embedding(t, dim, max_period=10000):
|
||||
"""
|
||||
Create sinusoidal timestep embeddings.
|
||||
:param t: a 1-D Tensor of N indices, one per batch element.
|
||||
These may be fractional.
|
||||
:param dim: the dimension of the output.
|
||||
:param max_period: controls the minimum frequency of the embeddings.
|
||||
:return: an (N, D) Tensor of positional embeddings.
|
||||
"""
|
||||
# https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
|
||||
half = dim // 2
|
||||
freqs = torch.exp(
|
||||
-math.log(max_period) *
|
||||
torch.arange(start=0, end=half, dtype=torch.float32) /
|
||||
half).to(device=t.device)
|
||||
args = t[:, None].float() * freqs[None]
|
||||
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
||||
if dim % 2:
|
||||
embedding = torch.cat(
|
||||
[embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
||||
return embedding
|
||||
|
||||
def forward(self, t):
|
||||
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
|
||||
t_emb = self.mlp(t_freq)
|
||||
return t_emb
|
||||
|
||||
|
||||
class SizeEmbedder(TimestepEmbedder):
|
||||
"""
|
||||
Embeds scalar timesteps into vector representations.
|
||||
"""
|
||||
def __init__(self, hidden_size, frequency_embedding_size=256):
|
||||
super().__init__(hidden_size=hidden_size,
|
||||
frequency_embedding_size=frequency_embedding_size)
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(frequency_embedding_size, hidden_size, bias=True),
|
||||
nn.SiLU(),
|
||||
nn.Linear(hidden_size, hidden_size, bias=True),
|
||||
)
|
||||
self.frequency_embedding_size = frequency_embedding_size
|
||||
self.outdim = hidden_size
|
||||
|
||||
def forward(self, s, bs):
|
||||
if s.ndim == 1:
|
||||
s = s[:, None]
|
||||
assert s.ndim == 2
|
||||
if s.shape[0] != bs:
|
||||
s = s.repeat(bs // s.shape[0], 1)
|
||||
assert s.shape[0] == bs
|
||||
b, dims = s.shape[0], s.shape[1]
|
||||
s = rearrange(s, 'b d -> (b d)')
|
||||
s_freq = self.timestep_embedding(s, self.frequency_embedding_size).to(
|
||||
self.dtype)
|
||||
s_emb = self.mlp(s_freq)
|
||||
s_emb = rearrange(s_emb,
|
||||
'(b d) d2 -> b (d d2)',
|
||||
b=b,
|
||||
d=dims,
|
||||
d2=self.outdim)
|
||||
return s_emb
|
||||
|
||||
@property
|
||||
def dtype(self):
|
||||
# 返回模型参数的数据类型
|
||||
return next(self.parameters()).dtype
|
||||
|
||||
|
||||
class LabelEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.
|
||||
"""
|
||||
def __init__(self, num_classes, hidden_size, dropout_prob):
|
||||
super().__init__()
|
||||
use_cfg_embedding = dropout_prob > 0
|
||||
self.embedding_table = nn.Embedding(num_classes + use_cfg_embedding,
|
||||
hidden_size)
|
||||
self.num_classes = num_classes
|
||||
self.dropout_prob = dropout_prob
|
||||
|
||||
def token_drop(self, labels, force_drop_ids=None):
|
||||
"""
|
||||
Drops labels to enable classifier-free guidance.
|
||||
"""
|
||||
if force_drop_ids is None:
|
||||
drop_ids = torch.rand(labels.shape[0]).cuda() < self.dropout_prob
|
||||
else:
|
||||
drop_ids = force_drop_ids == 1
|
||||
labels = torch.where(drop_ids, self.num_classes, labels)
|
||||
return labels
|
||||
|
||||
def forward(self, labels, train, force_drop_ids=None):
|
||||
use_dropout = self.dropout_prob > 0
|
||||
if (train and use_dropout) or (force_drop_ids is not None):
|
||||
labels = self.token_drop(labels, force_drop_ids)
|
||||
return self.embedding_table(labels)
|
||||
|
||||
|
||||
class CaptionEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.
|
||||
"""
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
hidden_size,
|
||||
uncond_prob,
|
||||
act_layer=nn.GELU(approximate='tanh'),
|
||||
token_num=120):
|
||||
super().__init__()
|
||||
self.y_proj = Mlp(in_features=in_channels,
|
||||
hidden_features=hidden_size,
|
||||
out_features=hidden_size,
|
||||
act_layer=act_layer,
|
||||
drop=0)
|
||||
self.register_buffer(
|
||||
'y_embedding',
|
||||
nn.Parameter(
|
||||
torch.randn(token_num, in_channels) / in_channels**0.5))
|
||||
self.uncond_prob = uncond_prob
|
||||
|
||||
def token_drop(self, caption, force_drop_ids=None):
|
||||
"""
|
||||
Drops labels to enable classifier-free guidance.
|
||||
"""
|
||||
if force_drop_ids is None:
|
||||
drop_ids = torch.rand(caption.shape[0]).cuda() < self.uncond_prob
|
||||
else:
|
||||
drop_ids = force_drop_ids == 1
|
||||
caption = torch.where(drop_ids[:, None, None], self.y_embedding,
|
||||
caption)
|
||||
return caption
|
||||
|
||||
def forward(self, caption, train, force_drop_ids=None):
|
||||
if train:
|
||||
assert caption.shape[1:] == self.y_embedding.shape
|
||||
use_dropout = self.uncond_prob > 0
|
||||
if (train and use_dropout) or (force_drop_ids is not None):
|
||||
caption = self.token_drop(caption, force_drop_ids)
|
||||
caption = self.y_proj(caption)
|
||||
return caption
|
||||
|
||||
|
||||
class CaptionEmbedderDoubleBr(nn.Module):
|
||||
"""
|
||||
Embeds class labels into vector representations. Also handles label dropout for classifier-free guidance.
|
||||
"""
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
hidden_size,
|
||||
uncond_prob,
|
||||
act_layer=nn.GELU(approximate='tanh'),
|
||||
token_num=120):
|
||||
super().__init__()
|
||||
self.proj = Mlp(in_features=in_channels,
|
||||
hidden_features=hidden_size,
|
||||
out_features=hidden_size,
|
||||
act_layer=act_layer,
|
||||
drop=0)
|
||||
self.embedding = nn.Parameter(torch.randn(1, in_channels) / 10**0.5)
|
||||
self.y_embedding = nn.Parameter(
|
||||
torch.randn(token_num, in_channels) / 10**0.5)
|
||||
self.uncond_prob = uncond_prob
|
||||
|
||||
def token_drop(self, global_caption, caption, force_drop_ids=None):
|
||||
"""
|
||||
Drops labels to enable classifier-free guidance.
|
||||
"""
|
||||
if force_drop_ids is None:
|
||||
drop_ids = torch.rand(
|
||||
global_caption.shape[0]).cuda() < self.uncond_prob
|
||||
else:
|
||||
drop_ids = force_drop_ids == 1
|
||||
global_caption = torch.where(drop_ids[:, None], self.embedding,
|
||||
global_caption)
|
||||
caption = torch.where(drop_ids[:, None, None, None], self.y_embedding,
|
||||
caption)
|
||||
return global_caption, caption
|
||||
|
||||
def forward(self, caption, train, force_drop_ids=None):
|
||||
assert caption.shape[2:] == self.y_embedding.shape
|
||||
global_caption = caption.mean(dim=2).squeeze()
|
||||
use_dropout = self.uncond_prob > 0
|
||||
if (train and use_dropout) or (force_drop_ids is not None):
|
||||
global_caption, caption = self.token_drop(global_caption, caption,
|
||||
force_drop_ids)
|
||||
y_embed = self.proj(global_caption)
|
||||
return y_embed, caption
|
||||
|
||||
|
||||
class Mlp(nn.Module):
|
||||
""" MLP as used in Vision Transformer, MLP-Mixer and related networks
|
||||
"""
|
||||
def __init__(self,
|
||||
in_features,
|
||||
hidden_features=None,
|
||||
out_features=None,
|
||||
act_layer=nn.GELU,
|
||||
drop=0.):
|
||||
super().__init__()
|
||||
out_features = out_features or in_features
|
||||
hidden_features = hidden_features or in_features
|
||||
self.fc1 = nn.Linear(in_features, hidden_features)
|
||||
self.act = act_layer()
|
||||
self.fc2 = nn.Linear(hidden_features, out_features)
|
||||
self.drop = nn.Dropout(drop)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.fc1(x)
|
||||
x = self.act(x)
|
||||
x = self.drop(x)
|
||||
x = self.fc2(x)
|
||||
x = self.drop(x)
|
||||
return x
|
||||
@@ -0,0 +1,54 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
# All rights reserved.
|
||||
# This file contains code that is adapted from
|
||||
# timm: https://github.com/huggingface/pytorch-image-models
|
||||
# pixart: https://github.com/PixArt-alpha/PixArt-alpha
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class PatchEmbed(nn.Module):
|
||||
""" 2D Image to Patch Embedding
|
||||
"""
|
||||
def __init__(
|
||||
self,
|
||||
patch_size=16,
|
||||
in_chans=3,
|
||||
embed_dim=768,
|
||||
norm_layer=None,
|
||||
flatten=True,
|
||||
bias=True,
|
||||
):
|
||||
super().__init__()
|
||||
self.flatten = flatten
|
||||
self.proj = nn.Conv2d(in_chans,
|
||||
embed_dim,
|
||||
kernel_size=patch_size,
|
||||
stride=patch_size,
|
||||
bias=bias)
|
||||
self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
|
||||
|
||||
def forward(self, x):
|
||||
x = self.proj(x)
|
||||
if self.flatten:
|
||||
x = x.flatten(2).transpose(1, 2) # BCHW -> BNC
|
||||
x = self.norm(x)
|
||||
return x
|
||||
|
||||
|
||||
def unpatchify(x, h, w, c, p_h, p_w):
|
||||
'''
|
||||
Args:
|
||||
x: input tensor for unpatchified with shape as (N, T, patch_size**2 * C).
|
||||
h: tokens' number align height
|
||||
w: tokens' number align width
|
||||
c: output channels
|
||||
p_h: patch size for h
|
||||
p_w: patch size for w
|
||||
Returns: unpatchified imgs with shape as (N, H, W, C)
|
||||
'''
|
||||
assert h * w == x.shape[1]
|
||||
x = x.reshape(shape=(x.shape[0], h, w, p_h, p_w, c))
|
||||
x = torch.einsum('nhwpqc->nchpwq', x)
|
||||
return x.reshape(shape=(x.shape[0], c, h * p_h, w * p_w))
|
||||
@@ -0,0 +1,129 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
# All rights reserved.
|
||||
# This file contains code that is adapted from
|
||||
# timm: https://github.com/huggingface/pytorch-image-models
|
||||
# pixart: https://github.com/PixArt-alpha/PixArt-alpha
|
||||
from itertools import repeat as iter_repeat
|
||||
from typing import Iterable
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
def _ntuple(n):
|
||||
def parse(x):
|
||||
if isinstance(x, Iterable) and not isinstance(x, str):
|
||||
return x
|
||||
return tuple(iter_repeat(x, n))
|
||||
|
||||
return parse
|
||||
|
||||
|
||||
to_1tuple = _ntuple(1)
|
||||
to_2tuple = _ntuple(2)
|
||||
|
||||
|
||||
def get_2d_sincos_pos_embed(embed_dim,
|
||||
grid_size,
|
||||
cls_token=False,
|
||||
extra_tokens=0,
|
||||
lewei_scale=1.0,
|
||||
base_h_size=16.,
|
||||
base_w_size=16):
|
||||
"""
|
||||
grid_size: int of the grid height and width
|
||||
return:
|
||||
pos_embed: [grid_size*grid_size, embed_dim] or [1+grid_size*grid_size, embed_dim] (w/ or w/o cls_token)
|
||||
"""
|
||||
if isinstance(grid_size, int):
|
||||
grid_size = to_2tuple(grid_size)
|
||||
grid_h = np.arange(grid_size[0], dtype=np.float32) / (
|
||||
grid_size[0] / base_h_size) / lewei_scale
|
||||
grid_w = np.arange(grid_size[1], dtype=np.float32) / (
|
||||
grid_size[1] / base_w_size) / lewei_scale
|
||||
grid = np.meshgrid(grid_w, grid_h) # here w goes first
|
||||
grid = np.stack(grid, axis=0)
|
||||
grid = grid.reshape([2, 1, grid_size[1], grid_size[0]])
|
||||
|
||||
pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid)
|
||||
if cls_token and extra_tokens > 0:
|
||||
pos_embed = np.concatenate(
|
||||
[np.zeros([extra_tokens, embed_dim]), pos_embed], axis=0)
|
||||
return pos_embed
|
||||
|
||||
|
||||
def get_2d_sincos_pos_embed_from_grid(embed_dim, grid):
|
||||
assert embed_dim % 2 == 0
|
||||
|
||||
# use half of dimensions to encode grid_h
|
||||
emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2,
|
||||
grid[0]) # (H*W, D/2)
|
||||
emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2,
|
||||
grid[1]) # (H*W, D/2)
|
||||
|
||||
return np.concatenate([emb_h, emb_w], axis=1)
|
||||
|
||||
|
||||
def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
|
||||
"""
|
||||
embed_dim: output dimension for each position
|
||||
pos: a list of positions to be encoded: size (M,)
|
||||
out: (M, D)
|
||||
"""
|
||||
assert embed_dim % 2 == 0
|
||||
omega = np.arange(embed_dim // 2, dtype=np.float64)
|
||||
omega /= embed_dim / 2.
|
||||
omega = 1. / 10000**omega # (D/2,)
|
||||
|
||||
pos = pos.reshape(-1) # (M,)
|
||||
out = np.einsum('m,d->md', pos, omega) # (M, D/2), outer product
|
||||
|
||||
emb_sin = np.sin(out) # (M, D/2)
|
||||
emb_cos = np.cos(out) # (M, D/2)
|
||||
return np.concatenate([emb_sin, emb_cos], axis=1)
|
||||
|
||||
|
||||
def apply_2d_rope(xq,
|
||||
xk,
|
||||
padded_pos_index,
|
||||
num_head,
|
||||
head_dim,
|
||||
rotary_base=10000):
|
||||
'''
|
||||
x query/key: [b, seq, num_head*head_dim]
|
||||
padded_pos_index: [b, seq, 2]
|
||||
'''
|
||||
b = xq.shape[0]
|
||||
assert head_dim % 4 == 0, 'the 2d_rope dims should be divided by 4'
|
||||
rope_dim = head_dim // 2 # 2d_rope_dim, 1d_rope_dim = head_dim
|
||||
# 1. theta_d = b ** (-2d/D)
|
||||
theta = 1.0 / (rotary_base**(
|
||||
torch.arange(0, rope_dim, 2)[:(rope_dim // 2)].float() / rope_dim))
|
||||
# 2. [h * Theta || w * Theta]
|
||||
theta = theta.to(xq.device).expand(b, 1, rope_dim // 2)
|
||||
freqs_h = torch.bmm(padded_pos_index[:, :, :1],
|
||||
theta).float() # h * \theta
|
||||
freqs_w = torch.bmm(padded_pos_index[:, :, 1:],
|
||||
theta).float() # w * \theta
|
||||
freqs = torch.cat([freqs_h, freqs_w], dim=2).repeat(1, 1,
|
||||
num_head) # multi-head
|
||||
# 3. as_complex for complex multiply
|
||||
# if freqs = [x, y] then freqs_cis = [cos(x) + sin(x)i, cos(y) + sin(y)i]
|
||||
freqs_cis = torch.polar(
|
||||
torch.ones_like(freqs),
|
||||
freqs) # torch.polar(abs, angle)=> abs⋅cos(angle)+abs⋅sin(angle)⋅j
|
||||
# xq.shape = [b, seq_len, dim]
|
||||
# xq_.shape = [b, seq_len, dim // 2, 2]
|
||||
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 2)
|
||||
xk_ = xk.float().reshape(*xk.shape[:-1], -1, 2)
|
||||
# 转为复数域
|
||||
xq_ = torch.view_as_complex(
|
||||
xq_) # [b, seq_len, dim // 2, 2]=>xq.shape = [b, seq_len, dim]
|
||||
xk_ = torch.view_as_complex(xk_)
|
||||
# 4. complex multiply and as real
|
||||
# xq_out.shape = [b, seq_len, dim]
|
||||
xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(
|
||||
2) # point_wise mul, then flatten eg[[1,2],[3,4],[5,6]]->[1,2,3,4,5,6]
|
||||
xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(2)
|
||||
return xq_out.type_as(xq), xk_out.type_as(xk)
|
||||
Reference in New Issue
Block a user