[CUDA] Fix performance bug in DecoderMaskedMultiheadAttention for BeamSearch (#17613)

This commit is contained in:
Hariharan Seshadri 2023-09-20 10:35:15 -07:00 committed by GitHub
parent e6301eee6a
commit c65e892089
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23

View file

@ -174,7 +174,6 @@ __global__ void masked_multihead_attention_kernel(DecoderMaskedMultiHeadAttentio
q = add_vec(q, q_bias);
}
T* params_k_cache = reinterpret_cast<T*>(params.k_cache);
const float inv_sqrt_dh = params.scale;
@ -350,24 +349,22 @@ __global__ void masked_multihead_attention_kernel(DecoderMaskedMultiHeadAttentio
// The keys loaded from the key cache.
K_vec_k k_vec[K_VECS_PER_THREAD];
if (ti < tlength) {
if (has_beams) {
const int beam_offset = beam_indices[ti] * params.num_heads * params.max_sequence_length * head_size;
if (has_beams) {
#pragma unroll
for (int ii = 0; ii < K_VECS_PER_THREAD; ++ii) {
int jj = ii * params.max_sequence_length + ti;
for (int ii = 0; ii < K_VECS_PER_THREAD; ++ii) {
int jj = ii * params.max_sequence_length + ti;
if (ti < tlength) {
const int beam_offset = beam_indices[ti] * params.num_heads * params.max_sequence_length * head_size;
k_vec[ii] = vec_conversion<K_vec_k, K_vec_m>(
(*reinterpret_cast<const K_vec_m*>(&k_cache_batch[beam_offset + jj * QK_ELTS_IN_16B])));
}
}
} else {
} else {
#pragma unroll
for (int ii = 0; ii < K_VECS_PER_THREAD; ++ii) {
int jj = ii * params.max_sequence_length + ti;
for (int ii = 0; ii < K_VECS_PER_THREAD; ++ii) {
int jj = ii * params.max_sequence_length + ti;
if (ti < tlength) {
k_vec[ii] = vec_conversion<K_vec_k, K_vec_m>(
(*reinterpret_cast<const K_vec_m*>(&k_cache_batch[jj * QK_ELTS_IN_16B])));
}