File size: 6,574 Bytes
6a176cb | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 | # ------------------------------------------------------------------------
# Copyright (c) 2023 megvii-model. All Rights Reserved.
# ------------------------------------------------------------------------
# Modified by Shihao Wang
# ------------------------------------------------------------------------
# flash-attention
import math
import torch
import torch.nn as nn
from torch.nn.init import (
xavier_uniform_,
constant_,
xavier_normal_
)
from torch.nn.functional import linear
from einops import rearrange
from mmcv.runner import auto_fp16
from mmcv.runner.base_module import BaseModule
try:
from flash_attn.flash_attn_interface import flash_attn_unpadded_kvpacked_func
except ImportError:
from flash_attn.flash_attn_interface import (
flash_attn_varlen_kvpacked_func as flash_attn_unpadded_kvpacked_func,
)
from flash_attn.bert_padding import unpad_input, pad_input, index_first_axis
def _in_projection_packed(q, k, v, w, b = None):
w_q, w_k, w_v = w.chunk(3)
if b is None:
b_q = b_k = b_v = None
else:
b_q, b_k, b_v = b.chunk(3)
return linear(q, w_q, b_q), linear(k, w_k, b_k), linear(v, w_v, b_v)
class FlashAttention(nn.Module):
"""Implement the scaled dot product attention with softmax.
Arguments
---------
softmax_scale: The temperature to use for the softmax attention.
(default: 1/sqrt(d_keys) where d_keys is computed at
runtime)
attention_dropout: The dropout rate to apply to the attention
(default: 0.1)
"""
def __init__(self, softmax_scale=None, attention_dropout=0.0, device=None, dtype=None):
super().__init__()
self.softmax_scale = softmax_scale
self.dropout_p = attention_dropout
self.fp16_enabled = True
@auto_fp16(apply_to=('q', 'kv'), out_fp32=True)
def forward(self, q, kv,
causal=False,
key_padding_mask=None):
"""Implements the multihead softmax attention.
Arguments
---------
q: The tensor containing the query. (B, T, H, D)
kv: The tensor containing the key, and value. (B, S, 2, H, D)
key_padding_mask: a bool tensor of shape (B, S)
"""
assert q.dtype in [torch.float16, torch.bfloat16] and kv.dtype in [torch.float16, torch.bfloat16]
assert q.is_cuda and kv.is_cuda
assert q.shape[0] == kv.shape[0] and q.shape[-2] == kv.shape[-2] and q.shape[-1] == kv.shape[-1]
batch_size = q.shape[0]
seqlen_q, seqlen_k = q.shape[1], kv.shape[1]
if key_padding_mask is None:
q, kv = rearrange(q, 'b s ... -> (b s) ...'), rearrange(kv, 'b s ... -> (b s) ...')
max_sq, max_sk = seqlen_q, seqlen_k
cu_seqlens_q = torch.arange(0, (batch_size + 1) * seqlen_q, step=seqlen_q, dtype=torch.int32,
device=q.device)
cu_seqlens_k = torch.arange(0, (batch_size + 1) * seqlen_k, step=seqlen_k, dtype=torch.int32,
device=kv.device)
output = flash_attn_unpadded_kvpacked_func(
q, kv, cu_seqlens_q, cu_seqlens_k, max_sq, max_sk,
self.dropout_p if self.training else 0.0,
softmax_scale=self.softmax_scale, causal=causal
)
output = rearrange(output, '(b s) ... -> b s ...', b=batch_size)
else:
nheads = kv.shape[-2]
q = rearrange(q, 'b s ... -> (b s) ...')
max_sq = seqlen_q
cu_seqlens_q = torch.arange(0, (batch_size + 1) * seqlen_q, step=seqlen_q, dtype=torch.int32,
device=q.device)
x = rearrange(kv, 'b s two h d -> b s (two h d)')
x_unpad, indices, cu_seqlens_k, max_sk = unpad_input(x, key_padding_mask)
x_unpad = rearrange(x_unpad, 'nnz (two h d) -> nnz two h d', two=2, h=nheads)
output_unpad = flash_attn_unpadded_kvpacked_func(
q, x_unpad, cu_seqlens_q, cu_seqlens_k, max_sq, max_sk,
self.dropout_p if self.training else 0.0,
softmax_scale=self.softmax_scale, causal=causal
)
output = rearrange(output_unpad, '(b s) ... -> b s ...', b=batch_size)
return output, None
class FlashMHA(nn.Module):
def __init__(self, embed_dim, num_heads, bias=True, batch_first=True, attention_dropout=0.0,
causal=False, device=None, dtype=None, **kwargs) -> None:
assert batch_first
factory_kwargs = {'device': device, 'dtype': dtype}
super().__init__()
self.embed_dim = embed_dim
self.causal = causal
self.bias = bias
self.num_heads = num_heads
assert self.embed_dim % num_heads == 0, "self.kdim must be divisible by num_heads"
self.head_dim = self.embed_dim // num_heads
assert self.head_dim % 8 == 0 and self.head_dim <= 128, "Only support head_dim <= 128 and divisible by 8"
self.in_proj_weight = nn.Parameter(torch.empty((3 * embed_dim, embed_dim)))
if bias:
self.in_proj_bias = nn.Parameter(torch.empty(3 * embed_dim))
else:
self.register_parameter('in_proj_bias', None)
self.inner_attn = FlashAttention(attention_dropout=attention_dropout, **factory_kwargs)
self.out_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
self._reset_parameters()
def _reset_parameters(self) -> None:
xavier_uniform_(self.in_proj_weight)
if self.in_proj_bias is not None:
constant_(self.in_proj_bias, 0.)
constant_(self.out_proj.bias, 0.)
def forward(self, q, k, v, key_padding_mask=None):
"""x: (batch, seqlen, hidden_dim) (where hidden_dim = num heads * head dim)
key_padding_mask: bool tensor of shape (batch, seqlen)
"""
# q, k, v = self.Wq(q), self.Wk(k), self.Wv(v)
q, k, v = _in_projection_packed(q, k, v, self.in_proj_weight, self.in_proj_bias)
q = rearrange(q, 'b s (h d) -> b s h d', h=self.num_heads)
k = rearrange(k, 'b s (h d) -> b s h d', h=self.num_heads)
v = rearrange(v, 'b s (h d) -> b s h d', h=self.num_heads)
kv = torch.stack([k, v], dim=2)
context, attn_weights = self.inner_attn(q, kv, key_padding_mask=key_padding_mask, causal=self.causal)
return self.out_proj(rearrange(context, 'b s h d -> b s (h d)')), attn_weights
|