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:
parent
a243ddb4f0
commit
fccdf38c45
6 changed files with 47 additions and 39 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue