mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
optimize gqa cpu (#20598)
### Description <!-- Describe your changes. --> optimize the GQA implementation on CPU. Mainly optimization are: 1. compute attention on real total sequence length instead of maximum sequence length in case past/present share same buffer 2. remove the mask 3. remove the transpose after attention x value It improve the phi3 model https://github.com/microsoft/onnxruntime-genai/blob/main/examples/python/phi3-qa.py with max sequence length 2k/4k from 10 tps to 20 tps. ### Motivation and Context <!-- - Why is this change required? What problem does it solve? - If it fixes an open issue, please link to the issue here. -->
This commit is contained in:
parent
1f509215bc
commit
156d52163d
3 changed files with 45 additions and 96 deletions
|
|
@ -153,49 +153,6 @@ void PrepareMask(const int32_t* mask_index,
|
|||
}
|
||||
}
|
||||
|
||||
// Applies causal mask and past seqlens right pad to the mask_data buffer.
|
||||
template <typename T>
|
||||
void PrepareMaskGQA(T* mask_data,
|
||||
int batch_size,
|
||||
int sequence_length,
|
||||
int buffer_sequence_length,
|
||||
int local_window_size,
|
||||
const int32_t* seqlens_k) {
|
||||
// mask_data has been filled with 0, and its shape is BxSxT
|
||||
T* p_mask = mask_data;
|
||||
// TODO: parallelize this
|
||||
for (int b_i = 0; b_i < batch_size; b_i++) {
|
||||
if (sequence_length > 1) {
|
||||
// Apply causal/local mask for prompt case.
|
||||
for (int s_i = 0; s_i < sequence_length; s_i++) {
|
||||
for (int m_i = s_i + 1; m_i < buffer_sequence_length; m_i++) {
|
||||
p_mask[s_i * buffer_sequence_length + m_i] = std::numeric_limits<T>::lowest();
|
||||
}
|
||||
// Apply local mask.
|
||||
if (local_window_size > 0) {
|
||||
for (int m_i = 0; m_i < s_i - local_window_size; m_i++) {
|
||||
p_mask[s_i * buffer_sequence_length + m_i] = std::numeric_limits<T>::lowest();
|
||||
}
|
||||
}
|
||||
}
|
||||
} else if (sequence_length == 1) {
|
||||
// Apply right padding to mask for token gen case.
|
||||
int total_seqlen = seqlens_k[b_i] + 1;
|
||||
for (int m_i = total_seqlen; m_i < buffer_sequence_length; m_i++) {
|
||||
p_mask[m_i] = std::numeric_limits<T>::lowest();
|
||||
}
|
||||
// Apply local mask.
|
||||
if (local_window_size > 0) {
|
||||
for (int m_i = 0; m_i < total_seqlen - local_window_size - 1; m_i++) {
|
||||
p_mask[m_i] = std::numeric_limits<T>::lowest();
|
||||
}
|
||||
}
|
||||
}
|
||||
ptrdiff_t mask_to_advance = SafeInt<ptrdiff_t>(sequence_length) * buffer_sequence_length;
|
||||
p_mask += mask_to_advance;
|
||||
}
|
||||
}
|
||||
|
||||
// Concatenate a past state chunk PxH with input state chunk LxH into present state chunk TxH
|
||||
// Returns a pointer to the start of present state chunk.
|
||||
template <typename T>
|
||||
|
|
|
|||
|
|
@ -55,12 +55,6 @@ class GQAAttentionBase : public AttentionBase {
|
|||
auto attention_probs = allocator->Alloc(bytes);
|
||||
BufferUniquePtr scratch_buffer(attention_probs, BufferDeleter(allocator));
|
||||
|
||||
void* mask_data = nullptr;
|
||||
size_t mask_data_bytes = SafeInt<size_t>(batch_size) * sequence_length * seqlen_present_kv_cache * sizeof(T);
|
||||
mask_data = allocator->Alloc(mask_data_bytes);
|
||||
memset(mask_data, 0, mask_data_bytes);
|
||||
BufferUniquePtr mask_data_buffer(mask_data, BufferDeleter(allocator));
|
||||
|
||||
const T* past_key_data = past_key != nullptr ? past_key->Data<T>() : nullptr;
|
||||
T* present_key_data = present_key != nullptr ? present_key->MutableData<T>() : nullptr;
|
||||
const T* past_value_data = past_value != nullptr ? past_value->Data<T>() : nullptr;
|
||||
|
|
@ -70,17 +64,13 @@ class GQAAttentionBase : public AttentionBase {
|
|||
|
||||
const T* k = packed_qkv ? Q + num_heads_ * sequence_length * head_size : K;
|
||||
ComputeAttentionProbs<T>(static_cast<T*>(attention_probs), Q, k,
|
||||
seqlens_k->Data<int32_t>(), static_cast<T*>(mask_data),
|
||||
seqlens_k->Data<int32_t>(),
|
||||
batch_size, sequence_length, seqlen_past_kv_cache, seqlen_present_kv_cache,
|
||||
head_size, past_key_data, present_key_data, past_present_share_buffer, packed_qkv, tp);
|
||||
|
||||
// Compute the attentionScore * Value: out_tmp(B, N, S, H_v) = attention_probs(B, N, S, T) x V(B, N, T, H_v)
|
||||
auto out_tmp_data =
|
||||
allocator->Alloc(SafeInt<size_t>(batch_size) * num_heads_ * sequence_length * head_size * sizeof(T));
|
||||
BufferUniquePtr out_tmp_buffer(out_tmp_data, BufferDeleter(allocator));
|
||||
|
||||
// Compute the attentionScore * Value: out(B, N, S, H_v) = attention_probs(B, N, S, T) x V(B, N, T, H_v)
|
||||
const T* v = packed_qkv ? Q + (num_heads_ + kv_num_heads_) * sequence_length * head_size : V;
|
||||
ComputeVxAttentionScore(output->MutableData<T>(), static_cast<T*>(out_tmp_data), static_cast<T*>(attention_probs),
|
||||
ComputeVxAttentionScore(output->MutableData<T>(), static_cast<T*>(attention_probs),
|
||||
v, seqlens_k->Data<int32_t>(), batch_size, sequence_length, seqlen_past_kv_cache,
|
||||
seqlen_present_kv_cache, head_size, hidden_size, past_value_data, present_value_data,
|
||||
past_present_share_buffer, packed_qkv, tp);
|
||||
|
|
@ -90,15 +80,13 @@ class GQAAttentionBase : public AttentionBase {
|
|||
|
||||
private:
|
||||
// Helper function to compute the attention probs. It does 2 things:
|
||||
// attention_probs(B, N, S, T) = 1/sqrt(H) x Q(B, N, S, H) x K'(B, N, T, H -> B, N, H, T) +
|
||||
// 1 x mask_data(B, N, S, T)
|
||||
// attention_probs(B, N, S, T) = 1/sqrt(H) x Q(B, N, S, H) x K'(B, N, T, H -> B, N, H, T)
|
||||
// attention_probs(B, N, S, T) = Softmax(attention_probs)
|
||||
template <typename T>
|
||||
void ComputeAttentionProbs(T* attention_probs, // output buffer with size BxNxSxT
|
||||
const T* Q, // Q data. Its size is BxNxSxH
|
||||
const T* K, // k data. Its size is BxNxLxH
|
||||
const int32_t* seqlens_k, // past sequence lengths tensor
|
||||
T* mask_data, // buffer for mask data.
|
||||
int batch_size, // batch size of self-attention
|
||||
int sequence_length, // sequence length of self-attention (S)
|
||||
int past_buffer_sequence_length, // sequence length of past state
|
||||
|
|
@ -117,7 +105,6 @@ class GQAAttentionBase : public AttentionBase {
|
|||
const size_t past_buff_chunk_length = static_cast<size_t>(past_buffer_sequence_length) * head_size; // L x H
|
||||
const size_t present_buff_chunk_length = static_cast<size_t>(present_buffer_sequence_length) * head_size; // T x H
|
||||
|
||||
PrepareMaskGQA(mask_data, batch_size, sequence_length, present_buffer_sequence_length, local_window_size_, seqlens_k);
|
||||
if (!past_present_share_buffer) {
|
||||
memset(present_key, 0, batch_size * kv_num_heads_ * present_buffer_sequence_length * head_size * sizeof(T));
|
||||
}
|
||||
|
|
@ -146,16 +133,11 @@ class GQAAttentionBase : public AttentionBase {
|
|||
const int head_index = static_cast<int>(i) % num_heads_;
|
||||
const int past_seqlen = sequence_length == 1 ? static_cast<int>(seqlens_k[batch_index]) : past_buffer_sequence_length;
|
||||
const size_t past_chunk_length = static_cast<size_t>(past_seqlen) * head_size;
|
||||
const int total_seqlen = seqlens_k[batch_index] + 1;
|
||||
|
||||
const int output_offset = static_cast<int>(i) * sequence_length * present_buffer_sequence_length;
|
||||
const int mask_offset = batch_index * sequence_length * present_buffer_sequence_length;
|
||||
T* output = attention_probs + output_offset;
|
||||
|
||||
// Broadcast mask data: (Bx)SxT -> (BxNx)SxT
|
||||
memcpy(output,
|
||||
mask_data + mask_offset,
|
||||
probs_matrix_bytes);
|
||||
|
||||
const T* k;
|
||||
if (packed_qkv) {
|
||||
k = K + packed_batch_stride * batch_index + kv_input_chunk_length * (head_index / kv_num_heads_factor);
|
||||
|
|
@ -168,7 +150,6 @@ class GQAAttentionBase : public AttentionBase {
|
|||
i / kv_num_heads_factor);
|
||||
}
|
||||
|
||||
// TODO: CblasTrans stuff what do?
|
||||
// Compute Q*K' + AttentionMask
|
||||
// original transposed each iteration
|
||||
// A: Q (B x N x) S x H (B x N x) S x H S x H
|
||||
|
|
@ -180,20 +161,38 @@ class GQAAttentionBase : public AttentionBase {
|
|||
} else {
|
||||
q = Q + q_input_chunk_length * i;
|
||||
}
|
||||
math::Gemm<T, ThreadPool>(CblasNoTrans, CblasTrans, sequence_length, present_buffer_sequence_length, head_size, alpha,
|
||||
q, k, mask_data != nullptr ? 1.0f : 0.0f, output, nullptr);
|
||||
math::GemmEx<T, ThreadPool>(CblasNoTrans, CblasTrans,
|
||||
sequence_length, total_seqlen, head_size, alpha,
|
||||
q, head_size, k, head_size,
|
||||
0.0f /*bata*/,
|
||||
output, present_buffer_sequence_length, nullptr);
|
||||
|
||||
// compute Softmax
|
||||
T* output_softmax = output;
|
||||
for (int seq = 0; seq < sequence_length; seq++) {
|
||||
int seq_causal_length = sequence_length == 1 ? total_seqlen : seq + 1;
|
||||
if (local_window_size_ > 0 && seq_causal_length > local_window_size_ + 1) {
|
||||
for (int total_seq_id = 0; total_seq_id < seq_causal_length - local_window_size_ - 1; total_seq_id++) {
|
||||
output_softmax[total_seq_id] = 0.f;
|
||||
}
|
||||
ComputeAttentionSoftmaxInplace(output_softmax + seq_causal_length - local_window_size_ - 1, 1, local_window_size_ + 1, nullptr);
|
||||
} else {
|
||||
ComputeAttentionSoftmaxInplace(output_softmax, 1, seq_causal_length, nullptr);
|
||||
}
|
||||
|
||||
// set causal [seq_causal_length, total_seqlen) to 0.f
|
||||
for (int total_seq_id = seq_causal_length; total_seq_id < total_seqlen; total_seq_id++) {
|
||||
output_softmax[total_seq_id] = 0.f;
|
||||
}
|
||||
|
||||
output_softmax += present_buffer_sequence_length;
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
// attention_probs(B, N, S, T) = Softmax(attention_probs)
|
||||
const int N = batch_size * num_heads_ * sequence_length;
|
||||
const int D = present_buffer_sequence_length;
|
||||
ComputeAttentionSoftmaxInplace(attention_probs, N, D, tp);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
void ComputeVxAttentionScore(T* output, // buffer for the result with size BxSxNxH
|
||||
T* tmp_buffer, // buffer for temp use with size is BxNxSxH
|
||||
const T* attention_probs, // Attention probs with size BxNxSxT
|
||||
const T* V, // V value with size BxN_kvxSxH
|
||||
const int32_t* seqlens_k, // past sequence lengths tensor
|
||||
|
|
@ -211,8 +210,7 @@ class GQAAttentionBase : public AttentionBase {
|
|||
const bool is_prompt = sequence_length != 1;
|
||||
const int packed_batch_stride = packed_qkv ? (num_heads_ + 2 * kv_num_heads_) * sequence_length * head_size : 0;
|
||||
const int kv_num_heads_factor = num_heads_ / kv_num_heads_;
|
||||
const size_t q_input_chunk_length = static_cast<size_t>(sequence_length) * head_size; // S x H
|
||||
const size_t kv_input_chunk_length = static_cast<size_t>(sequence_length) * head_size; // L x H
|
||||
const int kv_input_chunk_length = sequence_length * head_size; // L x H
|
||||
const size_t past_buff_chunk_length = static_cast<size_t>(past_buffer_sequence_length) * head_size; // L x H
|
||||
const size_t present_buff_chunk_length = static_cast<size_t>(present_buffer_sequence_length) * head_size; // T x H
|
||||
|
||||
|
|
@ -243,6 +241,7 @@ class GQAAttentionBase : public AttentionBase {
|
|||
const int head_index = static_cast<int>(i % num_heads_);
|
||||
const int past_seqlen = sequence_length == 1 ? static_cast<int>(seqlens_k[batch_index]) : past_buffer_sequence_length;
|
||||
const size_t past_chunk_length = static_cast<size_t>(past_seqlen) * head_size;
|
||||
const int total_seqlen = seqlens_k[batch_index] + 1;
|
||||
|
||||
const T* v;
|
||||
if (packed_qkv) {
|
||||
|
|
@ -256,21 +255,17 @@ class GQAAttentionBase : public AttentionBase {
|
|||
i / kv_num_heads_factor);
|
||||
}
|
||||
|
||||
T* current_tmp_data = reinterpret_cast<T*>(tmp_buffer) + q_input_chunk_length * static_cast<int>(i);
|
||||
const int attention_probs_offset = sequence_length * present_buffer_sequence_length * static_cast<int>(i);
|
||||
math::MatMul<T>(sequence_length, head_size, present_buffer_sequence_length,
|
||||
attention_probs + attention_probs_offset,
|
||||
v, current_tmp_data, nullptr);
|
||||
T* output_current = output + (batch_index * sequence_length * num_heads_ + head_index) * head_size;
|
||||
ptrdiff_t attention_probs_offset = SafeInt<ptrdiff_t>(sequence_length) * present_buffer_sequence_length * i;
|
||||
|
||||
// Transpose: out(B, S, N, H_v) -> out_tmp(B, N, S, H_v)
|
||||
T* src = current_tmp_data;
|
||||
const int dest_offset = (batch_index * sequence_length * num_heads_ + head_index) * head_size;
|
||||
T* dest = output + dest_offset;
|
||||
for (int j = 0; j < sequence_length; j++) {
|
||||
memcpy(dest, src, bytes_to_copy_trans);
|
||||
src += head_size;
|
||||
dest += hidden_size;
|
||||
}
|
||||
math::GemmEx<T, ThreadPool>(CblasNoTrans,
|
||||
CblasNoTrans,
|
||||
sequence_length, head_size, total_seqlen,
|
||||
1.f, /*alpha*/
|
||||
attention_probs + attention_probs_offset, present_buffer_sequence_length,
|
||||
v, head_size,
|
||||
0.0f /*beta*/,
|
||||
output_current, hidden_size, nullptr);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1294,7 +1294,7 @@ def parity_check_gqa_prompt_no_buff(
|
|||
None,
|
||||
cos,
|
||||
sin,
|
||||
cache_seqlens,
|
||||
cache_seqlens - 1,
|
||||
left_window_size,
|
||||
past_format,
|
||||
False,
|
||||
|
|
@ -1310,7 +1310,7 @@ def parity_check_gqa_prompt_no_buff(
|
|||
new_v,
|
||||
cos,
|
||||
sin,
|
||||
cache_seqlens,
|
||||
cache_seqlens - 1,
|
||||
left_window_size,
|
||||
past_format,
|
||||
False,
|
||||
|
|
@ -1766,9 +1766,6 @@ class TestGQA(unittest.TestCase):
|
|||
seqs = (
|
||||
[
|
||||
(127, 127),
|
||||
(35, 35),
|
||||
(2000, 2000),
|
||||
(200, 200),
|
||||
(240, 240),
|
||||
]
|
||||
if pipeline_mode
|
||||
|
|
|
|||
Loading…
Reference in a new issue