Update llama.py
This commit is contained in:
parent
302fe3e8fe
commit
e1ed9c85a0
1 changed files with 3 additions and 2 deletions
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue