Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 14 additions & 12 deletions colossalai/shardformer/modeling/llama.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,11 +25,8 @@ def apply_rotary_pos_emb(q, k, cos, sin, position_ids):

def get_llama_forward():

try:
from xformers.ops import memory_efficient_attention as me_attention
except:
raise ImportError("Error: xformers module is not installed. Please install it to use flash attention.")

from colossalai.kernel.cuda_native.flash_attention import AttnMaskType, ColoAttention

def llama_flash_attention_forward(
self,
hidden_states: torch.Tensor,
Expand Down Expand Up @@ -64,19 +61,24 @@ def llama_flash_attention_forward(
key_states = key_states.transpose(1, 2).contiguous().view(*me_input_shape)
value_states = value_states.transpose(1, 2).contiguous().view(*me_input_shape)

flash_attention_mask = None
attn_mask_type = AttnMaskType.causal
if attention_mask != 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()}")
attention_mask = attention_mask.expand(bsz, self.num_heads, q_len, kv_seq_len).contiguous()
flash_attention_mask = ~(attention_mask[:, :, -1].squeeze(1).to(torch.bool)).contiguous()
attn_mask_type = AttnMaskType.paddedcausal

attention = ColoAttention(embed_dim=self.hidden_size, num_heads=self.num_heads)
attn_output = attention(query_states,
key_states,
value_states,
attn_mask=flash_attention_mask,
attn_mask_type=attn_mask_type)

attn_output = me_attention(query_states, key_states, value_states, attn_bias=attention_mask)
if attn_output.size() != (bsz, q_len, self.num_heads, self.head_dim):
raise ValueError(f"`attn_output` should be of size {(bsz, q_len, self.num_heads, self.head_dim)}, but is"
f" {attn_output.size()}")
attn_output = attn_output.reshape(bsz, q_len, self.hidden_size)
attn_output = self.o_proj(attn_output)

return attn_output, None, past_key_value

return llama_flash_attention_forward
37 changes: 17 additions & 20 deletions colossalai/shardformer/modeling/opt.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,12 +6,9 @@


def get_opt_forward():

try:
from xformers.ops import memory_efficient_attention as me_attention
except:
raise ImportError("Error: xformers module is not installed. Please install it to use flash attention.")


from colossalai.kernel.cuda_native.flash_attention import AttnMaskType, ColoAttention

def opt_flash_attention_forward(
self,
hidden_states: torch.Tensor,
Expand All @@ -34,7 +31,7 @@ def opt_flash_attention_forward(
query_states = self.q_proj(hidden_states).view(*attention_input_shape)
# get key, value proj
if is_cross_attention and past_key_value is not None:
# reuse k,v, cross_attentions
# reuse k, v, cross_attentions
key_states = past_key_value[0].transpose(1, 2).contiguous().view(*attention_input_shape)
value_states = past_key_value[1].transpose(1, 2).contiguous().view(*attention_input_shape)
elif is_cross_attention:
Expand Down Expand Up @@ -66,27 +63,27 @@ def opt_flash_attention_forward(
if layer_head_mask != None:
if layer_head_mask.size() != (self.num_heads,):
raise ValueError(f"Head mask for a single layer should be of size {(self.num_heads,)}, but is"
f" {layer_head_mask.size()}")
f" {layer_head_mask.size()}")
flash_attention_mask = None
attn_mask_type = AttnMaskType.causal
if attention_mask != None:
if attention_mask.size() != (bsz, 1, tgt_len, src_len):
raise ValueError(
f"Attention mask should be of size {(bsz, 1, tgt_len, src_len)}, but is {attention_mask.size()}")
attention_mask = attention_mask.expand(bsz, self.num_heads, tgt_len, tgt_len).contiguous()
flash_attention_mask = ~(attention_mask[:, :, -1].squeeze(1).to(torch.bool)).contiguous()
attn_mask_type = AttnMaskType.paddedcausal

attn_output = me_attention(query_states,
attention = ColoAttention(embed_dim=self.embed_dim,
num_heads=self.num_heads,
dropout=self.dropout,
scale=self.scaling)
attn_output = attention(query_states,
key_states,
value_states,
attn_bias=attention_mask,
p=self.dropout,
scale=self.scaling)

attn_output = attn_output.view(bsz, tgt_len, self.num_heads, self.head_dim)

# Use the `embed_dim` from the config (stored in the class) rather than `hidden_state` because `attn_output` can be
# partitioned aross GPUs when using tensor-parallelism.
attn_output = attn_output.reshape(bsz, tgt_len, self.embed_dim)
attn_mask=flash_attention_mask,
attn_mask_type=attn_mask_type)

attn_output = self.out_proj(attn_output)
return attn_output, None, past_key_value

return opt_flash_attention_forward