DMMHA: add unit tests; fix CPU, CUDA kernel (#22567)

### Description

Fixes:
(1) cpu kernel: applying scale before bias and mask like other MHA ops
(2) cpu kernel: correct offset during appending past to present.
(3) cuda kernel: apply mask if provided; fix output_qk offset.

Add DMMHA unit tests
This commit is contained in:
mindest 2024-11-02 22:05:56 +09:00 committed by GitHub
parent 2e4e221da8
commit 4ffc1ff3b4
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 402 additions and 388 deletions

View file

@ -77,7 +77,7 @@ class AttentionCPUBase : public AttentionBase {
// Convert mask from boolean (0/1) to float (mask_filter_value/0.0f).
// Merge padding mask with causal mask, and broadcast to 3D (BxSxT).
PrepareMask(mask_index_data, mask_index_dims, static_cast<T*>(mask_data),
causal, batch_size, sequence_length, past_sequence_length, mask_filter_value_);
causal, batch_size, sequence_length, kv_sequence_length, past_sequence_length, mask_filter_value_);
DUMP_CPU_TENSOR("Mask3D", static_cast<T*>(mask_data), batch_size, sequence_length, total_sequence_length);
}

View file

@ -120,9 +120,10 @@ void PrepareMask(const int32_t* mask_index,
bool causal,
int batch_size,
int sequence_length,
int kv_sequence_length,
int past_sequence_length,
float mask_filter_value) {
const int all_sequence_length = past_sequence_length + sequence_length;
const int all_sequence_length = past_sequence_length + kv_sequence_length;
// mask_data has been filled with 0, and its shape is BxSxT
T* p_mask = mask_data;

View file

@ -339,6 +339,7 @@ void DecoderMaskedMultiHeadAttention<T>::ComputeAttentionProbsWithBeams(
T* attention_probs_ptr = reinterpret_cast<T*>(attention_probs) + last_offset;
math::Dot<float, CPUMathUtil>(head_size, q_vec, K + i * head_size, attention_probs_ptr, nullptr);
*attention_probs_ptr *= scale;
// Apply the attention bias and mask
if (attn_bias_data != nullptr) {
*attention_probs_ptr += attn_bias_data[attn_bias_base_offset + past_sequence_length];
@ -348,7 +349,6 @@ void DecoderMaskedMultiHeadAttention<T>::ComputeAttentionProbsWithBeams(
if (is_masked) {
*attention_probs_ptr += mask_filter_value_;
}
*attention_probs_ptr *= scale;
}
{
@ -362,6 +362,8 @@ void DecoderMaskedMultiHeadAttention<T>::ComputeAttentionProbsWithBeams(
const T* past_k_vec = past_key_data + beam_batch_offset + beam_offset + j * head_size;
T* output = reinterpret_cast<T*>(attention_probs) + j + i * probs_matrix_size;
math::Dot<float, CPUMathUtil>(head_size, q_vec, past_k_vec, output, nullptr);
*output *= scale;
// Apply the attention bias and mask
if (attn_bias_data != nullptr) {
*output += attn_bias_data[attn_bias_base_offset + j];
@ -371,11 +373,11 @@ void DecoderMaskedMultiHeadAttention<T>::ComputeAttentionProbsWithBeams(
if (is_masked) {
*output += mask_filter_value_;
}
*output *= scale;
}
}
// Append current key to present key (past_present_share_buffer_ is true)
memcpy(present_key_data + i * max_sequence_length * head_size, K + i * head_size, head_size * sizeof(T));
memcpy(present_key_data + (i * max_sequence_length + past_sequence_length) * head_size,
K + i * head_size, head_size * sizeof(T));
}
});
@ -460,7 +462,7 @@ void DecoderMaskedMultiHeadAttention<T>::ComputeVxAttentionScoreWithBeams(
}
}
// Append current value to present value (past_present_share_buffer_ is true)
memcpy(present_value_data + i * max_sequence_length * v_head_size,
memcpy(present_value_data + (i * max_sequence_length + past_sequence_length) * v_head_size,
V + i * v_head_size,
v_head_size * sizeof(T));
}

View file

@ -33,7 +33,7 @@ class DecoderMaskedMultiHeadAttention final : public OpKernel, public AttentionC
const Tensor* cache_indir,
OpKernelContext* context,
int beam_width,
Tensor* scaled_qk = nullptr) const;
Tensor* output_qk = nullptr) const;
void ComputeAttentionProbsWithBeams(T* attention_probs,
const T* Q,
const T* K,
@ -50,7 +50,7 @@ class DecoderMaskedMultiHeadAttention final : public OpKernel, public AttentionC
bool broadcast_attn_bias_dim_1,
const int32_t* cache_indir_data,
int beam_width,
T* scaled_qk_data = nullptr) const;
T* output_qk_data = nullptr) const;
void ComputeVxAttentionScoreWithBeams(T* output,
T* tmp_buffer,
const T* attention_probs,

View file

@ -298,6 +298,9 @@ __global__ void masked_multihead_attention_kernel(DecoderMaskedMultiHeadAttentio
if (params.attention_bias != nullptr) {
qk = add_vec(qk, reinterpret_cast<T*>(params.attention_bias)[attn_bias_offset + tlength]);
}
if (params.mask != nullptr && params.mask[bi_total_seq_length + params.past_sequence_length] == 0) {
qk += params.mask_filter_value;
}
qk_max = qk;
qk_smem[tlength] = qk;
}
@ -534,7 +537,7 @@ __global__ void masked_multihead_attention_kernel(DecoderMaskedMultiHeadAttentio
if (params.out_qk != nullptr) {
// store cross qk before softmax, out_qk has shape [B(batchxbeam), #Head, 1, total_sequence_length]
float* target = ((float*)params.out_qk) + ((int64_t)bhi * tlength);
float* target = (reinterpret_cast<float*>(params.out_qk)) + (static_cast<int64_t>(bhi) * (sum_tlength + 1));
for (int ti = tidx; ti <= sum_tlength; ti += THREADS_PER_BLOCK) {
target[ti] = (float)(qk_smem[ti]);
}

View file

@ -908,7 +908,6 @@ ONNX_MS_OPERATOR_SET_SCHEMA(
OpSchema::Optional)
.Input(9,
"cache_indirection",
// This input is useful for CUDA EP only.
"A buffer of shape [batch_size, beam_width, max_output_length] where an `[i, j, k]` entry specifies "
"which beam the `k`-th token came from for the `j`-th beam for batch `i` in the current iteration",
"M",

View file

@ -15,23 +15,20 @@ namespace onnxruntime {
namespace test {
// This op is currently only supported on CUDA- so test it only for CUDA
#ifdef USE_CUDA
template <typename T>
static std::vector<T> CreateOnes(int size) {
std::vector<T> f;
f.reserve(size);
for (int i = 0; i < size; ++i) {
f.push_back(T(1));
f.push_back(T(1.0f));
}
return f;
}
template <typename T>
static std::vector<T> CreateValues(int size, int val) {
static std::vector<T> CreateValues(int size, float val) {
std::vector<T> f;
f.reserve(size);
@ -72,39 +69,25 @@ static std::vector<T> CreateRandom(int size) {
return f;
}
// QKV
template <typename T>
static std::vector<T> QKV(std::vector<T>& input, std::vector<T>& weights, std::vector<T>& bias,
int batch_size, int sequence_length, int hidden_size);
float ToFloat(T val);
template <>
std::vector<float> QKV(std::vector<float>& input, std::vector<float>& weights, std::vector<float>& bias,
int batch_size, int sequence_length, int hidden_size) {
std::vector<float> qkv;
qkv.resize(batch_size * sequence_length * 3 * hidden_size, 0);
for (int b = 0; b < batch_size; ++b) {
for (int i = 0; i < sequence_length; ++i) {
for (int j = 0; j < 3 * hidden_size; ++j) {
float sum = 0;
for (int k = 0; k < hidden_size; ++k) {
sum += input[b * sequence_length * hidden_size + i * hidden_size + k] * weights[k * 3 * hidden_size + j];
}
qkv[b * sequence_length * 3 * hidden_size + i * 3 * hidden_size + j] = sum + bias[j];
}
}
}
return qkv;
constexpr float ToFloat(float val) {
return val;
}
template <>
std::vector<MLFloat16> QKV(std::vector<MLFloat16>& input, std::vector<MLFloat16>& weights, std::vector<MLFloat16>& bias,
int batch_size, int sequence_length, int hidden_size) {
std::vector<MLFloat16> qkv;
qkv.resize(batch_size * sequence_length * 3 * hidden_size, static_cast<MLFloat16>(0.f));
float ToFloat(MLFloat16 val) {
return val.ToFloat();
}
// QKV
template <typename T>
static std::vector<T> QKV(std::vector<T>& input, std::vector<T>& weights, std::vector<T>& bias,
int batch_size, int sequence_length, int hidden_size) {
std::vector<T> qkv;
qkv.resize(batch_size * sequence_length * 3 * hidden_size, static_cast<T>(0.f));
for (int b = 0; b < batch_size; ++b) {
for (int i = 0; i < sequence_length; ++i) {
@ -112,10 +95,11 @@ std::vector<MLFloat16> QKV(std::vector<MLFloat16>& input, std::vector<MLFloat16>
float sum = 0;
for (int k = 0; k < hidden_size; ++k) {
sum += input[b * sequence_length * hidden_size + i * hidden_size + k].ToFloat() * weights[k * 3 * hidden_size + j].ToFloat();
sum += ToFloat(input[b * sequence_length * hidden_size + i * hidden_size + k]) *
ToFloat(weights[k * 3 * hidden_size + j]);
}
qkv[b * sequence_length * 3 * hidden_size + i * 3 * hidden_size + j] = static_cast<MLFloat16>(sum + bias[j].ToFloat());
qkv[b * sequence_length * 3 * hidden_size + i * 3 * hidden_size + j] = static_cast<T>(sum + ToFloat(bias[j]));
}
}
}
@ -180,15 +164,17 @@ void CheckEquality(T* data_1, T* data_2, int batch_size, int num_heads, int num_
// Reorder 'K' from [B, N, S, H] to [B, N, H/x, S, x] where x = (sizeof(T) / 16);
// Copy 'V' over as is
template <typename T>
static std::vector<T> ReorderKVCache(std::vector<T>& unordered_k_cache,
static std::vector<T> ReorderKVCache(const std::vector<T>& unordered_k_cache,
int batch_size, int num_heads, int sequence_length,
int head_size, int max_sequence_length) {
int head_size, int max_sequence_length, bool merge_past_kv = true) {
std::vector<T> ordered(unordered_k_cache.size(), T{0.f});
// Copy V over
size_t v_start = unordered_k_cache.size() / 2;
for (size_t i = v_start; i < unordered_k_cache.size(); ++i) {
ordered[i] = unordered_k_cache[i];
if (merge_past_kv) {
size_t v_start = unordered_k_cache.size() / 2;
for (size_t i = v_start; i < unordered_k_cache.size(); ++i) {
ordered[i] = unordered_k_cache[i];
}
}
// Now let us re-order K and copy it over to the final buffer
@ -203,7 +189,8 @@ static std::vector<T> ReorderKVCache(std::vector<T>& unordered_k_cache,
(h * max_sequence_length * head_size);
int input_base_offset = base_offset + (s * head_size) + (c * num_inner_elements);
int output_base_offset = base_offset + (c * max_sequence_length * num_inner_elements) + (s * num_inner_elements);
int output_base_offset = base_offset + (c * max_sequence_length * num_inner_elements) +
(s * num_inner_elements);
for (int e = 0; e < num_inner_elements; ++e) {
ordered[output_base_offset + e] = unordered_k_cache[input_base_offset + e];
@ -224,7 +211,7 @@ static std::vector<T> MergeReorderedKVCacheWithK(std::vector<T>& ordered_k_cache
T* k,
int batch_size, int num_heads,
int past_sequence_length, int max_sequence_length,
int head_size) {
int head_size, bool merge_past_kv = true) {
std::vector<T> merged = ordered_k_cache;
int total_seq_length = past_sequence_length + 1;
@ -249,10 +236,11 @@ static std::vector<T> MergeReorderedKVCacheWithK(std::vector<T>& ordered_k_cache
input_value = ordered_k_cache[input_offset];
} else {
int hidden_size = num_heads * head_size;
int input_offset = (b * 3 * hidden_size) +
(n * num_chunks * chunk_size) +
(c * chunk_size) +
h;
int input_offset = merge_past_kv ? ((b * 3 * hidden_size) +
(n * num_chunks * chunk_size) +
(c * chunk_size) +
h)
: ((b * hidden_size) + n * head_size + c * chunk_size + h);
input_value = k[input_offset];
}
@ -272,7 +260,7 @@ static std::vector<T> MergeReorderedKVCacheWithK(std::vector<T>& ordered_k_cache
return merged;
}
// GIven a pointer to the 'V' component of the past cache, we will merge it
// Given a pointer to the 'V' component of the past cache, we will merge it
// with current 'V' in-place
template <typename T>
static void MergeReorderedKVCacheWithV(T* v_cache,
@ -299,7 +287,8 @@ static void MergeReorderedKVCacheWithV(T* v_cache,
template <typename T>
static std::pair<std::vector<T>, std::vector<T>> MergePastKWithPresentKAndTranspose(T* past_k, T* present_k,
int num_batch, int num_heads,
int past_sequence_length, int max_sequence_length,
int past_sequence_length,
int max_sequence_length,
int head_size) {
int total_seq_length = (past_sequence_length + 1);
std::vector<T> merged_k(num_batch * num_heads * total_seq_length * head_size, T{0.f});
@ -312,16 +301,18 @@ static std::pair<std::vector<T>, std::vector<T>> MergePastKWithPresentKAndTransp
T input_value{0.f};
if (s < past_sequence_length) {
int input_offset = b * num_heads * max_sequence_length * head_size + (n * max_sequence_length * head_size) + (s * head_size) + h;
int input_offset = b * num_heads * max_sequence_length * head_size +
(n * max_sequence_length * head_size) + (s * head_size) + h;
input_value = past_k[input_offset];
} else {
int hidden_size = num_heads * head_size;
// Offset by 3* hidden_size because QKV data contains Q, K, and V per batch
// Offset by 3 * hidden_size because QKV data contains Q, K, and V per batch
int input_offset = (b * 3 * hidden_size) + (n * head_size) + h;
input_value = present_k[input_offset];
}
int output_offset = b * num_heads * total_seq_length * head_size + (n * total_seq_length * head_size) + (s * head_size) + h;
int output_offset = b * num_heads * total_seq_length * head_size +
(n * total_seq_length * head_size) + (s * head_size) + h;
merged_k[output_offset] = input_value;
}
@ -383,15 +374,11 @@ void ValidateReorderedMergedKWithK(T* k, T* k_cache, int batch_size, int num_hea
// QK_Transpose
template <typename T>
std::vector<T> QK_Transpose(T* q_matrix, T* k_transpose_matrix,
int batch_size, int num_heads, int total_sequence_length, int head_size);
template <>
std::vector<float> QK_Transpose(float* q_matrix, float* k_transpose_matrix,
int batch_size, int num_heads, int total_sequence_length, int head_size) {
int batch_size, int num_heads, int total_sequence_length, int head_size) {
int hidden_size = num_heads * head_size;
std::vector<float> qk_transpose;
qk_transpose.resize(batch_size * num_heads * total_sequence_length, 0);
std::vector<T> qk_transpose;
qk_transpose.resize(batch_size * num_heads * total_sequence_length, static_cast<T>(0.f));
for (int b = 0; b < batch_size; ++b) {
for (int n = 0; n < num_heads; ++n) {
@ -409,50 +396,12 @@ std::vector<float> QK_Transpose(float* q_matrix, float* k_transpose_matrix,
for (int j = 0; j < total_sequence_length; ++j) {
float sum = 0;
for (int k = 0; k < head_size; ++k) {
sum += (q_matrix[input_1_base_offset + i * head_size + k] *
k_transpose_matrix[input_2_base_offset + k * total_sequence_length + j]);
sum += (ToFloat(q_matrix[input_1_base_offset + i * head_size + k]) *
ToFloat(k_transpose_matrix[input_2_base_offset + k * total_sequence_length + j]));
}
float scale = 1 / sqrt(static_cast<float>(head_size));
qk_transpose[output_base_offset + i * total_sequence_length + j] = scale * sum;
}
}
}
}
return qk_transpose;
}
template <>
std::vector<MLFloat16> QK_Transpose(MLFloat16* q_matrix, MLFloat16* k_transpose_matrix,
int batch_size, int num_heads, int total_sequence_length, int head_size) {
int hidden_size = num_heads * head_size;
std::vector<MLFloat16> qk_transpose;
qk_transpose.resize(batch_size * num_heads * total_sequence_length, MLFloat16(0.f));
for (int b = 0; b < batch_size; ++b) {
for (int n = 0; n < num_heads; ++n) {
int input_1_base_offset = (b * 3 * hidden_size) +
(n * head_size);
int input_2_base_offset = (b * num_heads * total_sequence_length * head_size) +
(n * total_sequence_length * head_size);
int output_base_offset = (b * num_heads * total_sequence_length) +
(n * total_sequence_length);
// sequence_length == 1
for (int i = 0; i < 1; ++i) {
for (int j = 0; j < total_sequence_length; ++j) {
float sum = 0;
for (int k = 0; k < head_size; ++k) {
sum += (q_matrix[input_1_base_offset + i * head_size + k].ToFloat() *
k_transpose_matrix[input_2_base_offset + k * total_sequence_length + j].ToFloat());
}
float scale = 1 / sqrt(static_cast<float>(head_size));
qk_transpose[output_base_offset + i * total_sequence_length + j] = MLFloat16(scale * sum);
qk_transpose[output_base_offset + i * total_sequence_length + j] = static_cast<T>(scale * sum);
}
}
}
@ -464,26 +413,23 @@ std::vector<MLFloat16> QK_Transpose(MLFloat16* q_matrix, MLFloat16* k_transpose_
// Softmax_QK_Transpose
template <typename T>
std::vector<T> Softmax_QK_Transpose(T* qk_transpose_matrix, int batch_size, int num_heads,
int sequence_length, int total_sequence_length, int head_size);
template <>
std::vector<float> Softmax_QK_Transpose(float* qk_transpose_matrix, int batch_size, int num_heads,
int sequence_length, int total_sequence_length, int /*head_size*/) {
int sequence_length, int total_sequence_length) {
if (sequence_length != 1) {
throw std::runtime_error("Not supported");
}
std::vector<float> softmax_qk_transpose;
softmax_qk_transpose.resize(batch_size * num_heads * sequence_length * total_sequence_length, 0);
std::vector<T> softmax_qk_transpose;
softmax_qk_transpose.resize(static_cast<size_t>(batch_size) * num_heads * sequence_length * total_sequence_length,
static_cast<T>(0.f));
for (int b = 0; b < batch_size; ++b) {
for (int n = 0; n < num_heads; ++n) {
int base_offset = (b * num_heads * sequence_length * total_sequence_length) +
(n * sequence_length * total_sequence_length);
float max = std::numeric_limits<float>::min();
float max = std::numeric_limits<float>::lowest();
for (int s = 0; s < total_sequence_length; ++s) {
auto val = qk_transpose_matrix[base_offset + s];
auto val = ToFloat(qk_transpose_matrix[base_offset + s]);
if (val > max) {
max = val;
}
@ -491,52 +437,13 @@ std::vector<float> Softmax_QK_Transpose(float* qk_transpose_matrix, int batch_si
float denom = 0;
for (int s = 0; s < total_sequence_length; ++s) {
auto val = qk_transpose_matrix[base_offset + s];
auto val = ToFloat(qk_transpose_matrix[base_offset + s]);
denom += std::exp(val - max);
}
for (int s = 0; s < total_sequence_length; ++s) {
auto val = qk_transpose_matrix[base_offset + s];
softmax_qk_transpose[base_offset + s] = std::exp(val - max) / (denom + (float)0.000001);
}
}
}
return softmax_qk_transpose;
}
template <>
std::vector<MLFloat16> Softmax_QK_Transpose(MLFloat16* qk_transpose_matrix, int batch_size, int num_heads,
int sequence_length, int total_sequence_length, int /*head_size*/) {
if (sequence_length != 1) {
throw std::runtime_error("Not supported");
}
std::vector<MLFloat16> softmax_qk_transpose;
softmax_qk_transpose.resize(batch_size * num_heads * sequence_length * total_sequence_length, MLFloat16(0.f));
for (int b = 0; b < batch_size; ++b) {
for (int n = 0; n < num_heads; ++n) {
int base_offset = (b * num_heads * sequence_length * total_sequence_length) +
(n * sequence_length * total_sequence_length);
float max = std::numeric_limits<float>::min();
for (int s = 0; s < total_sequence_length; ++s) {
auto val = qk_transpose_matrix[base_offset + s].ToFloat();
if (val > max) {
max = val;
}
}
float denom = 0;
for (int s = 0; s < total_sequence_length; ++s) {
auto val = qk_transpose_matrix[base_offset + s].ToFloat();
denom += std::exp(val - max);
}
for (int s = 0; s < total_sequence_length; ++s) {
auto val = qk_transpose_matrix[base_offset + s].ToFloat();
softmax_qk_transpose[base_offset + s] = MLFloat16(std::exp(val - max) / (denom + (float)0.000001));
auto val = ToFloat(qk_transpose_matrix[base_offset + s]);
softmax_qk_transpose[base_offset + s] = static_cast<T>(std::exp(val - max) / (denom + (float)0.000001));
}
}
}
@ -550,19 +457,13 @@ std::vector<T> Softmax_QK_Transpose_V(T* softmax_qk_transpose_matrix,
T* v_matrix,
int batch_size, int num_heads, int sequence_length,
int total_sequence_length, int max_sequence_length,
int head_size);
template <>
std::vector<float> Softmax_QK_Transpose_V(float* softmax_qk_transpose_matrix,
float* v_matrix,
int batch_size, int num_heads, int sequence_length,
int total_sequence_length, int max_sequence_length,
int head_size) {
int head_size) {
if (sequence_length != 1) {
throw std::runtime_error("Not supported");
}
std::vector<float> output;
output.resize(batch_size * sequence_length * num_heads * head_size, 0);
std::vector<T> output;
output.resize(batch_size * sequence_length * num_heads * head_size, static_cast<T>(0.f));
for (int b = 0; b < batch_size; ++b) {
for (int n = 0; n < num_heads; ++n) {
@ -580,11 +481,11 @@ std::vector<float> Softmax_QK_Transpose_V(float* softmax_qk_transpose_matrix,
float sum = 0;
for (int k = 0; k < total_sequence_length; ++k) {
sum += (softmax_qk_transpose_matrix[input_1_base_offset + i * total_sequence_length + k] *
v_matrix[input_2_base_offset + k * head_size + j]);
sum += (ToFloat(softmax_qk_transpose_matrix[input_1_base_offset + i * total_sequence_length + k]) *
ToFloat(v_matrix[input_2_base_offset + k * head_size + j]));
}
output[output_base_offset + i * head_size + j] = sum;
output[output_base_offset + i * head_size + j] = static_cast<T>(sum);
}
}
}
@ -593,48 +494,11 @@ std::vector<float> Softmax_QK_Transpose_V(float* softmax_qk_transpose_matrix,
return output;
}
template <>
std::vector<MLFloat16> Softmax_QK_Transpose_V(MLFloat16* softmax_qk_transpose_matrix,
MLFloat16* v_matrix,
int batch_size, int num_heads, int sequence_length,
int total_sequence_length, int max_sequence_length,
int head_size) {
if (sequence_length != 1) {
throw std::runtime_error("Not supported");
}
// Currently we only support CUDA for DecoderMaskedSelfAttention
#ifdef USE_CUDA
std::vector<MLFloat16> output;
output.resize(batch_size * sequence_length * num_heads * head_size, MLFloat16(0.f));
for (int b = 0; b < batch_size; ++b) {
for (int n = 0; n < num_heads; ++n) {
int input_1_base_offset = (b * num_heads * sequence_length * total_sequence_length) +
(n * sequence_length * total_sequence_length);
int input_2_base_offset = (b * num_heads * max_sequence_length * head_size) +
(n * max_sequence_length * head_size);
int output_base_offset = (b * num_heads * sequence_length * head_size) +
(n * sequence_length * head_size);
for (int i = 0; i < sequence_length; ++i) {
for (int j = 0; j < head_size; ++j) {
float sum = 0;
for (int k = 0; k < total_sequence_length; ++k) {
sum += (softmax_qk_transpose_matrix[input_1_base_offset + i * total_sequence_length + k].ToFloat() *
v_matrix[input_2_base_offset + k * head_size + j].ToFloat());
}
output[output_base_offset + i * head_size + j] = MLFloat16(sum);
}
}
}
}
return output;
}
TEST(DecoderMaskedSelfAttentionTest, Test_fp32) {
template <typename T>
static void TestDecoderMaskedSelfAttention() {
// The kernel is only supported on CC 5.3 or higher GPUs
if (NeedSkipIfCudaArchLowerThan(530)) {
return;
@ -661,19 +525,19 @@ TEST(DecoderMaskedSelfAttentionTest, Test_fp32) {
};
constexpr int sequence_length = 1;
constexpr int number_of_heads = 12;
constexpr int num_heads = 12;
for (MyTestCase test_case : test_cases) {
int batch_size = test_case.batch_size;
int past_sequence_length = test_case.past_sequence_length;
int hidden_size = test_case.hidden_size;
int head_size = (hidden_size / number_of_heads);
int head_size = (hidden_size / num_heads);
int total_sequence_length = sequence_length + past_sequence_length;
int max_sequence_length = past_sequence_length + 1; // Always keep > past_sequence_length
int max_sequence_length = past_sequence_length + 1; // Always keep > past_sequence_length
OpTester tester("DecoderMaskedSelfAttention", 1, onnxruntime::kMSDomain);
tester.AddAttribute<int64_t>("num_heads", static_cast<int64_t>(number_of_heads));
tester.AddAttribute<int64_t>("num_heads", static_cast<int64_t>(num_heads));
tester.AddAttribute<int64_t>("past_present_share_buffer", static_cast<int64_t>(1));
std::vector<int64_t> input_dims = {batch_size, sequence_length, hidden_size};
@ -681,38 +545,38 @@ TEST(DecoderMaskedSelfAttentionTest, Test_fp32) {
std::vector<int64_t> bias_dims = {3 * hidden_size};
std::vector<int64_t> output_dims = {batch_size, sequence_length, hidden_size};
auto input = CreateRandom<float>(batch_size * sequence_length * hidden_size);
tester.AddInput<float>("input", input_dims, input);
auto input = CreateRandom<T>(batch_size * sequence_length * hidden_size);
tester.AddInput<T>("input", input_dims, input);
auto weight = CreateRandom<float>(hidden_size * 3 * hidden_size);
tester.AddInput<float>("weight", weights_dims, weight);
auto weight = CreateRandom<T>(hidden_size * 3 * hidden_size);
tester.AddInput<T>("weight", weights_dims, weight);
auto bias = CreateRandom<float>(3 * hidden_size);
tester.AddInput<float>("bias", bias_dims, bias);
auto bias = CreateRandom<T>(3 * hidden_size);
tester.AddInput<T>("bias", bias_dims, bias);
// Mask
tester.AddOptionalInputEdge<int32_t>();
// Past
std::vector<int64_t> past_dims = {2, batch_size, number_of_heads, max_sequence_length, head_size};
int past_present_size = 2 * batch_size * number_of_heads * max_sequence_length * head_size;
std::vector<int64_t> past_dims = {2, batch_size, num_heads, max_sequence_length, head_size};
int past_present_size = 2 * batch_size * num_heads * max_sequence_length * head_size;
auto kv_cache = CreateRandom<float>(past_present_size);
auto kv_cache = CreateRandom<T>(past_present_size);
auto reordered_kv_cache = ReorderKVCache<float>(kv_cache, batch_size,
number_of_heads, past_sequence_length, head_size, max_sequence_length);
auto reordered_kv_cache = ReorderKVCache<T>(kv_cache, batch_size,
num_heads, past_sequence_length, head_size, max_sequence_length);
// Validate if reordering went well - by transposing and checking equality
int chunk_size = 16 / sizeof(float);
int chunk_size = 16 / sizeof(T);
int num_chunks = head_size / chunk_size;
auto transposed = Transpose<float>(kv_cache.data(), batch_size, number_of_heads, num_chunks, max_sequence_length, chunk_size);
CheckEquality<float>(transposed.data(), reordered_kv_cache.data(), batch_size, number_of_heads, num_chunks,
max_sequence_length, past_sequence_length, chunk_size);
auto transposed = Transpose<T>(kv_cache.data(), batch_size, num_heads, num_chunks, max_sequence_length, chunk_size);
CheckEquality<T>(transposed.data(), reordered_kv_cache.data(), batch_size, num_heads, num_chunks,
max_sequence_length, past_sequence_length, chunk_size);
tester.AddInput<float>("past", past_dims, reordered_kv_cache);
tester.AddInput<T>("past", past_dims, reordered_kv_cache);
// Rel
tester.AddOptionalInputEdge<float>();
tester.AddOptionalInputEdge<T>();
// Past sequence length
std::vector<int32_t> arr_past_sequence_len(1, past_sequence_length);
@ -722,41 +586,44 @@ TEST(DecoderMaskedSelfAttentionTest, Test_fp32) {
auto qkv = QKV(input, weight, bias, batch_size, sequence_length, hidden_size);
auto* qkv_matrix = qkv.data();
auto pair = MergePastKWithPresentKAndTranspose<float>(kv_cache.data(), qkv_matrix + hidden_size, batch_size,
number_of_heads, past_sequence_length,
max_sequence_length, head_size);
auto pair = MergePastKWithPresentKAndTranspose<T>(kv_cache.data(), qkv_matrix + hidden_size, batch_size, num_heads,
past_sequence_length, max_sequence_length, head_size);
auto k_merged = pair.first;
auto k_transpose = pair.second;
auto qk_transpose = QK_Transpose<float>(qkv_matrix, k_transpose.data(), batch_size, number_of_heads,
total_sequence_length, head_size);
auto qk_transpose = QK_Transpose<T>(qkv_matrix, k_transpose.data(), batch_size, num_heads,
total_sequence_length, head_size);
auto softmax_qk_transpose = Softmax_QK_Transpose<float>(qk_transpose.data(), batch_size, number_of_heads,
sequence_length, total_sequence_length, head_size);
auto softmax_qk_transpose = Softmax_QK_Transpose<T>(qk_transpose.data(), batch_size, num_heads,
sequence_length, total_sequence_length);
auto present = MergeReorderedKVCacheWithK<float>(reordered_kv_cache, qkv_matrix + hidden_size, batch_size,
number_of_heads, past_sequence_length, max_sequence_length, head_size);
auto present = MergeReorderedKVCacheWithK<T>(reordered_kv_cache, qkv_matrix + hidden_size, batch_size,
num_heads, past_sequence_length, max_sequence_length, head_size);
// Validate our test logic
// We want to validate if our merged "unordered" K is the same as
// the merged "ordered" K so that the QKT we do in our test code
// is equivalent to the QKT we do in the kernel
ValidateReorderedMergedKWithK<float>(k_merged.data(), present.data(), batch_size, number_of_heads, total_sequence_length, max_sequence_length, head_size);
ValidateReorderedMergedKWithK<T>(k_merged.data(), present.data(), batch_size, num_heads, total_sequence_length,
max_sequence_length, head_size);
MergeReorderedKVCacheWithV<float>(present.data() + (past_present_size / 2), qkv_matrix + 2 * hidden_size, batch_size,
number_of_heads, past_sequence_length, max_sequence_length, head_size);
MergeReorderedKVCacheWithV<T>(present.data() + (past_present_size / 2), qkv_matrix + 2 * hidden_size, batch_size,
num_heads, past_sequence_length, max_sequence_length, head_size);
auto output = Softmax_QK_Transpose_V<float>(softmax_qk_transpose.data(), present.data() + (past_present_size / 2),
batch_size, number_of_heads,
sequence_length, total_sequence_length,
max_sequence_length, head_size);
auto output = Softmax_QK_Transpose_V<T>(softmax_qk_transpose.data(), present.data() + (past_present_size / 2),
batch_size, num_heads, sequence_length, total_sequence_length,
max_sequence_length, head_size);
// Output(s)
tester.AddOutput<float>("output", input_dims, output);
tester.AddOutput<float>("present", past_dims, present);
tester.AddOutput<T>("output", input_dims, output);
tester.AddOutput<T>("present", past_dims, present);
tester.SetOutputTolerance(0.001f, 0.001f);
if (std::is_same<T, MLFloat16>::value) {
tester.SetOutputTolerance(0.005f);
} else {
tester.SetOutputTolerance(0.001f, 0.001f);
}
// Run - Regular kernel execution path
{
@ -778,150 +645,292 @@ TEST(DecoderMaskedSelfAttentionTest, Test_fp32) {
}
}
#endif // USE_CUDA
template <typename T>
static std::vector<T> CalculateOutputQK(const std::vector<T>& q, const std::vector<T>& k,
const std::vector<int32_t>& mask_index, const std::vector<T>& attention_bias,
int batch_size, int num_heads,
int sequence_length, int max_sequence_length, int head_size) {
// q (B, 1, NH), k (B, N, L(M), H) -> qk (B, N, 1, L)
// mask_index (B, L), (optional) attention_bias (1, 1, 1, L)
float scale = 1 / sqrt(static_cast<float>(head_size));
std::vector<T> output_qk;
output_qk.resize(static_cast<size_t>(batch_size) * num_heads * sequence_length, static_cast<T>(0.f));
for (int b = 0; b < batch_size; ++b) {
for (int n = 0; n < num_heads; ++n) {
for (int s = 0; s < sequence_length; ++s) {
float mask_value = (mask_index[b * sequence_length + s] == 0) ? -10000.f : 0.f;
float bias_value = (attention_bias.empty()) ? 0.f : ToFloat(attention_bias[s]);
float sum = 0;
for (int h = 0; h < head_size; ++h) {
sum += ToFloat(q[b * num_heads * head_size + n * head_size + h]) *
ToFloat(k[b * num_heads * max_sequence_length * head_size +
n * max_sequence_length * head_size + s * head_size + h]);
}
output_qk[b * num_heads * sequence_length + n * sequence_length + s] =
static_cast<T>(scale * sum + mask_value + bias_value);
}
}
}
return output_qk;
}
template <typename T>
static std::vector<T> CalculateOutput(const std::vector<T>& softmax, const std::vector<T>& v, int batch_size,
int num_heads, int sequence_length, int max_sequence_length, int head_size) {
// softmax (B, N, 1, L) v (B, N, L(M), H) -> output (B, N, 1, H)
std::vector<T> output;
output.resize(static_cast<size_t>(batch_size) * num_heads * head_size, static_cast<T>(0.f));
for (int b = 0; b < batch_size; ++b) {
for (int n = 0; n < num_heads; ++n) {
for (int h = 0; h < head_size; ++h) {
float sum = 0;
for (int s = 0; s < sequence_length; ++s) {
sum += ToFloat(softmax[b * num_heads * sequence_length + n * sequence_length + s]) *
ToFloat(v[b * num_heads * max_sequence_length * head_size +
n * max_sequence_length * head_size + s * head_size + h]);
}
output[b * num_heads * head_size + n * head_size + h] = static_cast<T>(sum);
}
}
}
return output;
}
template <typename T>
static std::vector<T> MergePast(const std::vector<T>& past, const std::vector<T>& current, int batch_size,
int num_heads, int past_seq_len, int max_seq_len, int head_size) {
// past (B, N, S(M), H), current (B, 1, NH) -> merged (B, N, S+1(M), H)
std::vector<T> merged = past;
for (int b = 0; b < batch_size; ++b) {
for (int n = 0; n < num_heads; ++n) {
for (int h = 0; h < head_size; ++h) {
merged[b * num_heads * max_seq_len * head_size + n * max_seq_len * head_size + past_seq_len * head_size + h] =
current[b * num_heads * head_size + n * head_size + h];
}
}
}
return merged;
}
template <typename T>
static std::vector<T> ReorderKVByCacheIndirection(const std::vector<T>& key_or_value,
const int32_t* cache_indirection,
int batch_size, int beam_width, int max_sequence_length,
int num_heads, int head_size, int past_sequence_length) {
std::vector<T> reordered = key_or_value;
for (int b = 0; b < batch_size; ++b) {
int beam_batch_index = b / beam_width;
const int* beam_indices = cache_indirection + b * max_sequence_length;
for (int n = 0; n < num_heads; ++n) {
for (int s = 0; s < past_sequence_length; ++s) {
int beam_offset = beam_indices[s] * num_heads * max_sequence_length * head_size;
int beam_batch_offset = (beam_batch_index * beam_width * num_heads + n) * max_sequence_length * head_size;
for (int h = 0; h < head_size; ++h) {
reordered[b * num_heads * max_sequence_length * head_size +
n * max_sequence_length * head_size + s * head_size + h] =
key_or_value[beam_offset + beam_batch_offset + s * head_size + h];
}
}
}
}
return reordered;
}
template <typename T>
static void TestDecoderMaskedMultiHeadAttention(bool is_cross_attn = true, bool use_cuda = true) {
int batch_size = 8;
int past_sequence_length = 2;
int kv_sequence_length = 16;
int head_size = 32;
int num_heads = 12;
int beam_width = 4;
int hidden_size = head_size * num_heads;
OpTester tester("DecoderMaskedMultiHeadAttention", 1, onnxruntime::kMSDomain);
FixedPatternValueGenerator generator{};
RandomValueGenerator random{};
// Attributes
tester.AddAttribute<int64_t>("num_heads", static_cast<int64_t>(num_heads));
tester.AddAttribute<int64_t>("past_present_share_buffer", static_cast<int64_t>(!is_cross_attn));
// Output scaled Q * K^T by default for cross-attention
tester.AddAttribute<int64_t>("output_qk", static_cast<int64_t>(is_cross_attn));
// Inputs and outputs
auto query = CreateRandom<T>(batch_size * 1 * hidden_size);
tester.AddInput<T>("query", {batch_size, 1, hidden_size}, query);
if (is_cross_attn) {
auto key = CreateRandom<T>(batch_size * num_heads * kv_sequence_length * head_size);
std::vector<T> reordered_key;
if (use_cuda) {
reordered_key = ReorderKVCache<T>(key, batch_size, num_heads,
kv_sequence_length, head_size, kv_sequence_length, false);
}
auto value = CreateRandom<T>(batch_size * num_heads * kv_sequence_length * head_size);
tester.AddInput<T>("key", {batch_size, num_heads, kv_sequence_length, head_size}, (use_cuda ? reordered_key : key));
tester.AddInput<T>("value", {batch_size, num_heads, kv_sequence_length, head_size},
CreateRandom<T>(batch_size * num_heads * kv_sequence_length * head_size));
const std::vector<int64_t> mask_index_dims = {batch_size, kv_sequence_length};
auto mask_index = generator.Discrete<int32_t>(mask_index_dims, AsSpan({0, 1}));
tester.AddInput<int32_t>("mask_index", {batch_size, kv_sequence_length}, mask_index);
// Calculate Softmax(Q * K^T + (Optional) mask) * V
std::vector<T> empty_attention_bias;
auto output_qk = CalculateOutputQK(query, key, mask_index, empty_attention_bias, batch_size, num_heads,
kv_sequence_length, kv_sequence_length, head_size);
std::vector<float> output_qk_float(output_qk.size());
for (size_t i = 0; i < output_qk.size(); ++i) {
output_qk_float[i] = static_cast<float>(output_qk[i]);
}
auto softmax = Softmax_QK_Transpose<T>(output_qk.data(), batch_size, num_heads, 1, kv_sequence_length);
auto output = CalculateOutput<T>(softmax, value, batch_size, num_heads,
kv_sequence_length, kv_sequence_length, head_size);
tester.AddOutput<T>("output", {batch_size, 1, hidden_size}, output);
tester.AddOptionalOutputEdge<T>(); // optional present_key
tester.AddOptionalOutputEdge<T>(); // optional present_value
tester.AddOutput<float>("qk", {batch_size, num_heads, 1, kv_sequence_length}, output_qk_float);
} else {
int max_sequence_length = past_sequence_length + 10;
int total_sequence_length = past_sequence_length + 1;
auto key = CreateRandom<T>(batch_size * hidden_size);
auto value = CreateRandom<T>(batch_size * hidden_size);
tester.AddInput<T>("key", {batch_size, 1, hidden_size}, key);
tester.AddInput<T>("value", {batch_size, 1, hidden_size}, value);
const std::vector<int64_t> mask_index_dims = {batch_size, total_sequence_length};
auto mask_index = generator.Discrete<int32_t>(mask_index_dims, AsSpan({0, 1}));
tester.AddInput<int32_t>("mask_index", {batch_size, total_sequence_length}, mask_index);
std::vector<int64_t> attention_bias_dims = {1, 1, 1, total_sequence_length};
auto attention_bias_float = random.Gaussian<float>(attention_bias_dims, 0.0f, 0.3f);
std::vector<T> attention_bias(attention_bias_float.size());
for (size_t i = 0; i < attention_bias.size(); ++i) {
attention_bias[i] = static_cast<T>(attention_bias_float[i]);
}
tester.AddInput<T>("attention_bias", {1, 1, 1, total_sequence_length}, attention_bias);
auto past_key = CreateRandom<T>(batch_size * num_heads * max_sequence_length * head_size);
auto past_value = CreateRandom<T>(batch_size * num_heads * max_sequence_length * head_size);
std::vector<T> reordered_past_key; // For CUDA, we need to reorder past key
if (use_cuda) {
reordered_past_key = ReorderKVCache<T>(past_key, batch_size, num_heads,
past_sequence_length, head_size, max_sequence_length, false);
}
tester.AddInput<T>("past_key", {batch_size, num_heads, max_sequence_length, head_size},
(use_cuda ? reordered_past_key : past_key));
tester.AddInput<T>("past_value", {batch_size, num_heads, max_sequence_length, head_size}, past_value);
// merge past key and value with current key and value
auto merged_key = MergePast<T>(past_key, key, batch_size, num_heads,
past_sequence_length, max_sequence_length, head_size);
std::vector<T> merged_reordered_key;
if (use_cuda) {
merged_reordered_key = MergeReorderedKVCacheWithK<T>(reordered_past_key, key.data(), batch_size, num_heads,
past_sequence_length, max_sequence_length, head_size, false);
}
auto merged_value = MergePast<T>(past_value, value, batch_size, num_heads,
past_sequence_length, max_sequence_length, head_size);
tester.AddInput<int32_t>("past_sequence_length", {1}, {past_sequence_length});
std::vector<T> mod_merged_key, mod_merged_value;
if (beam_width > 1) {
tester.AddInput<int32_t>("beam_width", {1}, {beam_width});
const std::vector<int64_t> cache_indir_dims = {batch_size, beam_width, max_sequence_length};
auto value_candidates = ValueRange<int32_t>(beam_width);
auto cache_indir = generator.Discrete<int32_t>(cache_indir_dims, value_candidates);
tester.AddInput<int32_t>("cache_indirection", cache_indir_dims, cache_indir);
// Modify merged_key and merged_value according to cache_indirection
mod_merged_key = ReorderKVByCacheIndirection<T>(merged_key, cache_indir.data(),
batch_size, beam_width, max_sequence_length,
num_heads, head_size, past_sequence_length);
mod_merged_value = ReorderKVByCacheIndirection<T>(merged_value, cache_indir.data(),
batch_size, beam_width, max_sequence_length,
num_heads, head_size, past_sequence_length);
}
// Calculate Softmax(Q * K^T + (Optional) mask) * V
auto output_qk = CalculateOutputQK<T>(query, (beam_width > 1 ? mod_merged_key : merged_key),
mask_index, attention_bias,
batch_size, num_heads, total_sequence_length, max_sequence_length, head_size);
auto softmax = Softmax_QK_Transpose<T>(output_qk.data(), batch_size, num_heads, 1, total_sequence_length);
auto output = CalculateOutput<T>(softmax, (beam_width > 1 ? mod_merged_value : merged_value),
batch_size, num_heads, total_sequence_length, max_sequence_length, head_size);
tester.AddOutput<T>("output", {batch_size, 1, hidden_size}, output);
tester.AddOutput<T>("present_key", {batch_size, num_heads, max_sequence_length, head_size},
(use_cuda ? merged_reordered_key : merged_key));
tester.AddOutput<T>("present_value", {batch_size, num_heads, max_sequence_length, head_size}, merged_value);
}
if (std::is_same<T, MLFloat16>::value) {
tester.SetOutputTolerance(0.02f);
} else {
tester.SetOutputTolerance(0.0001f, 0.0001f);
}
{
std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
if (use_cuda) {
execution_providers.push_back(DefaultCudaExecutionProvider());
} else {
execution_providers.push_back(DefaultCpuExecutionProvider());
}
tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
}
}
#ifdef USE_CUDA
TEST(DecoderMaskedSelfAttentionTest, Test_fp32) {
TestDecoderMaskedSelfAttention<float>();
}
TEST(DecoderMaskedSelfAttentionTest, Test_fp16) {
// The kernel is only supported on CC 5.3 or higher GPUs
if (NeedSkipIfCudaArchLowerThan(530)) {
return;
}
TestDecoderMaskedSelfAttention<MLFloat16>();
}
// Buckets for test data:
// batch_size: 1, >=2
// past_sequence_length 0, 1~30, 31~2046, >=2047 (so that total_sequence_length: 1, 2-31, 32~2047, >=2048)
// head_size: 32, 64, 128
struct MyTestCase {
int batch_size;
int past_sequence_length;
int hidden_size;
} test_cases[] = {
{1, 0, 768},
{1, 1, 768},
{3, 30, 384},
{8, 31, 1536},
{4, 256, 384},
{3, 1024, 768},
{2, 2046, 1536},
{1, 2047, 384},
{2, 3000, 768},
};
TEST(DecoderMaskedMultiHeadAttentionTest, cuda_cross_attn_fp32) {
TestDecoderMaskedMultiHeadAttention<float>();
}
constexpr int sequence_length = 1;
constexpr int number_of_heads = 12;
TEST(DecoderMaskedMultiHeadAttentionTest, cuda_cross_attn_fp16) {
TestDecoderMaskedMultiHeadAttention<MLFloat16>();
}
for (MyTestCase test_case : test_cases) {
int batch_size = test_case.batch_size;
int past_sequence_length = test_case.past_sequence_length;
int hidden_size = test_case.hidden_size;
TEST(DecoderMaskedMultiHeadAttentionTest, cuda_self_attn_fp32) {
TestDecoderMaskedMultiHeadAttention<float>(/* is_cross_attn = */ false);
}
int head_size = (hidden_size / number_of_heads);
int total_sequence_length = sequence_length + past_sequence_length;
int max_sequence_length = past_sequence_length + 1; // Always keep > past_sequence_length
OpTester tester("DecoderMaskedSelfAttention", 1, onnxruntime::kMSDomain);
tester.AddAttribute<int64_t>("num_heads", static_cast<int64_t>(number_of_heads));
tester.AddAttribute<int64_t>("past_present_share_buffer", static_cast<int64_t>(1));
std::vector<int64_t> input_dims = {batch_size, sequence_length, hidden_size};
std::vector<int64_t> weights_dims = {hidden_size, 3 * hidden_size};
std::vector<int64_t> bias_dims = {3 * hidden_size};
std::vector<int64_t> output_dims = {batch_size, sequence_length, hidden_size};
auto input = CreateRandom<MLFloat16>(batch_size * sequence_length * hidden_size);
tester.AddInput<MLFloat16>("input", input_dims, input);
auto weight = CreateRandom<MLFloat16>(hidden_size * 3 * hidden_size);
tester.AddInput<MLFloat16>("weight", weights_dims, weight);
auto bias = CreateRandom<MLFloat16>(3 * hidden_size);
tester.AddInput<MLFloat16>("bias", bias_dims, bias);
// Mask
tester.AddOptionalInputEdge<int32_t>();
// Past
std::vector<int64_t> past_dims = {2, batch_size, number_of_heads, max_sequence_length, head_size};
int past_present_size = 2 * batch_size * number_of_heads * max_sequence_length * head_size;
auto kv_cache = CreateRandom<MLFloat16>(past_present_size);
auto reordered_kv_cache = ReorderKVCache<MLFloat16>(kv_cache, batch_size,
number_of_heads, past_sequence_length, head_size, max_sequence_length);
// Validate if reordering went well - by transposing and checking equality
int chunk_size = 16 / sizeof(MLFloat16);
int num_chunks = head_size / chunk_size;
auto transposed = Transpose<MLFloat16>(kv_cache.data(), batch_size, number_of_heads, num_chunks, max_sequence_length, chunk_size);
CheckEquality<MLFloat16>(transposed.data(), reordered_kv_cache.data(), batch_size, number_of_heads, num_chunks,
max_sequence_length, past_sequence_length, chunk_size);
tester.AddInput<MLFloat16>("past", past_dims, reordered_kv_cache);
// Rel
tester.AddOptionalInputEdge<MLFloat16>();
// Past sequence length
std::vector<int32_t> arr_past_sequence_len(1, past_sequence_length);
tester.AddInput<int32_t>("past_sequence_length", {1}, arr_past_sequence_len);
// QKV MatMul
auto qkv = QKV(input, weight, bias, batch_size, sequence_length, hidden_size);
auto* qkv_matrix = qkv.data();
auto pair = MergePastKWithPresentKAndTranspose<MLFloat16>(kv_cache.data(), qkv_matrix + hidden_size, batch_size,
number_of_heads, past_sequence_length,
max_sequence_length, head_size);
auto k_merged = pair.first;
auto k_transpose = pair.second;
auto qk_transpose = QK_Transpose<MLFloat16>(qkv_matrix, k_transpose.data(), batch_size, number_of_heads,
total_sequence_length, head_size);
auto softmax_qk_transpose = Softmax_QK_Transpose<MLFloat16>(qk_transpose.data(), batch_size, number_of_heads,
sequence_length, total_sequence_length, head_size);
auto present = MergeReorderedKVCacheWithK<MLFloat16>(reordered_kv_cache, qkv_matrix + hidden_size, batch_size,
number_of_heads, past_sequence_length, max_sequence_length, head_size);
// Validate our test logic
// We want to validate if our merged "unordered" K is the same as
// the merged "ordered" K so that the QKT we do in our test code
// is equivalent to the QKT we do in the kernel
ValidateReorderedMergedKWithK<MLFloat16>(k_merged.data(), present.data(), batch_size, number_of_heads, total_sequence_length, max_sequence_length, head_size);
MergeReorderedKVCacheWithV<MLFloat16>(present.data() + (past_present_size / 2), qkv_matrix + 2 * hidden_size, batch_size,
number_of_heads, past_sequence_length, max_sequence_length, head_size);
auto output = Softmax_QK_Transpose_V(softmax_qk_transpose.data(), present.data() + (past_present_size / 2),
batch_size, number_of_heads,
sequence_length, total_sequence_length,
max_sequence_length, head_size);
// Output(s)
tester.AddOutput<MLFloat16>("output", input_dims, output);
tester.AddOutput<MLFloat16>("present", past_dims, present);
tester.SetOutputTolerance(0.005f);
// Run - Regular kernel execution path
{
std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
execution_providers.push_back(DefaultCudaExecutionProvider());
tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
}
// Test alternate kernel path of loading more KV data "in flight"
{
ScopedEnvironmentVariables scoped_env_vars{
EnvVarMap{{onnxruntime::contrib::attention::kDecoderMaskedAttentionLoadKVDataInFlight, "1"}}};
std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
execution_providers.push_back(DefaultCudaExecutionProvider());
tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
}
}
TEST(DecoderMaskedMultiHeadAttentionTest, cuda_self_attn_fp16) {
TestDecoderMaskedMultiHeadAttention<MLFloat16>(/* is_cross_attn = */ false);
}
#endif
TEST(DecoderMaskedMultiHeadAttentionTest, cpu_cross_attn_fp32) {
TestDecoderMaskedMultiHeadAttention<float>(/* is_cross_attn = */ true, /* use_cuda = */ false);
}
TEST(DecoderMaskedMultiHeadAttentionTest, cpu_self_attn_fp32) {
TestDecoderMaskedMultiHeadAttention<float>(/* is_cross_attn = */ false, /* use_cuda = */ false);
}
} // namespace test
} // namespace onnxruntime