mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-27 20:02:15 +00:00
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:
parent
2e4e221da8
commit
4ffc1ff3b4
7 changed files with 402 additions and 388 deletions
|
|
@ -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);
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Reference in a new issue