diff --git a/onnxruntime/contrib_ops/cpu/bert/attention_cpu_base.h b/onnxruntime/contrib_ops/cpu/bert/attention_cpu_base.h index b42ad93f1d..ed0bbc57b9 100644 --- a/onnxruntime/contrib_ops/cpu/bert/attention_cpu_base.h +++ b/onnxruntime/contrib_ops/cpu/bert/attention_cpu_base.h @@ -18,18 +18,18 @@ class AttentionCPUBase : public AttentionBase { AttentionCPUBase(const OpKernelInfo& info) : AttentionBase(info) {} template - Status ApplyAttention(const T* Q, // Q data. Its size is BxNxSxH - const T* K, // K data. Its size is BxNxSxH - const T* V, // V value with size BxNxSxH - const Tensor* mask_index, // mask index. nullptr if no mask or its size is B - const Tensor* past, // past state - Tensor* output, // output tensor - int batch_size, // batch size - int sequence_length, // sequence length - int qk_head_size, // qk_head_size - int v_head_size, // head_size - int v_hidden_size, // hidden_size - const Tensor* extra_add_qk,// extra add in QK. Its size is BxNxSxS + Status ApplyAttention(const T* Q, // Q data. Its size is BxNxSxH + const T* K, // K data. Its size is BxNxSxH + const T* V, // V value with size BxNxSxH + const Tensor* mask_index, // mask index. nullptr if no mask or its size is B + const Tensor* past, // past state + Tensor* output, // output tensor + int batch_size, // batch size + int sequence_length, // sequence length + int qk_head_size, // qk_head_size + int v_head_size, // head_size + int v_hidden_size, // hidden_size + const Tensor* extra_add_qk, // extra add in QK. Its size is BxNxSxS OpKernelContext* context) const { AllocatorPtr allocator; ORT_RETURN_IF_ERROR(context->GetTempSpaceAllocator(&allocator)); @@ -50,8 +50,10 @@ class AttentionCPUBase : public AttentionBase { auto attention_probs = allocator->Alloc(attention_probs_bytes); BufferUniquePtr scratch_buffer(attention_probs, BufferDeleter(allocator)); + bool has_unidirectional = (is_unidirectional_ && sequence_length > 1); + void* mask_data = nullptr; - if (mask_index != nullptr || (is_unidirectional_ && sequence_length > 1)) { + if (mask_index != nullptr || has_unidirectional) { size_t mask_data_bytes = SafeInt(batch_size) * sequence_length * all_sequence_length * sizeof(T); mask_data = allocator->Alloc(mask_data_bytes); memset(mask_data, 0, mask_data_bytes); @@ -69,7 +71,7 @@ class AttentionCPUBase : public AttentionBase { } ComputeAttentionProbs(static_cast(attention_probs), Q, K, - mask_index_data, mask_index_dims, static_cast(mask_data), + mask_index_data, mask_index_dims, static_cast(mask_data), has_unidirectional, batch_size, sequence_length, past_sequence_length, qk_head_size == 0 ? v_head_size : qk_head_size, past_data, present_data, tp, extra_add_qk_data); @@ -91,20 +93,22 @@ class AttentionCPUBase : public AttentionBase { // 1 x mask_data(B, N, S, S*) // II.attention_probs(B, N, S, S*) = Softmax(attention_probs) template - void ComputeAttentionProbs(T* attention_probs, // output buffer for the attention probs. Its size is BxNxSxS + void ComputeAttentionProbs(T* attention_probs, // output buffer for the attention probs. Its size is BxNxSxS* const T* Q, // Q data. Its size is BxNxSxH const T* K, // k data. Its size is BxNxSxH const int32_t* mask_index, // mask index. nullptr if no mask or its size is B const std::vector* mask_index_dims, // mask index shape - T* mask_data, // buffer for mask data. It is nullptr if mask_index is nullptr, otherwise its shape is BxSxS* + T* mask_data, // buffer for mask data. It is nullptr if mask_index is nullptr and not unidirectional, otherwise its shape is BxSxS* + bool has_unidirectional, // has unidirectional mask int batch_size, // batch size of self-attention int sequence_length, // sequence length of self-attention int past_sequence_length, // sequence length of past state int head_size, // head size of self-attention const T* past, // past state T* present, // present state - ThreadPool* tp, - const T* extra_add_qk_data) const { + ThreadPool* tp, // thread pool + const T* extra_add_qk_data // extra add matrix with shape BxNxSxS* + ) const { const int all_sequence_length = past_sequence_length + sequence_length; // S* = S' + S const size_t past_chunk_length = static_cast(past_sequence_length) * head_size; // S' x H const size_t input_chunk_length = static_cast(sequence_length) * head_size; // S x H @@ -112,12 +116,11 @@ class AttentionCPUBase : public AttentionBase { { if (mask_data != nullptr) { - // Convert attention mask data from int to float (0 to -10000.0f). The mask_data shape is BxSxS*. - PrepareMask(mask_index, mask_index_dims, mask_data, batch_size, sequence_length, past_sequence_length); + PrepareMask(mask_index, mask_index_dims, mask_data, has_unidirectional, batch_size, sequence_length, past_sequence_length); + } else { // no any mask + memset(attention_probs, 0, static_cast(batch_size) * num_heads_ * sequence_length * all_sequence_length * sizeof(T)); } - memset(attention_probs, 0, static_cast(batch_size) * num_heads_ * sequence_length * all_sequence_length * sizeof(T)); - const int loop_len = batch_size * num_heads_; const float alpha = 1.0f / sqrt(static_cast(head_size)); @@ -126,7 +129,16 @@ class AttentionCPUBase : public AttentionBase { ThreadPool::TryParallelFor(tp, loop_len, cost, [&](std::ptrdiff_t begin, std::ptrdiff_t end) { for (std::ptrdiff_t i = begin; i != end; ++i) { - const std::ptrdiff_t batch_index = i / num_heads_; + const int batch_index = static_cast(i) / num_heads_; + + const int output_offset = static_cast(i) * sequence_length * all_sequence_length; + const int mask_offset = batch_index * sequence_length * all_sequence_length; + T* output = attention_probs + output_offset; + + // Broadcast mask data: (Bx)SxS* -> (BxNx)SxS* + if (mask_data != nullptr) { + memcpy(output, mask_data + mask_offset, sequence_length * all_sequence_length * sizeof(T)); + } const T* k = K + input_chunk_length * i; if (nullptr != present) { @@ -134,10 +146,7 @@ class AttentionCPUBase : public AttentionBase { k = ConcatStateChunk(past, k, present, past_chunk_length, present_chunk_length, i); } - int offset = sequence_length * all_sequence_length * static_cast(i); - T* output = reinterpret_cast(attention_probs) + offset; - - // Compute Q*K' + // 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 // B: K' (B x N x) S* x H (B x N x) H x S* H x S* @@ -146,27 +155,19 @@ class AttentionCPUBase : public AttentionBase { Q + input_chunk_length * i, k, 1.0, output, nullptr); - // Apply unidirectional mask and set future words to -10000.0f. - if (is_unidirectional_) { - for (int s = 0; s < sequence_length - 1; s++) { - for (int t = past_sequence_length + s + 1; t < all_sequence_length; t++) { - output[s * all_sequence_length + t] = static_cast(-10000.0f); + // Fix unidirectional mask to be parity with huggingface implementation. + if (has_unidirectional) { + for (int s_i = 0; s_i < sequence_length - 1; s_i++) { + for (int m_i = past_sequence_length + s_i + 1; m_i < all_sequence_length; m_i++) { + int j = s_i * all_sequence_length + m_i; + output[j] = mask_data[mask_offset + j]; } } } - // Apply attention mask - if (mask_data != nullptr) { - const T* attention_mask = reinterpret_cast(mask_data) + batch_index * sequence_length * all_sequence_length; - for (int j = 0; j < sequence_length * all_sequence_length ; j++) { - output[j] += attention_mask[j]; - } - } - if (extra_add_qk_data != nullptr) { - const T* extra = extra_add_qk_data + offset; - for (int j = 0; j < sequence_length * all_sequence_length ; j++) { - output[j] += extra[j]; + for (int j = 0; j < sequence_length * all_sequence_length; j++) { + output[j] += extra_add_qk_data[output_offset + j]; } } } diff --git a/onnxruntime/contrib_ops/cpu/bert/attention_helper.h b/onnxruntime/contrib_ops/cpu/bert/attention_helper.h index c274dee7ea..cf9408a2db 100644 --- a/onnxruntime/contrib_ops/cpu/bert/attention_helper.h +++ b/onnxruntime/contrib_ops/cpu/bert/attention_helper.h @@ -64,6 +64,7 @@ template void PrepareMask(const int32_t* mask_index, const std::vector* mask_index_dims, T* mask_data, + bool is_unidirectional, int batch_size, int sequence_length, int past_sequence_length) { @@ -83,6 +84,18 @@ void PrepareMask(const int32_t* mask_index, for (int i = 0; i < batch_size * sequence_length * all_sequence_length; i++) { p_mask[i] = (mask_index[i] > 0) ? static_cast(0.0f) : static_cast(-10000.0f); } + + if (is_unidirectional) { + for (int b_i = 0; b_i < batch_size; b_i++) { + for (int s_i = 0; s_i < sequence_length - 1; s_i++) { + for (int m_i = past_sequence_length + s_i + 1; m_i < all_sequence_length; m_i++) { + p_mask[s_i * all_sequence_length + m_i] += static_cast(-10000.0f); + } + } + p_mask += sequence_length * all_sequence_length; + } + } + return; } @@ -122,6 +135,15 @@ void PrepareMask(const int32_t* mask_index, memcpy(p_mask + s_i * all_sequence_length, p_mask, all_sequence_length * sizeof(T)); } + // Apply unidirectional mask. + if (is_unidirectional) { + for (int s_i = 0; s_i < sequence_length - 1; s_i++) { + for (int m_i = past_sequence_length + s_i + 1; m_i < all_sequence_length; m_i++) { + p_mask[s_i * all_sequence_length + m_i] += static_cast(-10000.0f); + } + } + } + p_mask += sequence_length * all_sequence_length; } } diff --git a/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu b/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu index e9c41f7754..6b72c1e023 100644 --- a/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu +++ b/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu @@ -178,11 +178,10 @@ bool LaunchAttentionKernel( const void* past, void* present) { - // GPT-2 model is more sensitive on parity since error will accumulate in text generation. - // So use persistent softmax for GPT-2 model by default. - // For testing, environment variable ORT_TRANSFORMER_OPTIONS=1 or 2 could enable or disable it explicitly. + + // For testing, environment variable ORT_TRANSFORMER_OPTIONS=1 could enable persistent softmax const TransformerOptions* options = TransformerOptions::GetInstance(); - bool use_persistent_softmax = (is_unidirectional || options->IsPrecisionMode()) && !options->DisablePersistentSoftmax(); + bool use_persistent_softmax = options->IsPrecisionMode() && !options->DisablePersistentSoftmax(); if (element_size == 2) { return QkvToContext(prop, cublas, stream,