diff --git a/src/lmflow/utils/flash_attention/llama_flash_attention.py b/src/lmflow/utils/flash_attention/llama_flash_attention.py index 4159629c6..5855c9146 100644 --- a/src/lmflow/utils/flash_attention/llama_flash_attention.py +++ b/src/lmflow/utils/flash_attention/llama_flash_attention.py @@ -2,17 +2,18 @@ import torch from torch import nn +import math import transformers -from transformers.models.llama.modeling_llama import apply_rotary_pos_emb +from transformers.models.llama.modeling_llama import apply_rotary_pos_emb,_make_causal_mask,_expand_mask from einops import rearrange #try to import flash_attn 2.x.x, if not, import flash_attn 1.x.x try: - from flash_attn.flash_attn_interface import flash_attn_varlen_qkvpacked_func as flash_attn_unpadded_qkvpacked_func + from flash_attn.flash_attn_interface import flash_attn_varlen_qkvpacked_func except: - from flash_attn.flash_attn_interface import flash_attn_unpadded_qkvpacked_func + from flash_attn.flash_attn_interface import flash_attn_unpadded_qkvpacked_func as flash_attn_varlen_qkvpacked_func from flash_attn.bert_padding import unpad_input, pad_input @@ -51,19 +52,69 @@ def forward( # [bsz, nh, q_len, hd] kv_seq_len = key_states.shape[-2] - assert past_key_value is None, "past_key_value is not supported" + # assert past_key_value is None, "past_key_value is not supported" + if past_key_value is not None: + kv_seq_len += past_key_value[0].shape[-2] cos, sin = self.rotary_emb(value_states, seq_len=kv_seq_len) query_states, key_states = apply_rotary_pos_emb( query_states, key_states, cos, sin, position_ids ) # [bsz, nh, t, hd] - assert not output_attentions, "output_attentions is not supported" - assert not use_cache, "use_cache is not supported" + # assert not output_attentions, "output_attentions is not supported" + # assert not use_cache, "use_cache is not supported" + if past_key_value is not None: + # reuse k, v, self_attention + key_states = torch.cat([past_key_value[0], key_states], dim=2) + value_states = torch.cat([past_key_value[1], value_states], dim=2) + + past_key_value = (key_states, value_states) if use_cache else None # Flash attention codes from # https://github.com/HazyResearch/flash-attention/blob/main/flash_attn/flash_attention.py + # transform the data into the format required by flash attention + # import pdb; pdb.set_trace() + + if query_states.shape[-2] == 1 or query_states.shape[-2] != key_states.shape[-2]: + # decode token-by-token, do not use flash attention + # in incremental state, do not use flash attention + attn_weights = torch.matmul(query_states, key_states.transpose(2, 3)) / math.sqrt(self.head_dim) + + if attn_weights.size() != (bsz, self.num_heads, q_len, kv_seq_len): + raise ValueError( + f"Attention weights should be of size {(bsz * self.num_heads, q_len, kv_seq_len)}, but is" + f" {attn_weights.size()}" + ) + + if attention_mask is not None: + if attention_mask.size() != (bsz, 1, q_len, kv_seq_len): + raise ValueError( + f"Attention mask should be of size {(bsz, 1, q_len, kv_seq_len)}, but is {attention_mask.size()}" + ) + attn_weights = attn_weights + attention_mask + attn_weights = torch.max(attn_weights, torch.tensor(torch.finfo(attn_weights.dtype).min)) + + # upcast attention to fp32 + attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query_states.dtype) + attn_output = torch.matmul(attn_weights, value_states) + + if attn_output.size() != (bsz, self.num_heads, q_len, self.head_dim): + raise ValueError( + f"`attn_output` should be of size {(bsz, self.num_heads, q_len, self.head_dim)}, but is" + f" {attn_output.size()}" + ) + + attn_output = attn_output.transpose(1, 2) + attn_output = attn_output.reshape(bsz, q_len, self.hidden_size) + + attn_output = self.o_proj(attn_output) + + if not output_attentions: + attn_weights = None + + return attn_output, attn_weights, past_key_value + # transform the data into the format required by flash attention qkv = torch.stack( [query_states, key_states, value_states], dim=2 @@ -79,7 +130,7 @@ def forward( cu_q_lens = torch.arange( 0, (bsz + 1) * q_len, step=q_len, dtype=torch.int32, device=qkv.device ) - output = flash_attn_unpadded_qkvpacked_func( + output = flash_attn_varlen_qkvpacked_func( qkv, cu_q_lens, max_s, 0.0, softmax_scale=None, causal=True ) output = rearrange(output, "(b s) ... -> b s ...", b=bsz) @@ -90,7 +141,7 @@ def forward( x_unpad = rearrange( x_unpad, "nnz (three h d) -> nnz three h d", three=3, h=nheads ) - output_unpad = flash_attn_unpadded_qkvpacked_func( + output_unpad = flash_attn_varlen_qkvpacked_func( x_unpad, cu_q_lens, max_s, 0.0, softmax_scale=None, causal=True ) output = rearrange( @@ -100,7 +151,7 @@ def forward( "b s (h d) -> b s h d", h=nheads, ) - return self.o_proj(rearrange(output, "b s h d -> b s (h d)")), None, None + return self.o_proj(rearrange(output, "b s h d -> b s (h d)")), None, past_key_value # Disable the transformation of the attention mask in LlamaModel as the flash attention @@ -109,11 +160,31 @@ def _prepare_decoder_attention_mask( self, attention_mask, input_shape, inputs_embeds, past_key_values_length ): # [bsz, seq_len] - return attention_mask + if input_shape[-1] > 1 and past_key_values_length == 0: # encode + return attention_mask + # create causal mask + # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len] + combined_attention_mask = None + if input_shape[-1] > 1: + combined_attention_mask = _make_causal_mask( + input_shape, + inputs_embeds.dtype, + device=inputs_embeds.device, + past_key_values_length=past_key_values_length, + ) + + if attention_mask is not None: + # [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len] + expanded_attn_mask = _expand_mask(attention_mask, inputs_embeds.dtype, tgt_len=input_shape[-1]).to( + inputs_embeds.device + ) + combined_attention_mask = ( + expanded_attn_mask if combined_attention_mask is None else expanded_attn_mask + combined_attention_mask + ) + + return combined_attention_mask def replace_llama_attn_with_flash_attn(): - transformers.models.llama.modeling_llama.LlamaModel._prepare_decoder_attention_mask = ( - _prepare_decoder_attention_mask - ) + transformers.models.llama.modeling_llama.LlamaModel._prepare_decoder_attention_mask = _prepare_decoder_attention_mask transformers.models.llama.modeling_llama.LlamaAttention.forward = forward \ No newline at end of file