Update llama.py

This commit is contained in:
Daniel Han 2024-11-05 01:55:17 -08:00
commit e1ed9c85a0

View file

@ -390,7 +390,7 @@ def LlamaAttention_fast_forward(
past_key_value = (K, V) if use_cache else None
# Attention module
if False:#(not HAS_FLASH_ATTENTION and attention_mask is None):
if (not HAS_FLASH_ATTENTION and attention_mask is None):
# Xformers memory efficient attention
# Also has Flash Attention v2 dispatching
Q = Q.transpose(1, 2)
@ -430,7 +430,7 @@ def LlamaAttention_fast_forward(
Q, K, V = Q.contiguous(), K.contiguous(), V.contiguous()
# Needs (batch_size, n_heads, seq_len, head_dim)
# is_casual and attention_mask must not be both set!
A = scaled_dot_product_attention(Q, K, V, is_causal = True)
A = scaled_dot_product_attention(Q, K, V, attn_mask = attention_mask, is_causal = False)
# Go back to (batch_size, seq_len, n_heads, head_dim)
A = A.transpose(1, 2).contiguous()
pass
@ -527,6 +527,7 @@ __DTYPE_MAP = {
}
# https://github.com/huggingface/transformers/blob/main/src/transformers/models/llama/modeling_llama.py#L825
@torch._disable_dynamo
def LlamaModel_fast_forward(
self,
input_ids: torch.LongTensor,