All attention refactor fix (#1491)

* change initilization of n_heads, n_kv_heads, hidden_size in llama.py

* do the same for cohere, mistral, gemma2, granite

* do the same for flexattention,cohere, mistral, granite
This commit is contained in:
Kareem 2025-01-07 17:41:15 +07:00 committed by GitHub
commit fccdf38c45
6 changed files with 47 additions and 39 deletions

View file

@ -43,9 +43,9 @@ if not HAS_FLEX_ATTENTION:
# Logit softcapping
@torch.compile(fullgraph = True, dynamic = True, options = torch_compile_options)
def slow_attention_softcapping(Q, K, V, causal_mask, self, bsz, q_len):
n_heads = self.num_heads
n_heads = self.config.num_attention_heads
head_dim = self.head_dim
n_kv_heads = self.num_key_value_heads
n_kv_heads = self.config.num_key_value_heads
n_groups = self.num_key_value_groups
# Grouped query attention
@ -130,7 +130,7 @@ else:
pass
def slow_attention_softcapping(Q, K, V, causal_mask, self, bsz, q_len):
n_heads = self.num_heads
n_heads = self.config.num_attention_heads
head_dim = self.head_dim
s = self.config.query_pre_attn_scalar
t = self.config.attn_logit_softcapping
@ -147,9 +147,9 @@ torch_matmul = torch.matmul
torch_tanh = torch.tanh
torch_nn_functional_softmax = torch.nn.functional.softmax
def slow_inference_attention_softcapping(Q, K, V, causal_mask, self, bsz, q_len):
n_heads = self.num_heads
n_heads = self.config.num_attention_heads
head_dim = self.head_dim
n_kv_heads = self.num_key_value_heads
n_kv_heads = self.config.num_key_value_heads
n_groups = self.num_key_value_groups
# Grouped query attention

View file

@ -94,9 +94,9 @@ def CohereAttention_fast_forward(
bsz, q_len, _ = hidden_states.size()
n_heads = self.num_heads
n_heads = self.config.num_attention_heads
n_groups = self.num_key_value_groups
n_kv_heads = self.num_key_value_heads
n_kv_heads = self.config.num_key_value_heads
head_dim = self.head_dim
assert(n_kv_heads * n_groups == n_heads)
@ -259,12 +259,14 @@ def CohereAttention_fast_forward_inference(
K1, V1 = past_key_value
dtype = Xn.dtype
n_heads = self.num_heads
n_heads = self.config.num_attention_heads
n_groups = self.num_key_value_groups
n_kv_heads = self.num_key_value_heads
n_kv_heads = self.config.num_key_value_heads
head_dim = self.head_dim
attention_size = n_heads*head_dim
# assert(n_kv_heads * n_groups == n_heads)
hidden_size = self.config.hidden_size
attention_size = n_heads*head_dim
seq_len = K1.shape[-2]
kv_seq_len = seq_len + 1
@ -281,10 +283,10 @@ def CohereAttention_fast_forward_inference(
self.RH_Q = torch.empty((bsz, n_heads, 1, head_dim), dtype = dtype, device = "cuda:0")
# Mistral Nemo 12b has weird dimensions
if attention_size != self.hidden_size:
self.temp_O = torch.empty((1, bsz, self.hidden_size), dtype = dtype, device = "cuda:0")
if attention_size != hidden_size:
self.temp_O = torch.empty((1, bsz, hidden_size), dtype = dtype, device = "cuda:0")
else:
self.temp_O = self.temp_QA[1][:,:,:self.hidden_size]
self.temp_O = self.temp_QA[1][:,:,:hidden_size]
pass
self.attention = torch.empty((bsz, n_heads, 1, KV_CACHE_INCREMENT+seq_len), dtype = dtype, device = "cuda:0")

View file

@ -98,9 +98,9 @@ def Gemma2Attention_fast_forward(
bsz, q_len, _ = hidden_states.size()
n_heads = self.num_heads
n_heads = self.config.num_attention_heads
n_groups = self.num_key_value_groups
n_kv_heads = self.num_key_value_heads
n_kv_heads = self.config.num_key_value_heads
head_dim = self.head_dim
assert(n_kv_heads * n_groups == n_heads)
@ -255,12 +255,14 @@ def Gemma2Attention_fast_forward_inference(
K1, V1 = past_key_value
dtype = Xn.dtype
n_heads = self.num_heads
n_heads = self.config.num_attention_heads
n_groups = self.num_key_value_groups
n_kv_heads = self.num_key_value_heads
n_kv_heads = self.config.num_key_value_heads
head_dim = self.head_dim
attention_size = n_heads*head_dim
# assert(n_kv_heads * n_groups == n_heads)
hidden_size = self.config.hidden_size
attention_size = n_heads*head_dim
seq_len = K1.shape[-2]
kv_seq_len = seq_len + 1
@ -276,7 +278,7 @@ def Gemma2Attention_fast_forward_inference(
self.temp_KV = torch.empty((2, bsz, 1, n_kv_heads*head_dim), dtype = dtype, device = "cuda:0")
self.RH_Q = torch.empty((bsz, n_heads, 1, head_dim), dtype = dtype, device = "cuda:0")
# Only for Gemma2
self.temp_O = torch.empty((1, bsz, self.hidden_size), dtype = dtype, device = "cuda:0")
self.temp_O = torch.empty((1, bsz, hidden_size), dtype = dtype, device = "cuda:0")
self.attention = torch.empty((bsz, n_heads, 1, KV_CACHE_INCREMENT+seq_len), dtype = dtype, device = "cuda:0")
# See https://github.com/google/gemma_pytorch/commit/03e657582d17cb5a8617ebf333c1c16f3694670e

View file

@ -84,9 +84,9 @@ def GraniteAttention_fast_forward(
bsz, q_len, _ = hidden_states.size()
n_heads = self.num_heads
n_heads = self.config.num_attention_heads
n_groups = self.num_key_value_groups
n_kv_heads = self.num_key_value_heads
n_kv_heads = self.config.num_key_value_heads
head_dim = self.head_dim
assert(n_kv_heads * n_groups == n_heads)
@ -257,12 +257,14 @@ def GraniteAttention_fast_forward_inference(
K1, V1 = past_key_value
dtype = Xn.dtype
n_heads = self.num_heads
n_heads = self.config.num_attention_heads
n_groups = self.num_key_value_groups
n_kv_heads = self.num_key_value_heads
n_kv_heads = self.config.num_key_value_heads
head_dim = self.head_dim
attention_size = n_heads*head_dim
# assert(n_kv_heads * n_groups == n_heads)
hidden_size = self.config.hidden_size
attention_size = n_heads*head_dim
seq_len = K1.shape[-2]
kv_seq_len = seq_len + 1
@ -278,7 +280,7 @@ def GraniteAttention_fast_forward_inference(
self.temp_KV = torch.empty((2, bsz, 1, n_kv_heads*head_dim), dtype = dtype, device = "cuda:0")
self.RH_Q = torch.empty((bsz, n_heads, 1, head_dim), dtype = dtype, device = "cuda:0")
# Only for Gemma2
self.temp_O = torch.empty((1, bsz, self.hidden_size), dtype = dtype, device = "cuda:0")
self.temp_O = torch.empty((1, bsz, hidden_size), dtype = dtype, device = "cuda:0")
self.attention = torch.empty((bsz, n_heads, 1, KV_CACHE_INCREMENT+seq_len), dtype = dtype, device = "cuda:0")

View file

@ -146,12 +146,14 @@ def LlamaAttention_fast_forward_inference(
K1, V1 = past_key_value
dtype = Xn.dtype
n_heads = self.num_heads
n_heads = self.config.num_attention_heads
n_groups = self.num_key_value_groups
n_kv_heads = self.num_key_value_heads
n_kv_heads = self.config.num_key_value_heads
head_dim = self.head_dim
attention_size = n_heads*head_dim
# assert(n_kv_heads * n_groups == n_heads)
hidden_size = self.config.hidden_size
attention_size = n_heads*head_dim
seq_len = K1.shape[-2]
kv_seq_len = seq_len + 1
@ -168,10 +170,10 @@ def LlamaAttention_fast_forward_inference(
self.RH_Q = torch.empty((bsz, n_heads, 1, head_dim), dtype = dtype, device = "cuda:0")
# Mistral Nemo 12b has weird dimensions
if attention_size != self.hidden_size:
self.temp_O = torch.empty((1, bsz, self.hidden_size), dtype = dtype, device = "cuda:0")
if attention_size != hidden_size:
self.temp_O = torch.empty((1, bsz, hidden_size), dtype = dtype, device = "cuda:0")
else:
self.temp_O = self.temp_QA[1][:,:,:self.hidden_size]
self.temp_O = self.temp_QA[1][:,:,:hidden_size]
pass
self.attention = torch.empty((bsz, n_heads, 1, KV_CACHE_INCREMENT+seq_len), dtype = dtype, device = "cuda:0")
@ -356,9 +358,9 @@ def LlamaAttention_fast_forward(
bsz, q_len, _ = hidden_states.size()
n_heads = self.num_heads
n_heads = self.config.num_attention_heads
n_groups = self.num_key_value_groups
n_kv_heads = self.num_key_value_heads
n_kv_heads = self.config.num_key_value_heads
head_dim = self.head_dim
assert(n_kv_heads * n_groups == n_heads)

View file

@ -64,9 +64,9 @@ def MistralAttention_fast_forward(
bsz, q_len, _ = hidden_states.size()
n_heads = self.num_heads
n_heads = self.config.num_attention_heads
n_groups = self.num_key_value_groups
n_kv_heads = self.num_key_value_heads
n_kv_heads = self.config.num_key_value_heads
head_dim = self.head_dim
assert(n_kv_heads * n_groups == n_heads)
@ -278,16 +278,16 @@ pass
# Transformers had to update for Mistral Nemo 12b since Attention is (5120, 4096) now.
def patch_mistral_nemo_attention(function):
function = function.replace(
"(self.head_dim * self.num_heads) != self.hidden_size",
"(self.head_dim * self.config.num_attention_heads) != self.config.hidden_size",
"False",
)
function = function.replace(
"self.head_dim = self.hidden_size // self.num_heads",
"self.head_dim = self.config.hidden_size // self.config.num_attention_heads",
"self.head_dim = config.head_dim",
)
function = function.replace(
"self.o_proj = nn.Linear(self.hidden_size, self.hidden_size, bias=False)",
"self.o_proj = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=False)",
"self.o_proj = nn.Linear(self.config.hidden_size, self.config.hidden_size, bias=False)",
"self.o_proj = nn.Linear(self.config.num_attention_heads * self.head_dim, self.config.hidden_size, bias=False)",
)
return function
pass