diff --git a/docs/ContribOperators.md b/docs/ContribOperators.md
index 3e7ad91369..5862716638 100644
--- a/docs/ContribOperators.md
+++ b/docs/ContribOperators.md
@@ -144,14 +144,14 @@ This version of the operator has been available since version 1 of the 'com.micr
Whether every token can only attend to previous tokens. Default value is 0.
-#### Inputs (3 - 7)
+#### Inputs (2 - 7)
- input : T
- Input tensor with shape (batch_size, sequence_length, input_hidden_size)
- weights : T
- Merged Q/K/V weights with shape (input_hidden_size, hidden_size + hidden_size + v_hidden_size)
-- bias : T
+- bias (optional) : T
- Bias tensor with shape (hidden_size + hidden_size + v_hidden_size) for input projection
- mask_index (optional) : M
- Attention mask with shape (batch_size, 1, max_sequence_length, max_sequence_length), (batch_size, total_sequence_length) or (batch_size, sequence_length, total_sequence_length), or index with shape (batch_size) or (2 * batch_size)
@@ -2381,30 +2381,40 @@ This version of the operator has been available since version 1 of the 'com.micr
- The value to be filled in the attention mask. Default value is -10000.0f
- num_heads : int (required)
- Number of attention heads
+- scale : float
+- Custom scale will be used if specified. Default value is 1/sqrt(head_size)
-#### Inputs (1 - 6)
+#### Inputs (1 - 8)
- query : T
- Query with shape (batch_size, sequence_length, hidden_size), or packed QKV with shape (batch_size, kv_sequence_length, num_heads, 3, head_size)
- key (optional) : T
-- Key with shape (batch_size, kv_sequence_length, hidden_size), or packed KV with shape (batch_size, kv_sequence_length, num_heads, 2, head_size)
+- Key with shape (batch_size, kv_sequence_length, hidden_size), or packed KV with shape (batch_size, kv_sequence_length, num_heads, 2, head_size), or past_key with shape (batch_size, num_heads, kv_sequence_length, head_size)
- value (optional) : T
-- Value with shape (batch_size, kv_sequence_length, v_hidden_size)
+- Value with shape (batch_size, kv_sequence_length, v_hidden_size), or past_value with shape (batch_size, num_heads, kv_sequence_length, head_size)
- bias (optional) : T
- Bias tensor with shape (hidden_size + hidden_size + v_hidden_size) from input projection
- key_padding_mask (optional) : M
- Key padding mask with shape (batch_size) or (batch_size, kv_sequence_length)
- relative_position_bias (optional) : T
- relative position bias: addition to QxK' with shape (batch_size, num_heads, sequence_length, total_sequence_length) or (1, num_heads, sequence_length, total_sequence_length)
+- past_key (optional) : T
+- past state for self attention key with shape (batch_size, num_heads, past_sequence_length, head_size)
+- past_value (optional) : T
+- past state for self attention value with shape (batch_size, num_heads, past_sequence_length, head_size)
-#### Outputs
+#### Outputs (1 - 3)
- output : T
- 3D output tensor with shape (batch_size, sequence_length, v_hidden_size)
+- present_key (optional) : T
+- present state for cross attention key with shape (batch_size, num_heads, kv_sequence_length, head_size)or present state for self attention key with shape (batch_size, num_heads, total_sequence_length, head_size)
+- present_value (optional) : T
+- present state for cross attention value with shape (batch_size, num_heads, kv_sequence_length, head_size)or present state for self attention value with shape (batch_size, num_heads, total_sequence_length, head_size)
#### Type Constraints
diff --git a/docs/OperatorKernels.md b/docs/OperatorKernels.md
index 5cf5b3500e..1e687e1cb4 100644
--- a/docs/OperatorKernels.md
+++ b/docs/OperatorKernels.md
@@ -813,7 +813,7 @@ Do not modify directly.*
|Inverse|*in* X:**T**
*out* Y:**T**|1+|**T** = tensor(double), tensor(float), tensor(float16)|
|Irfft|*in* X:**T**
*out* Y:**T**|1+|**T** = tensor(double), tensor(float), tensor(float16)|
|LongformerAttention|*in* input:**T**
*in* weight:**T**
*in* bias:**T**
*in* mask:**T**
*in* global_weight:**T**
*in* global_bias:**T**
*in* global:**G**
*out* output:**T**|1+|**T** = tensor(float), tensor(float16)|
-|MultiHeadAttention|*in* query:**T**
*in* key:**T**
*in* value:**T**
*in* bias:**T**
*in* key_padding_mask:**M**
*in* relative_position_bias:**T**
*out* output:**T**|1+|**T** = tensor(float), tensor(float16)|
+|MultiHeadAttention|*in* query:**T**
*in* key:**T**
*in* value:**T**
*in* bias:**T**
*in* key_padding_mask:**M**
*in* relative_position_bias:**T**
*in* past_key:**T**
*in* past_value:**T**
*out* output:**T**
*out* present_key:**T**
*out* present_value:**T**|1+|**T** = tensor(float), tensor(float16)|
|NGramRepeatBlock|*in* input_ids:**Tid**
*in* scores:**T**
*out* scores_out:**T**|1+|**T** = tensor(float)
**Tid** = tensor(int64)|
|NhwcConv|*in* X:**T**
*in* W:**T**
*in* B:**T**
*out* Y:**T**|1+|**T** = tensor(float), tensor(float16)|
|QAttention|*in* input:**T1**
*in* weight:**T2**
*in* bias:**T3**
*in* input_scale:**T3**
*in* weight_scale:**T3**
*in* mask_index:**T4**
*in* input_zero_point:**T1**
*in* weight_zero_point:**T2**
*in* past:**T3**
*out* output:**T3**
*out* present:**T3**|1+|**T1** = tensor(int8)
**T2** = tensor(int8)
**T3** = tensor(float), tensor(float16)
**T4** = tensor(int32)|
diff --git a/onnxruntime/contrib_ops/cpu/bert/attention_base.cc b/onnxruntime/contrib_ops/cpu/bert/attention_base.cc
index f3bbb4f039..07f3b49b4e 100644
--- a/onnxruntime/contrib_ops/cpu/bert/attention_base.cc
+++ b/onnxruntime/contrib_ops/cpu/bert/attention_base.cc
@@ -193,6 +193,7 @@ Status AttentionBase::CheckInputs(const TensorShape& input_shape,
}
}
+ bool broadcast_res_pos_bias = false;
if (relative_position_bias != nullptr) {
const auto& relative_position_bias_dims = relative_position_bias->Shape().GetDims();
@@ -202,11 +203,14 @@ Status AttentionBase::CheckInputs(const TensorShape& input_shape,
relative_position_bias_dims.size());
}
- if (relative_position_bias_dims[0] != batch_size) {
+ if (relative_position_bias_dims[0] != batch_size && relative_position_bias_dims[0] != 1) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
- "Input 'relative_position_bias' dimension 0 should be same as batch_size, got ",
+ "Input 'relative_position_bias' dimension 0 should be same as batch_size or 1, got ",
relative_position_bias_dims[0]);
}
+ if (relative_position_bias_dims[0] == 1) {
+ broadcast_res_pos_bias = true;
+ }
if (relative_position_bias_dims[1] != num_heads_) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"Input 'relative_position_bias' dimension 1 should be same as number of heads, got ",
@@ -255,6 +259,8 @@ Status AttentionBase::CheckInputs(const TensorShape& input_shape,
output_parameters->mask_filter_value = mask_filter_value_;
output_parameters->scale = scale_;
output_parameters->mask_type = mask_type;
+ output_parameters->broadcast_res_pos_bias = broadcast_res_pos_bias;
+ output_parameters->pass_past_in_kv = false;
}
return Status::OK();
diff --git a/onnxruntime/contrib_ops/cpu/bert/attention_common.h b/onnxruntime/contrib_ops/cpu/bert/attention_common.h
index adb6805632..292c1aae1e 100644
--- a/onnxruntime/contrib_ops/cpu/bert/attention_common.h
+++ b/onnxruntime/contrib_ops/cpu/bert/attention_common.h
@@ -52,6 +52,8 @@ struct AttentionParameters {
bool is_unidirectional;
bool past_present_share_buffer;
bool do_rotary;
+ bool broadcast_res_pos_bias;
+ bool pass_past_in_kv;
float mask_filter_value;
float scale;
AttentionMaskType mask_type;
diff --git a/onnxruntime/contrib_ops/cpu/bert/multihead_attention_helper.h b/onnxruntime/contrib_ops/cpu/bert/multihead_attention_helper.h
index 34a615a880..80e506b499 100644
--- a/onnxruntime/contrib_ops/cpu/bert/multihead_attention_helper.h
+++ b/onnxruntime/contrib_ops/cpu/bert/multihead_attention_helper.h
@@ -18,16 +18,21 @@ Status CheckInputs(const T* query,
const T* bias,
const T* key_padding_mask,
const T* relative_position_bias,
+ const T* past_key,
+ const T* past_value,
void* parameters,
int num_heads,
float mask_filter_value,
+ float scale,
int max_threads_per_block) {
// key_padding_mask (K/V) : (B) or (B, L) or None
// relative_position_bias : (B, 1, S, L)
+ // past_key : (B, N, S*, H)
+ // past_value : (B, N, S*, H)
// When no packing for q/k/v:
// query (Q) : (B, S, D)
- // key (K) : (B, L, D)
- // value (V) : (B, L, D_v)
+ // key (K) : (B, L, D) or (B, N, S*, H)
+ // value (V) : (B, L, D_v) or (B, N, S*, H)
// bias (Q/K/V) : (D + D + D_v)
// When packed kv is used:
// query (Q) : (B, S, D)
@@ -40,6 +45,7 @@ Status CheckInputs(const T* query,
// value (V) : None
// bias (Q/K/V) : None
+
const auto& query_dims = query->Shape().GetDims();
if (query_dims.size() != 3 && query_dims.size() != 5) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Input 'query' is expected to have 3 or 5 dimensions, got ",
@@ -52,6 +58,63 @@ Status CheckInputs(const T* query,
int head_size = static_cast(hidden_size) / num_heads;
int kv_sequence_length = sequence_length;
+ int past_sequence_length = 0;
+ if (past_key != nullptr && past_value != nullptr) {
+ const auto& past_key_dims = past_key->Shape().GetDims();
+ const auto& past_value_dims = past_value->Shape().GetDims();
+
+ if (past_key_dims.size() != 4) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Input 'past_key' is expected to have 4 dimensions, got ",
+ past_key_dims.size());
+ }
+ if (past_value_dims.size() != 4) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Input 'past_value' is expected to have 4 dimensions, got ",
+ past_value_dims.size());
+ }
+
+ if (past_key_dims[0] != batch_size) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Input 'past_key' dimension 0 should be batch_size, got ",
+ past_key_dims[0]);
+ }
+ if (past_value_dims[0] != batch_size) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Input 'past_value' dimension 0 should be batch_size, got ",
+ past_value_dims[0]);
+ }
+
+ if (past_key_dims[1] != num_heads) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Input 'past_key' dimension 1 should be same as number of heads, got ",
+ past_key_dims[1]);
+ }
+ if (past_value_dims[1] != num_heads) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Input 'past_value' dimension 1 should be same as number of heads, got ",
+ past_value_dims[1]);
+ }
+ if (past_key_dims[2] != past_value_dims[2]) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Input 'past_key' and 'past_value' shall have same dim 2 (past_sequence_length)");
+ }
+ if (past_key_dims[3] != head_size) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Input 'past_key' dimension 3 should be same as head_size, got ",
+ past_key_dims[3]);
+ }
+ if (past_value_dims[3] != head_size) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Input 'past_value' dimension 3 should be same as head_size, got ",
+ past_value_dims[3]);
+ }
+ past_sequence_length = static_cast(past_key_dims[2]);
+ } else if (past_key != nullptr || past_value != nullptr) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Input 'past_key' and 'past_value' shall be both present or both absent");
+ }
+
if (key != nullptr) {
if (query_dims.size() != 3) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Input 'query' is expected to have 3 dimensions when key is given, got ",
@@ -59,8 +122,8 @@ Status CheckInputs(const T* query,
}
const auto& key_dims = key->Shape().GetDims();
- if (key_dims.size() != 3 && key_dims.size() != 5) {
- return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Input 'key' is expected to have 3 or 5 dimensions, got ",
+ if (key_dims.size() != 3 && key_dims.size() != 4 && key_dims.size() != 5) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Input 'key' is expected to have 3, 4, or 5 dimensions, got ",
key_dims.size());
}
if (query_dims[0] != key_dims[0]) {
@@ -73,8 +136,9 @@ Status CheckInputs(const T* query,
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"Input 'query' and 'key' shall have same dim 2 (hidden_size)");
}
- } else // if (key_dims.size() == 5)
- {
+
+ kv_sequence_length = static_cast(key_dims[1]);
+ } else if (key_dims.size() == 5) {
if (static_cast(key_dims[2]) != num_heads || static_cast(key_dims[3]) != 2 || static_cast(key_dims[4]) != head_size) {
return ORT_MAKE_STATUS(
ONNXRUNTIME, INVALID_ARGUMENT,
@@ -83,9 +147,17 @@ Status CheckInputs(const T* query,
if (value != nullptr) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Expect 'value' be none when 'key' has packed kv format.");
}
- }
- kv_sequence_length = static_cast(key_dims[1]);
+ kv_sequence_length = static_cast(key_dims[1]);
+ } else { // key_dims.size() == 4 (cross-attention with past_key)
+ if (static_cast(key_dims[1]) != num_heads || static_cast(key_dims[3]) != head_size) {
+ return ORT_MAKE_STATUS(
+ ONNXRUNTIME, INVALID_ARGUMENT,
+ "Expect 'key' shape (batch_size, num_heads, kv_sequence_length, head_size) for past_key");
+ }
+
+ kv_sequence_length = static_cast(key_dims[2]);
+ }
} else { // packed QKV
if (query_dims.size() != 5) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Input 'query' is expected to have 5 dimensions when key is empty, got ",
@@ -128,11 +200,13 @@ Status CheckInputs(const T* query,
}
}
+ // NOTE: In Cross-Attention, we pass the past key and value to 'key' and 'value' instead of 'past_key' and 'past_value'.
+ bool pass_past_in_kv = false;
int v_hidden_size = hidden_size;
if (value != nullptr) {
const auto& value_dims = value->Shape().GetDims();
- if (value_dims.size() != 3) {
- return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Input 'value' is expected to have 3 dimensions, got ",
+ if (value_dims.size() != 3 && value_dims.size() != 4) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Input 'value' is expected to have 3 or 4 dimensions, got ",
value_dims.size());
}
@@ -141,13 +215,24 @@ Status CheckInputs(const T* query,
"Input 'query' and 'value' shall have same dim 0 (batch_size)");
}
- if (static_cast(kv_sequence_length) != value_dims[1]) {
- return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
- "Input 'key' and 'value' shall have same same dim 1 (kv_sequence_length)");
+ if (value_dims.size() == 3) {
+ if (static_cast(kv_sequence_length) != value_dims[1]) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Input 'key' and 'value' shall have the same dim 1 (kv_sequence_length)");
+ }
+ v_hidden_size = static_cast(value_dims[2]);
+ } else { // value_dims.size() == 4
+ if (static_cast(kv_sequence_length) != value_dims[2]) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Input 'past_key' and 'past_value' shall have the same dim 2 (kv_sequence_length)");
+ }
+ v_hidden_size = static_cast(value_dims[1]) * static_cast(value_dims[3]);
+ pass_past_in_kv = true;
}
- v_hidden_size = static_cast(value_dims[2]);
}
+ int total_sequence_length = past_sequence_length + kv_sequence_length;
+ bool broadcast_res_pos_bias = false;
if (relative_position_bias != nullptr) {
const auto& relative_position_bias_dims = relative_position_bias->Shape().GetDims();
@@ -161,6 +246,9 @@ Status CheckInputs(const T* query,
"Input 'relative_position_bias' dimension 0 should be batch_size or 1, got ",
relative_position_bias_dims[0]);
}
+ if (relative_position_bias_dims[0] == 1) {
+ broadcast_res_pos_bias = true;
+ }
if (relative_position_bias_dims[1] != num_heads) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"Input 'relative_position_bias' dimension 1 should be same as number of heads, got ",
@@ -171,7 +259,7 @@ Status CheckInputs(const T* query,
"Input 'relative_position_bias' dimension 2 should be same as sequence_length, got ",
relative_position_bias_dims[2]);
}
- if (relative_position_bias_dims[3] != kv_sequence_length) {
+ if (relative_position_bias_dims[3] != total_sequence_length) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"Input 'relative_position_bias' dimension 3 should be same as total_sequence_length, got ",
relative_position_bias_dims[3]);
@@ -182,9 +270,9 @@ Status CheckInputs(const T* query,
AttentionParameters* output_parameters = reinterpret_cast(parameters);
output_parameters->batch_size = batch_size;
output_parameters->sequence_length = sequence_length;
- output_parameters->past_sequence_length = 0;
+ output_parameters->past_sequence_length = past_sequence_length;
output_parameters->kv_sequence_length = kv_sequence_length;
- output_parameters->total_sequence_length = kv_sequence_length;
+ output_parameters->total_sequence_length = total_sequence_length;
output_parameters->max_sequence_length = 0;
output_parameters->input_hidden_size = 0;
output_parameters->hidden_size = hidden_size;
@@ -196,7 +284,9 @@ Status CheckInputs(const T* query,
output_parameters->past_present_share_buffer = false;
output_parameters->mask_filter_value = mask_filter_value;
output_parameters->mask_type = mask_type;
- output_parameters->scale = 0.0f;
+ output_parameters->scale = scale;
+ output_parameters->broadcast_res_pos_bias = broadcast_res_pos_bias;
+ output_parameters->pass_past_in_kv = pass_past_in_kv;
}
if (max_threads_per_block > 0 && num_heads > max_threads_per_block) {
diff --git a/onnxruntime/contrib_ops/cuda/bert/attention.cc b/onnxruntime/contrib_ops/cuda/bert/attention.cc
index f0669f6dc3..c5285ae486 100644
--- a/onnxruntime/contrib_ops/cuda/bert/attention.cc
+++ b/onnxruntime/contrib_ops/cuda/bert/attention.cc
@@ -67,9 +67,12 @@ Status Attention::ComputeInternal(OpKernelContext* context) const {
auto& device_prop = GetDeviceProp();
AttentionParameters parameters;
+ // Use the second dimension from weight for bias to get q_hidden_size when bias is nullptr
+ std::vector bias_dims{weights->Shape().GetDims()[1]};
+ const TensorShape bias_shape{bias_dims};
ORT_RETURN_IF_ERROR(CheckInputs(input->Shape(),
weights->Shape(),
- bias->Shape(),
+ bias != nullptr ? bias->Shape() : bias_shape,
mask_index,
past,
relative_position_bias,
@@ -200,18 +203,22 @@ Status Attention::ComputeInternal(OpKernelContext* context) const {
typedef typename ToCudaType::MappedType CudaT;
AttentionData data;
data.gemm_buffer = reinterpret_cast(gemm_buffer.get());
- data.bias = reinterpret_cast(bias->Data());
+ data.bias = nullptr == bias ? nullptr : reinterpret_cast(bias->Data());
data.query = nullptr;
data.key = nullptr;
data.value = nullptr;
data.mask_index = (nullptr == mask_index) ? nullptr : mask_index->Data();
data.mask_index_dims = (nullptr == mask_index) ? gsl::span() : mask_index->Shape().GetDims();
data.past = (nullptr == past) ? nullptr : reinterpret_cast(past->Data());
+ data.past_key = nullptr;
+ data.past_value = nullptr;
data.relative_position_bias = (nullptr == relative_position_bias) ? nullptr : reinterpret_cast(relative_position_bias->Data());
data.has_qkv_workspace = true;
data.workspace = reinterpret_cast(work_space.get());
data.output = reinterpret_cast(output->MutableData());
data.present = (nullptr == present) ? nullptr : reinterpret_cast(present->MutableData());
+ data.present_key = nullptr;
+ data.present_value = nullptr;
data.fused_runner = reinterpret_cast(fused_runner);
data.fused_cross_attention_kernel = nullptr;
data.use_memory_efficient_attention = use_memory_efficient_attention;
diff --git a/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu b/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu
index 214c7dfefd..9f655bf569 100644
--- a/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu
+++ b/onnxruntime/contrib_ops/cuda/bert/attention_impl.cu
@@ -333,6 +333,56 @@ Status PrepareQkv(contrib::AttentionParameters& parameters,
data.gemm_buffer, data.bias, qkv, true, v_head_size, qkv_add_bias,
3, parameters.do_rotary, parameters.original_past_sequence_length);
}
+ }
+ // cross attention with past/present state
+ else if (data.past_key != nullptr || data.present_key != nullptr) {
+ // no bias for T5 cross attention
+ assert(data.bias == nullptr);
+ // cross attention with past state
+ if (data.past_key != nullptr && data.present_key == nullptr) {
+ assert(data.past_value != nullptr);
+ assert(data.query != nullptr);
+ assert(data.key == nullptr);
+ assert(data.value == nullptr);
+ ORT_RETURN_IF_ERROR(LaunchTransQkv(stream, 1, sequence_length, batch_size, qk_head_size, num_heads,
+ max_threads_per_block, false, data.query, q));
+ }
+ // cross attention with present state or self attention with present state
+ else if (data.past_key == nullptr && data.present_key != nullptr) {
+ assert(data.past_value == nullptr);
+ assert(data.present_value != nullptr);
+ assert(data.query != nullptr);
+ assert(data.key != nullptr);
+ assert(data.value != nullptr);
+
+ // TODO: supporting packed qkv for self attention may benefit performance
+ ORT_RETURN_IF_ERROR(LaunchTransQkv(stream, 1, sequence_length, batch_size, qk_head_size, num_heads,
+ max_threads_per_block, false, data.query, q));
+
+ // TODO: supporting packed kv for cross attention may benefit performance
+ ORT_RETURN_IF_ERROR(LaunchTransQkv(stream, 1, kv_sequence_length, batch_size, qk_head_size, num_heads,
+ max_threads_per_block, false, data.key, data.present_key));
+ ORT_RETURN_IF_ERROR(LaunchTransQkv(stream, 1, kv_sequence_length, batch_size, v_head_size, num_heads,
+ max_threads_per_block, false, data.value, data.present_value));
+ }
+ // self attention with past and present state
+ else {
+ assert(data.past_key != nullptr);
+ assert(data.past_value != nullptr);
+ assert(data.present_key != nullptr);
+ assert(data.present_value != nullptr);
+ assert(data.query != nullptr);
+ assert(data.key != nullptr);
+ assert(data.value != nullptr);
+ // TODO: supporting packed qkv for self attention may benefit performance
+ ORT_RETURN_IF_ERROR(LaunchTransQkv(stream, 1, sequence_length, batch_size, qk_head_size, num_heads,
+ max_threads_per_block, false, data.query, q));
+ ORT_RETURN_IF_ERROR(LaunchTransQkv(stream, 1, kv_sequence_length, batch_size, qk_head_size, num_heads,
+ max_threads_per_block, false, data.key, k));
+ ORT_RETURN_IF_ERROR(LaunchTransQkv(stream, 1, kv_sequence_length, batch_size, v_head_size, num_heads,
+ max_threads_per_block, false, data.value, v));
+ }
+ qkv_format = AttentionQkvFormat::Q_K_V_BNSH;
} else if (data.key == nullptr) { // gemm_buffer == nullptr and packed qkv
assert(data.bias == nullptr);
assert(qk_head_size == v_head_size);
@@ -536,6 +586,28 @@ Status QkvToContext(
k = data.present;
v = data.present + batches * present_size_per_batch_k;
}
+
+ if (nullptr != data.past_key || nullptr != data.present_key) {
+ assert(qkv_format == AttentionQkvFormat::Q_K_V_BNSH);
+ if (nullptr != data.past_key && nullptr == data.present_key) {
+ k = const_cast(data.past_key);
+ v = const_cast(data.past_value);
+ } else if (nullptr == data.past_key && nullptr != data.present_key) {
+ k = data.present_key;
+ v = data.present_value;
+ } else {
+ ORT_RETURN_IF_ERROR(
+ LaunchConcatTensorToTensor(stream, parameters.total_sequence_length, sequence_length, batch_size, qk_head_size, num_heads,
+ max_threads_per_block, 1, data.past_key, k, data.present_key));
+ ORT_RETURN_IF_ERROR(
+ LaunchConcatTensorToTensor(stream, parameters.total_sequence_length, sequence_length, batch_size, v_head_size, num_heads,
+ max_threads_per_block, 1, data.past_value, v, data.present_value));
+
+ // Update pointers to present_k and present_v.
+ k = data.present_key;
+ v = data.present_value;
+ }
+ }
} else {
assert(qk_head_size == v_head_size);
assert(data.fused_cross_attention_kernel == nullptr);
@@ -544,6 +616,11 @@ Status QkvToContext(
assert(!data.use_memory_efficient_attention);
assert(data.has_qkv_workspace);
+ if (nullptr != data.past_key || nullptr != data.present_key) {
+ // TODO: support this case.
+ ORT_THROW("buffer sharing for no bias case between past and present is not supported yet.");
+ }
+
if (data.present != data.past) {
// For easy testing. Production should better avoid this path.
int64_t kv_size = 2LL * (int64_t)batch_size * num_heads * parameters.max_sequence_length * qk_head_size;
@@ -719,6 +796,8 @@ Status QkvToContext(
cublasSetStream(cublas, stream);
+ DUMP_TENSOR_D("q[BNSH]", q, batch_size, num_heads, sequence_length, qk_head_size);
+ DUMP_TENSOR_D("k[BNSH]", k, batch_size, num_heads, total_sequence_length, qk_head_size);
CUBLAS_RETURN_IF_ERROR(cublasGemmStridedBatchedHelper(
cublas, CUBLAS_OP_T, CUBLAS_OP_N,
total_sequence_length, sequence_length, qk_head_size,
@@ -743,21 +822,22 @@ Status QkvToContext(
T* persistent_softmax_workspace = scratch1; // replace Q*K' in place with masked score for persistent softmax.
ORT_RETURN_IF_ERROR(
ComputeSoftmaxWithRawMask(stream, total_sequence_length, sequence_length, batch_size, num_heads,
- mask_index, nullptr, data.relative_position_bias, scratch1, scratch2,
- parameters.is_unidirectional, scale, mask_dimension,
- parameters.max_sequence_length, use_persistent_softmax,
- persistent_softmax_workspace, mask_filter_value));
+ mask_index, nullptr, data.relative_position_bias, parameters.broadcast_res_pos_bias,
+ scratch1, scratch2, parameters.is_unidirectional, scale, mask_dimension,
+ parameters.max_sequence_length, use_persistent_softmax, persistent_softmax_workspace,
+ mask_filter_value));
} else if (nullptr != mask_index) { // 1d mask index
assert(mask_index_dims.size() == 1);
// mask_index has 1D shape: either (batch_size) or (2*batch_size). Only the later one has start postions.
const int* mask_start = (mask_index_dims[0] > batch_size) ? mask_index + batch_size : nullptr;
ORT_RETURN_IF_ERROR(ComputeSoftmaxWithMask1D(
stream, total_sequence_length, sequence_length, batch_size, num_heads,
- mask_index, mask_start, data.relative_position_bias, scratch1, scratch2, parameters.is_unidirectional));
+ mask_index, mask_start, data.relative_position_bias, parameters.broadcast_res_pos_bias,
+ scratch1, scratch2, parameters.is_unidirectional));
} else { // no mask
ORT_RETURN_IF_ERROR(
ComputeSoftmax(stream, total_sequence_length, sequence_length, batch_size, num_heads, data.relative_position_bias,
- scratch1, scratch2, parameters.is_unidirectional));
+ parameters.broadcast_res_pos_bias, scratch1, scratch2, parameters.is_unidirectional));
}
DUMP_TENSOR_D("Softmax", scratch2, batch_size * num_heads, sequence_length, total_sequence_length);
@@ -908,13 +988,15 @@ Status DecoderQkvToContext(
if (has_key_padding_mask) {
constexpr int mask_dimension = 2;
constexpr int max_sequence_length = 0;
- ORT_RETURN_IF_ERROR(ComputeSoftmaxWithRawMask(stream, kv_sequence_length, sequence_length, batch_size, num_heads,
- nullptr, key_padding_mask, add_before_softmax, scratch1, scratch2,
- is_unidirectional, 1.0f, mask_dimension, max_sequence_length,
- false, nullptr, mask_filter_value));
+ ORT_RETURN_IF_ERROR(ComputeSoftmaxWithRawMask(stream, kv_sequence_length, sequence_length, batch_size,
+ num_heads, nullptr, key_padding_mask, add_before_softmax,
+ false/*broadcast rpb*/, scratch1, scratch2, is_unidirectional,
+ 1.0f, mask_dimension, max_sequence_length, false, nullptr,
+ mask_filter_value));
} else {
ORT_RETURN_IF_ERROR(ComputeSoftmax(stream, kv_sequence_length, sequence_length, batch_size, num_heads,
- add_before_softmax, scratch1, scratch2, is_unidirectional));
+ add_before_softmax, false/*broadcast rpb*/, scratch1, scratch2,
+ is_unidirectional));
}
// compute P*V (as V*P), and store in scratch3: BxNxSxH
diff --git a/onnxruntime/contrib_ops/cuda/bert/attention_impl.h b/onnxruntime/contrib_ops/cuda/bert/attention_impl.h
index ec7371db4c..0df9547ba7 100644
--- a/onnxruntime/contrib_ops/cuda/bert/attention_impl.h
+++ b/onnxruntime/contrib_ops/cuda/bert/attention_impl.h
@@ -57,6 +57,8 @@ struct AttentionData {
const int* mask_index;
gsl::span mask_index_dims;
const T* past;
+ const T* past_key;
+ const T* past_value;
const T* relative_position_bias;
bool has_qkv_workspace;
@@ -64,6 +66,8 @@ struct AttentionData {
T* output;
T* present;
+ T* present_key;
+ T* present_value;
void* fused_runner;
const void* fused_cross_attention_kernel;
diff --git a/onnxruntime/contrib_ops/cuda/bert/attention_softmax.h b/onnxruntime/contrib_ops/cuda/bert/attention_softmax.h
index 92851c446d..d2cb5d2d58 100644
--- a/onnxruntime/contrib_ops/cuda/bert/attention_softmax.h
+++ b/onnxruntime/contrib_ops/cuda/bert/attention_softmax.h
@@ -39,7 +39,8 @@ __device__ inline void Softmax(const int all_sequence_length,
const int sequence_length,
const int valid_end,
const int valid_start,
- const T* add_before_softmax,
+ const T* rel_pos_bias,
+ const bool broadcast_rel_pos_bias,
const T* input,
T* output) {
using BlockReduce = cub::BlockReduce;
@@ -50,17 +51,22 @@ __device__ inline void Softmax(const int all_sequence_length,
float thread_data_max(-CUDART_INF_F);
- const bool no_add = (add_before_softmax == nullptr);
+ const bool no_rpb = (rel_pos_bias == nullptr);
// e^x is represented as infinity if x is large enough, like 100.f.
// Infinity divided by Infinity is a NAN. Thus, softmax gets a NAN if one or more item are large enough.
// a math transform as below is leveraged to get a stable softmax:
// e^xi/(e^x1 + ...e^xn) = e^(xi - max) / (e^(x1 - max) + ... + e^(xn - max))
const int offset = (blockIdx.y * gridDim.x + blockIdx.x) * all_sequence_length;
+ const int size_per_batch = gridDim.x * all_sequence_length;
for (int i = threadIdx.x; i < valid_end; i += TPB) {
if (i >= valid_start) {
const int index = offset + i;
- float input_at_idx = no_add ? float(input[index]) : float(input[index] + add_before_softmax[index]);
+ float input_at_idx = no_rpb
+ ? float(input[index])
+ : float(input[index] + (broadcast_rel_pos_bias
+ ? rel_pos_bias[index % size_per_batch]
+ : rel_pos_bias[index]));
if (thread_data_max < input_at_idx) {
thread_data_max = input_at_idx;
}
@@ -79,7 +85,7 @@ __device__ inline void Softmax(const int all_sequence_length,
for (int i = threadIdx.x; i < valid_end; i += TPB) {
if (i >= valid_start) {
const int index = offset + i;
- float val = no_add ? input[index] : input[index] + add_before_softmax[index];
+ float val = no_rpb ? input[index] : input[index] + rel_pos_bias[index % size_per_batch];
thread_data_sum += expf(val - max_block);
}
}
@@ -92,7 +98,7 @@ __device__ inline void Softmax(const int all_sequence_length,
for (int i = threadIdx.x; i < all_sequence_length; i += TPB) {
const int index = offset + i;
- float input_at_idx = no_add ? float(input[index]) : float(input[index] + add_before_softmax[index]);
+ float input_at_idx = no_rpb ? float(input[index]) : float(input[index] + rel_pos_bias[index % size_per_batch]);
const float val = (i >= valid_start && i < valid_end) ? expf(input_at_idx - max_block) * sum_reverse_block : 0.f;
output[index] = T(val);
}
@@ -103,7 +109,8 @@ __device__ inline void SoftmaxSmall(const int all_sequence_length,
const int sequence_length,
const int valid_end,
const int valid_start,
- const T* add_before_softmax,
+ const T* rel_pos_bias,
+ const bool broadcast_rel_pos_bias,
const T* input,
T* output,
bool is_unidirectional) {
@@ -139,8 +146,13 @@ __device__ inline void SoftmaxSmall(const int all_sequence_length,
// Infinity divided by Infinity is a NAN. Thus, softmax gets a NAN if one or more item are large enough.
// a math transform as below is leveraged to get a stable softmax:
// e^xi/(e^x1 + ...e^xn) = e^(xi - max) / (e^(x1 - max) + ... + e^(xn - max))
- const bool no_add = (add_before_softmax == nullptr);
- float input_data = no_add ? float(input[index]) : float(input[index] + add_before_softmax[index]);
+ const bool no_rpb = (rel_pos_bias == nullptr);
+ const int size_per_batch = gridDim.x * all_sequence_length;
+ float input_data = no_rpb
+ ? float(input[index])
+ : float(input[index] + (broadcast_rel_pos_bias
+ ? rel_pos_bias[index % size_per_batch]
+ : rel_pos_bias[index]));
float thread_data_max = is_valid ? input_data : float(-CUDART_INF_F);
const auto max = BlockReduce(tmp_storage).Reduce(thread_data_max, cub::Max(), end);
@@ -174,7 +186,8 @@ __global__ void SoftmaxLargeKernel(const int all_sequence_length,
const int sequence_length,
const int valid_end,
const int valid_start,
- const T* add_before_softmax,
+ const T* rel_pos_bias,
+ const bool broadcast_rel_pos_bias,
const T* input,
T* output,
bool is_unidirectional) {
@@ -205,6 +218,7 @@ __global__ void SoftmaxLargeKernel(const int all_sequence_length,
// Input dimension is BxNxSxS*; blockIdx.y is batch index b; gridDim.x=N*S; blockIdx.x is index within N*S;
const int offset = (blockIdx.y * gridDim.x + blockIdx.x) * all_sequence_length;
+ const int size_per_batch = gridDim.x * all_sequence_length;
float thread_data_max = -CUDART_INF_F;
for (int seq_idx = threadIdx.x; seq_idx < all_sequence_length; seq_idx += TPB) {
@@ -216,8 +230,12 @@ __global__ void SoftmaxLargeKernel(const int all_sequence_length,
// a math transform as below is leveraged to get a stable softmax:
// e^xi/(e^x1 + ...e^xn) = e^(xi - max) / (e^(x1 - max) + ... + e^(xn - max))
float input_data = is_valid
- ? (add_before_softmax ? float(input[index] + add_before_softmax[index]) : float(input[index]))
- : float(-CUDART_INF_F);
+ ? (rel_pos_bias
+ ? float(input[index] + (broadcast_rel_pos_bias
+ ? rel_pos_bias[index % size_per_batch]
+ : rel_pos_bias[index]))
+ : float(input[index]))
+ : float(-CUDART_INF_F);
cached_data[seq_idx] = input_data;
thread_data_max = max(thread_data_max, input_data);
}
@@ -254,7 +272,8 @@ __global__ void SoftmaxWithRawMaskLargeKernel(const int all_sequence_length,
const int sequence_length,
const int* attention_mask, // 2D, 3D or 4D attention mask
const bool* key_padding_mask,
- const T* add_before_softmax,
+ const T* rel_pos_bias,
+ const bool broadcast_rel_pos_bias,
const T* input,
T* output,
const bool is_unidirectional,
@@ -272,16 +291,19 @@ __global__ void SoftmaxWithRawMaskLargeKernel(const int all_sequence_length,
__shared__ float max_block;
float max_thread_data = -CUDART_INF_F;
+ const int size_per_batch = gridDim.x * all_sequence_length;
// Input dimension is BxNxSxS*; blockIdx.y is batch index b; gridDim.x=N*S; blockIdx.x is index within N*S;
int base_index = (blockIdx.y * gridDim.x + blockIdx.x) * all_sequence_length;
for (int seq_idx = threadIdx.x; seq_idx < all_sequence_length; seq_idx += TPB) {
float thread_data = -CUDART_INF_F;
int index = base_index + seq_idx;
- if (add_before_softmax == nullptr) {
+ if (rel_pos_bias == nullptr) {
thread_data = float(input[index]) * rsqrt_head_size;
} else {
- thread_data = float(input[index] + add_before_softmax[index]) * rsqrt_head_size;
+ T rel_pos_bias_value = broadcast_rel_pos_bias ?
+ rel_pos_bias[index % size_per_batch] : rel_pos_bias[index];
+ thread_data = float(input[index] + rel_pos_bias_value) * rsqrt_head_size;
}
const int sequence_index = blockIdx.x % sequence_length;
@@ -357,7 +379,8 @@ __device__ inline void SoftmaxWithRawMaskSmall(const int all_sequence_length,
const int sequence_length,
const int* attention_mask, // 2D, 3D or 4D attention mask
const bool* key_padding_mask,
- const T* add_before_softmax,
+ const T* rel_pos_bias,
+ const bool broadcast_rel_pos_bias,
const T* input,
T* output,
const bool is_unidirectional,
@@ -374,6 +397,7 @@ __device__ inline void SoftmaxWithRawMaskSmall(const int all_sequence_length,
// Input dimension is BxNxSxS*; blockIdx.y is batch index b; gridDim.x=N*S; blockIdx.x is index within N*S;
int index = (blockIdx.y * gridDim.x + blockIdx.x) * all_sequence_length + threadIdx.x;
+ const int size_per_batch = gridDim.x * all_sequence_length;
float thread_data = -CUDART_INF_F;
if (threadIdx.x < all_sequence_length) {
@@ -409,8 +433,10 @@ __device__ inline void SoftmaxWithRawMaskSmall(const int all_sequence_length,
}
}
- if (add_before_softmax != nullptr) {
- thread_data += float(add_before_softmax[index]);
+ if (rel_pos_bias != nullptr) {
+ float rel_pos_bias_value = broadcast_rel_pos_bias ?
+ float(rel_pos_bias[index % size_per_batch]) : float(rel_pos_bias[index]);
+ thread_data += rel_pos_bias_value;
}
}
@@ -446,62 +472,65 @@ __device__ inline void SoftmaxWithRawMaskSmall(const int all_sequence_length,
template
__global__ void SoftmaxKernelSmall(const int all_sequence_length,
const int sequence_length,
- const T* add_before_softmax,
+ const T* rel_pos_bias,
+ const bool broadcast_rel_pos_bias,
const T* input,
T* output,
bool is_unidirectional) {
SoftmaxSmall(all_sequence_length, sequence_length, all_sequence_length, 0,
- add_before_softmax, input, output, is_unidirectional);
+ rel_pos_bias, broadcast_rel_pos_bias, input, output, is_unidirectional);
}
template
__global__ void SoftmaxKernel(const int all_sequence_length,
const int sequence_length,
- const T* add_before_softmax,
+ const T* rel_pos_bias,
+ const bool broadcast_rel_pos_bias,
const T* input,
T* output) {
Softmax(all_sequence_length, sequence_length, all_sequence_length, 0,
- add_before_softmax, input, output);
+ rel_pos_bias, broadcast_rel_pos_bias, input, output);
}
template
Status ComputeSoftmax(cudaStream_t stream, const int all_sequence_length, const int sequence_length,
- const int batch_size, const int num_heads,
- const T* add_before_softmax, const T* input, T* output, bool is_unidirectional) {
+ const int batch_size, const int num_heads, const T* rel_pos_bias,
+ const bool broadcast_rel_pos_bias, T* input, T* output, bool is_unidirectional) {
const dim3 grid(sequence_length * num_heads, batch_size, 1);
if (all_sequence_length <= 32) {
const int blockSize = 32;
SoftmaxKernelSmall<<>>(
- all_sequence_length, sequence_length, add_before_softmax, input, output, is_unidirectional);
+ all_sequence_length, sequence_length, rel_pos_bias, broadcast_rel_pos_bias, input, output, is_unidirectional);
} else if (all_sequence_length <= 64) {
const int blockSize = 64;
SoftmaxKernelSmall<<>>(
- all_sequence_length, sequence_length, add_before_softmax, input, output, is_unidirectional);
+ all_sequence_length, sequence_length, rel_pos_bias, broadcast_rel_pos_bias, input, output, is_unidirectional);
} else if (all_sequence_length <= 128) {
const int blockSize = 128;
SoftmaxKernelSmall<<>>(
- all_sequence_length, sequence_length, add_before_softmax, input, output, is_unidirectional);
+ all_sequence_length, sequence_length, rel_pos_bias, broadcast_rel_pos_bias, input, output, is_unidirectional);
} else if (all_sequence_length <= 256) {
const int blockSize = 256;
SoftmaxKernelSmall<<>>(
- all_sequence_length, sequence_length, add_before_softmax, input, output, is_unidirectional);
+ all_sequence_length, sequence_length, rel_pos_bias, broadcast_rel_pos_bias, input, output, is_unidirectional);
} else if (all_sequence_length <= 512) {
const int blockSize = 512;
SoftmaxKernelSmall<<>>(
- all_sequence_length, sequence_length, add_before_softmax, input, output, is_unidirectional);
+ all_sequence_length, sequence_length, rel_pos_bias, broadcast_rel_pos_bias, input, output, is_unidirectional);
} else if (all_sequence_length <= 1024) {
const int blockSize = 1024;
SoftmaxKernelSmall<<>>(
- all_sequence_length, sequence_length, add_before_softmax, input, output, is_unidirectional);
+ all_sequence_length, sequence_length, rel_pos_bias, broadcast_rel_pos_bias, input, output, is_unidirectional);
} else if (!is_unidirectional) {
const int blockSize = 1024;
SoftmaxKernel<<>>(
- all_sequence_length, sequence_length, add_before_softmax, input, output);
+ all_sequence_length, sequence_length, rel_pos_bias, broadcast_rel_pos_bias, input, output);
} else {
const int blockSize = 256;
const int sh_bytes = sizeof(float) * all_sequence_length;
SoftmaxLargeKernel<<>>(
- all_sequence_length, sequence_length, all_sequence_length, 0, add_before_softmax, input, output, true);
+ all_sequence_length, sequence_length, all_sequence_length, 0, rel_pos_bias, broadcast_rel_pos_bias,
+ input, output, true);
}
return CUDA_CALL(cudaGetLastError());
@@ -512,7 +541,8 @@ __global__ void MaskedSoftmaxKernelSmall(const int all_sequence_length,
const int sequence_length,
const int* mask_end,
const int* mask_start,
- const T* add_before_softmax,
+ const T* rel_pos_bias,
+ const bool broadcast_rel_pos_bias,
const T* input,
T* output,
bool is_unidirectional) {
@@ -533,7 +563,7 @@ __global__ void MaskedSoftmaxKernelSmall(const int all_sequence_length,
__syncthreads();
SoftmaxSmall(all_sequence_length, sequence_length, end_position, start_position,
- add_before_softmax, input, output, is_unidirectional);
+ rel_pos_bias, broadcast_rel_pos_bias, input, output, is_unidirectional);
}
template
@@ -541,7 +571,8 @@ __global__ void MaskedSoftmaxKernel(const int all_sequence_length,
const int sequence_length,
const int* mask_end,
const int* mask_start,
- const T* add_before_softmax,
+ const T* rel_pos_bias,
+ const bool broadcast_rel_pos_bias,
const T* input, T* output) {
__shared__ int start_position;
__shared__ int end_position;
@@ -560,7 +591,7 @@ __global__ void MaskedSoftmaxKernel(const int all_sequence_length,
__syncthreads();
Softmax(all_sequence_length, sequence_length, end_position, start_position,
- add_before_softmax, input, output);
+ rel_pos_bias, broadcast_rel_pos_bias, input, output);
}
template
@@ -568,7 +599,8 @@ __global__ void SoftmaxWithRawMaskSmallKernel(const int all_sequence_length,
const int sequence_length,
const int* attention_mask,
const bool* key_padding_mask,
- const T* add_before_softmax,
+ const T* rel_pos_bias,
+ const bool broadcast_rel_pos_bias,
const T* input,
T* output,
const bool is_unidirectional,
@@ -579,7 +611,7 @@ __global__ void SoftmaxWithRawMaskSmallKernel(const int all_sequence_length,
const float mask_filter_value) {
SoftmaxWithRawMaskSmall(
all_sequence_length, sequence_length,
- attention_mask, key_padding_mask, add_before_softmax, input, output,
+ attention_mask, key_padding_mask, rel_pos_bias, broadcast_rel_pos_bias, input, output,
is_unidirectional, rsqrt_head_size, mask_dimension, max_sequence_length,
skip_softmax, mask_filter_value);
}
@@ -592,7 +624,8 @@ Status ComputeSoftmaxWithMask1D(cudaStream_t stream,
const int num_heads,
const int* mask_index,
const int* mask_start,
- const T* add_before_softmax,
+ const T* rel_pos_bias,
+ const bool broadcast_rel_pos_bias,
const T* input,
T* output,
const bool is_unidirectional) {
@@ -602,37 +635,37 @@ Status ComputeSoftmaxWithMask1D(cudaStream_t stream,
const int blockSize = 32;
MaskedSoftmaxKernelSmall
<<>>(all_sequence_length, sequence_length, mask_index, mask_start,
- add_before_softmax, input, output, is_unidirectional);
+ rel_pos_bias, broadcast_rel_pos_bias, input, output, is_unidirectional);
} else if (all_sequence_length <= 64) {
const int blockSize = 64;
MaskedSoftmaxKernelSmall
<<>>(all_sequence_length, sequence_length, mask_index, mask_start,
- add_before_softmax, input, output, is_unidirectional);
+ rel_pos_bias, broadcast_rel_pos_bias, input, output, is_unidirectional);
} else if (all_sequence_length <= 128) {
const int blockSize = 128;
MaskedSoftmaxKernelSmall
<<>>(all_sequence_length, sequence_length, mask_index, mask_start,
- add_before_softmax, input, output, is_unidirectional);
+ rel_pos_bias, broadcast_rel_pos_bias, input, output, is_unidirectional);
} else if (all_sequence_length <= 256) {
const int blockSize = 256;
MaskedSoftmaxKernelSmall
<<>>(all_sequence_length, sequence_length, mask_index, mask_start,
- add_before_softmax, input, output, is_unidirectional);
+ rel_pos_bias, broadcast_rel_pos_bias, input, output, is_unidirectional);
} else if (all_sequence_length <= 512) {
const int blockSize = 512;
MaskedSoftmaxKernelSmall
<<>>(all_sequence_length, sequence_length, mask_index, mask_start,
- add_before_softmax, input, output, is_unidirectional);
+ rel_pos_bias, broadcast_rel_pos_bias, input, output, is_unidirectional);
} else if (all_sequence_length <= 1024) {
const int blockSize = 1024;
MaskedSoftmaxKernelSmall
<<>>(all_sequence_length, sequence_length, mask_index, mask_start,
- add_before_softmax, input, output, is_unidirectional);
+ rel_pos_bias, broadcast_rel_pos_bias, input, output, is_unidirectional);
} else if (!is_unidirectional) {
const int blockSize = 1024;
MaskedSoftmaxKernel
<<>>(all_sequence_length, sequence_length, mask_index, mask_start,
- add_before_softmax, input, output);
+ rel_pos_bias, broadcast_rel_pos_bias, input, output);
} else {
return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Attention CUDA operator does not support total sequence length > 1024.");
}
@@ -648,7 +681,8 @@ Status ComputeSoftmaxWithRawMask(cudaStream_t stream,
const int num_heads,
const int* attention_mask,
const bool* key_padding_mask,
- const T* add_before_softmax,
+ const T* rel_pos_bias,
+ const bool broadcast_rel_pos_bias,
const T* input,
T* output,
const bool is_unidirectional,
@@ -665,51 +699,51 @@ Status ComputeSoftmaxWithRawMask(cudaStream_t stream,
const int blockSize = 32;
SoftmaxWithRawMaskSmallKernel
<<>>(all_sequence_length, sequence_length,
- attention_mask, key_padding_mask, add_before_softmax, input, out,
- is_unidirectional, rsqrt_head_size, mask_dimension, max_sequence_length,
+ attention_mask, key_padding_mask, rel_pos_bias, broadcast_rel_pos_bias, input,
+ out, is_unidirectional, rsqrt_head_size, mask_dimension, max_sequence_length,
use_persistent_softmax, mask_filter_value);
} else if (all_sequence_length <= 64) {
const int blockSize = 64;
SoftmaxWithRawMaskSmallKernel
<<>>(all_sequence_length, sequence_length,
- attention_mask, key_padding_mask, add_before_softmax, input, out,
- is_unidirectional, rsqrt_head_size, mask_dimension, max_sequence_length,
+ attention_mask, key_padding_mask, rel_pos_bias, broadcast_rel_pos_bias, input,
+ out, is_unidirectional, rsqrt_head_size, mask_dimension, max_sequence_length,
use_persistent_softmax, mask_filter_value);
} else if (all_sequence_length <= 128) {
const int blockSize = 128;
SoftmaxWithRawMaskSmallKernel
<<>>(all_sequence_length, sequence_length,
- attention_mask, key_padding_mask, add_before_softmax, input, out,
- is_unidirectional, rsqrt_head_size, mask_dimension, max_sequence_length,
+ attention_mask, key_padding_mask, rel_pos_bias, broadcast_rel_pos_bias, input,
+ out, is_unidirectional, rsqrt_head_size, mask_dimension, max_sequence_length,
use_persistent_softmax, mask_filter_value);
} else if (all_sequence_length <= 256) {
const int blockSize = 256;
SoftmaxWithRawMaskSmallKernel
<<>>(all_sequence_length, sequence_length,
- attention_mask, key_padding_mask, add_before_softmax, input, out,
- is_unidirectional, rsqrt_head_size, mask_dimension, max_sequence_length,
+ attention_mask, key_padding_mask, rel_pos_bias, broadcast_rel_pos_bias, input,
+ out, is_unidirectional, rsqrt_head_size, mask_dimension, max_sequence_length,
use_persistent_softmax, mask_filter_value);
} else if (all_sequence_length <= 512) {
const int blockSize = 512;
SoftmaxWithRawMaskSmallKernel
<<>>(all_sequence_length, sequence_length,
- attention_mask, key_padding_mask, add_before_softmax, input, out,
- is_unidirectional, rsqrt_head_size, mask_dimension, max_sequence_length,
+ attention_mask, key_padding_mask, rel_pos_bias, broadcast_rel_pos_bias, input,
+ out, is_unidirectional, rsqrt_head_size, mask_dimension, max_sequence_length,
use_persistent_softmax, mask_filter_value);
} else if (all_sequence_length <= 1024) {
const int blockSize = 1024;
SoftmaxWithRawMaskSmallKernel
<<>>(all_sequence_length, sequence_length,
- attention_mask, key_padding_mask, add_before_softmax, input, out,
- is_unidirectional, rsqrt_head_size, mask_dimension, max_sequence_length,
+ attention_mask, key_padding_mask, rel_pos_bias, broadcast_rel_pos_bias, input,
+ out, is_unidirectional, rsqrt_head_size, mask_dimension, max_sequence_length,
use_persistent_softmax, mask_filter_value);
} else {
const int blockSize = 256;
const int sh_bytes = sizeof(float) * all_sequence_length;
SoftmaxWithRawMaskLargeKernel
<<>>(all_sequence_length, sequence_length,
- attention_mask, key_padding_mask, add_before_softmax, input, out,
- is_unidirectional, rsqrt_head_size, mask_dimension, max_sequence_length,
+ attention_mask, key_padding_mask, rel_pos_bias, broadcast_rel_pos_bias, input,
+ out, is_unidirectional, rsqrt_head_size, mask_dimension, max_sequence_length,
use_persistent_softmax, mask_filter_value);
}
diff --git a/onnxruntime/contrib_ops/cuda/bert/multihead_attention.cc b/onnxruntime/contrib_ops/cuda/bert/multihead_attention.cc
index d87f122045..06bf815389 100644
--- a/onnxruntime/contrib_ops/cuda/bert/multihead_attention.cc
+++ b/onnxruntime/contrib_ops/cuda/bert/multihead_attention.cc
@@ -42,6 +42,8 @@ MultiHeadAttention::MultiHeadAttention(const OpKernelInfo& info)
mask_filter_value_ = info.GetAttrOrDefault("mask_filter_value", -10000.0f);
+ scale_ = info.GetAttrOrDefault("scale", 0.0f);
+
disable_fused_self_attention_ = sizeof(T) != 2 ||
ParseEnvironmentVariableWithDefault(attention::kDisableFusedSelfAttention, false);
@@ -73,6 +75,8 @@ Status MultiHeadAttention::ComputeInternal(OpKernelContext* context) const {
const Tensor* bias = context->Input(3);
const Tensor* key_padding_mask = context->Input(4);
const Tensor* relative_position_bias = context->Input(5);
+ const Tensor* past_key = context->Input(6);
+ const Tensor* past_value = context->Input(7);
auto& device_prop = GetDeviceProp();
AttentionParameters parameters;
@@ -82,9 +86,12 @@ Status MultiHeadAttention::ComputeInternal(OpKernelContext* context) const {
bias,
key_padding_mask,
relative_position_bias,
+ past_key,
+ past_value,
¶meters,
num_heads_,
mask_filter_value_,
+ scale_,
device_prop.maxThreadsPerBlock));
int sequence_length = parameters.sequence_length;
@@ -95,6 +102,12 @@ Status MultiHeadAttention::ComputeInternal(OpKernelContext* context) const {
output_shape[2] = static_cast(parameters.v_hidden_size);
Tensor* output = context->Output(0, output_shape);
+ std::vector present_dims{
+ parameters.batch_size, parameters.num_heads, parameters.total_sequence_length, parameters.head_size};
+ TensorShape present_shape(present_dims);
+ Tensor* present_key = context->Output(1, present_shape);
+ Tensor* present_value = context->Output(2, present_shape);
+
MHARunner* fused_runner = nullptr;
const FusedMultiHeadCrossAttentionKernel* fused_cross_attention_kernel = nullptr;
@@ -107,6 +120,7 @@ Status MultiHeadAttention::ComputeInternal(OpKernelContext* context) const {
bool use_fused_cross_attention = !disable_fused_cross_attention_ &&
nullptr == key_padding_mask &&
nullptr == relative_position_bias &&
+ (nullptr == past_key && nullptr == past_value && !parameters.pass_past_in_kv) &&
key != nullptr &&
(value != nullptr || bias == nullptr) && // TODO: new kernel for adding bias to packed KV
parameters.hidden_size == parameters.v_hidden_size &&
@@ -128,6 +142,7 @@ Status MultiHeadAttention::ComputeInternal(OpKernelContext* context) const {
fused_cross_attention_kernel == nullptr &&
nullptr == relative_position_bias &&
(value != nullptr || key == nullptr) &&
+ (nullptr == past_key && nullptr == past_value && !parameters.pass_past_in_kv) &&
(nullptr == key_padding_mask || is_mask_1d_seq_len) &&
parameters.hidden_size == parameters.v_hidden_size &&
parameters.sequence_length == parameters.kv_sequence_length &&
@@ -195,16 +210,22 @@ Status MultiHeadAttention::ComputeInternal(OpKernelContext* context) const {
data.gemm_buffer = nullptr;
data.bias = (nullptr == bias) ? nullptr : reinterpret_cast(bias->Data());
data.query = reinterpret_cast(query->Data());
- data.key = (nullptr == key) ? nullptr : reinterpret_cast(key->Data());
- data.value = (nullptr == value) ? nullptr : reinterpret_cast(value->Data());
+ data.key = (nullptr == key || parameters.pass_past_in_kv) ? nullptr : reinterpret_cast(key->Data());
+ data.value = (nullptr == value || parameters.pass_past_in_kv) ? nullptr : reinterpret_cast(value->Data());
data.mask_index = (nullptr == key_padding_mask) ? nullptr : key_padding_mask->Data();
data.mask_index_dims = (nullptr == key_padding_mask) ? gsl::span() : key_padding_mask->Shape().GetDims();
data.past = nullptr;
+ data.past_key = (parameters.pass_past_in_kv) ? reinterpret_cast(key->Data())
+ : (nullptr == past_key) ? nullptr : reinterpret_cast(past_key->Data());
+ data.past_value = (parameters.pass_past_in_kv) ? reinterpret_cast(value->Data())
+ : (nullptr == past_value) ? nullptr : reinterpret_cast(past_value->Data());
data.relative_position_bias = (nullptr == relative_position_bias) ? nullptr : reinterpret_cast(relative_position_bias->Data());
data.has_qkv_workspace = !no_qkv_workspace;
data.workspace = reinterpret_cast(work_space.get());
data.output = reinterpret_cast(output->MutableData());
data.present = nullptr;
+ data.present_key = (nullptr == present_key) ? nullptr : reinterpret_cast(present_key->MutableData());
+ data.present_value = (nullptr == present_value) ? nullptr : reinterpret_cast(present_value->MutableData());
data.fused_runner = reinterpret_cast(fused_runner);
data.fused_cross_attention_kernel = fused_cross_attention_kernel;
data.use_memory_efficient_attention = use_memory_efficient_attention;
diff --git a/onnxruntime/contrib_ops/cuda/bert/multihead_attention.h b/onnxruntime/contrib_ops/cuda/bert/multihead_attention.h
index b9cf271db8..af5045e70d 100644
--- a/onnxruntime/contrib_ops/cuda/bert/multihead_attention.h
+++ b/onnxruntime/contrib_ops/cuda/bert/multihead_attention.h
@@ -24,6 +24,7 @@ class MultiHeadAttention final : public CudaKernel {
protected:
int num_heads_; // number of attention heads
float mask_filter_value_;
+ float scale_;
bool disable_fused_self_attention_;
bool enable_trt_flash_attention_;
bool disable_fused_cross_attention_;
diff --git a/onnxruntime/contrib_ops/cuda/quantization/attention_quantization.cc b/onnxruntime/contrib_ops/cuda/quantization/attention_quantization.cc
index 90ec1a35ac..c155e0b641 100644
--- a/onnxruntime/contrib_ops/cuda/quantization/attention_quantization.cc
+++ b/onnxruntime/contrib_ops/cuda/quantization/attention_quantization.cc
@@ -200,11 +200,15 @@ Status QAttention::ComputeInternal(OpKernelContext* context) const {
data.mask_index = (nullptr == mask_index) ? nullptr : mask_index->Data();
data.mask_index_dims = (nullptr == mask_index) ? gsl::span() : mask_index->Shape().GetDims();
data.past = (nullptr == past_tensor) ? nullptr : reinterpret_cast(past_tensor->Data());
+ data.past_key = nullptr;
+ data.past_value = nullptr;
data.relative_position_bias = nullptr; // add_qk is not supported in quantized attention
data.has_qkv_workspace = true;
data.workspace = reinterpret_cast(work_space.get());
data.output = reinterpret_cast(output->MutableData());
data.present = (nullptr == present) ? nullptr : reinterpret_cast(present->MutableData());
+ data.present_key = nullptr;
+ data.present_value = nullptr;
data.fused_runner = fused_runner;
data.fused_cross_attention_kernel = nullptr;
data.use_memory_efficient_attention = use_memory_efficient_attention;
diff --git a/onnxruntime/contrib_ops/cuda/transformers/generation_device_helper.cc b/onnxruntime/contrib_ops/cuda/transformers/generation_device_helper.cc
index b5ae73c829..e4d75971c9 100644
--- a/onnxruntime/contrib_ops/cuda/transformers/generation_device_helper.cc
+++ b/onnxruntime/contrib_ops/cuda/transformers/generation_device_helper.cc
@@ -997,8 +997,8 @@ Status UpdateDecoderFeeds(
for (int i = 0; i < num_present_tensors; ++i) {
next_inputs[t5_decoder_first_past_input_idx + i] =
last_outputs[t5_decoder_first_present_output_idx + i];
- return Status::OK();
}
+ return Status::OK();
}
return PickT5PastState(last_outputs, next_inputs, num_present_tensors, beam_indices, allocator,
diff --git a/onnxruntime/core/graph/contrib_ops/bert_defs.cc b/onnxruntime/core/graph/contrib_ops/bert_defs.cc
index 75d8a652ba..0b1d319272 100644
--- a/onnxruntime/core/graph/contrib_ops/bert_defs.cc
+++ b/onnxruntime/core/graph/contrib_ops/bert_defs.cc
@@ -133,6 +133,11 @@ void MultiHeadAttentionTypeAndShapeInference(ONNX_NAMESPACE::InferenceContext& c
// Input 1 (key) has shape (batch_size, kv_sequence_length, hidden_size)
// Input 2 (value) has shape (batch_size, kv_sequence_length, v_hidden_size)
+ // Q, K and V without packing and past (cross attention):
+ // Input 0 (query) has shape (batch_size, sequence_length, hidden_size)
+ // Input 1 (key) has shape (batch_size, num_head, kv_sequence_length, head_size)
+ // Input 2 (value) has shape (batch_size, num_head, kv_sequence_length, head_size)
+
// Packed KV:
// Input 0 (query) has shape (batch_size, sequence_length, hidden_size)
// Input 1 (batch_size, kv_sequence_length, num_heads, 2, head_size)
@@ -167,14 +172,14 @@ void MultiHeadAttentionTypeAndShapeInference(ONNX_NAMESPACE::InferenceContext& c
if (hasInputShape(ctx, 2)) {
auto& value_shape = getInputShape(ctx, 2);
auto& value_dims = value_shape.dim();
- if (value_dims.size() != 3) {
- fail_shape_inference("Inputs 2 (value) shall be 3 dimensions");
+ if (value_dims.size() != 3 && value_dims.size() != 4) {
+ fail_shape_inference("Inputs 2 (value) shall be 3 or 4 dimensions");
}
ONNX_NAMESPACE::TensorShapeProto output_shape;
*output_shape.add_dim() = query_dims[0];
*output_shape.add_dim() = query_dims[1];
- *output_shape.add_dim() = value_dims[2];
+ *output_shape.add_dim() = value_dims.size() == 3 ? value_dims[2] : value_dims[1] * value_dims[3];
updateOutputShape(ctx, 0, output_shape);
return;
}
@@ -261,7 +266,8 @@ ONNX_MS_OPERATOR_SET_SCHEMA(
.Input(2,
"bias",
"Bias tensor with shape (hidden_size + hidden_size + v_hidden_size) for input projection",
- "T")
+ "T",
+ OpSchema::Optional)
.Input(3,
"mask_index",
"Attention mask with shape (batch_size, 1, max_sequence_length, max_sequence_length), "
@@ -412,18 +418,23 @@ ONNX_MS_OPERATOR_SET_SCHEMA(
.Attr("num_heads", "Number of attention heads", AttributeProto::INT)
.Attr("mask_filter_value", "The value to be filled in the attention mask. Default value is -10000.0f",
AttributeProto::FLOAT, OPTIONAL_VALUE)
+ .Attr("scale",
+ "Custom scale will be used if specified. Default value is 1/sqrt(head_size)",
+ AttributeProto::FLOAT,
+ OPTIONAL_VALUE)
.Input(0,
"query",
"Query with shape (batch_size, sequence_length, hidden_size), or packed QKV with shape (batch_size, kv_sequence_length, num_heads, 3, head_size)",
"T")
.Input(1,
"key",
- "Key with shape (batch_size, kv_sequence_length, hidden_size), or packed KV with shape (batch_size, kv_sequence_length, num_heads, 2, head_size)",
+ "Key with shape (batch_size, kv_sequence_length, hidden_size), or packed KV with shape (batch_size, kv_sequence_length, num_heads, 2, head_size), "
+ "or past_key with shape (batch_size, num_heads, kv_sequence_length, head_size)",
"T",
OpSchema::Optional)
.Input(2,
"value",
- "Value with shape (batch_size, kv_sequence_length, v_hidden_size)",
+ "Value with shape (batch_size, kv_sequence_length, v_hidden_size), or past_value with shape (batch_size, num_heads, kv_sequence_length, head_size)",
"T",
OpSchema::Optional)
.Input(3,
@@ -442,10 +453,32 @@ ONNX_MS_OPERATOR_SET_SCHEMA(
" or (1, num_heads, sequence_length, total_sequence_length)",
"T",
OpSchema::Optional)
+ .Input(6,
+ "past_key",
+ "past state for self attention key with shape (batch_size, num_heads, past_sequence_length, head_size)",
+ "T",
+ OpSchema::Optional)
+ .Input(7,
+ "past_value",
+ "past state for self attention value with shape (batch_size, num_heads, past_sequence_length, head_size)",
+ "T",
+ OpSchema::Optional)
.Output(0,
"output",
"3D output tensor with shape (batch_size, sequence_length, v_hidden_size)",
"T")
+ .Output(1,
+ "present_key",
+ "present state for cross attention key with shape (batch_size, num_heads, kv_sequence_length, head_size)"
+ "or present state for self attention key with shape (batch_size, num_heads, total_sequence_length, head_size)",
+ "T",
+ OpSchema::Optional)
+ .Output(2,
+ "present_value",
+ "present state for cross attention value with shape (batch_size, num_heads, kv_sequence_length, head_size)"
+ "or present state for self attention value with shape (batch_size, num_heads, total_sequence_length, head_size)",
+ "T",
+ OpSchema::Optional)
.TypeConstraint("T", {"tensor(float)", "tensor(float16)"}, "Constrain input and output to float tensors.")
.TypeConstraint("M", {"tensor(int32)"}, "Constrain mask to integer types")
.TypeAndShapeInferenceFunction([](ONNX_NAMESPACE::InferenceContext& ctx) {
diff --git a/onnxruntime/core/graph/contrib_ops/shape_inference_functions.cc b/onnxruntime/core/graph/contrib_ops/shape_inference_functions.cc
index 791a69412f..9c9e73bf1f 100644
--- a/onnxruntime/core/graph/contrib_ops/shape_inference_functions.cc
+++ b/onnxruntime/core/graph/contrib_ops/shape_inference_functions.cc
@@ -122,9 +122,16 @@ void AttentionTypeAndShapeInference(ONNX_NAMESPACE::InferenceContext& ctx, int p
// When past and present share buffer, they shape is same: (2, batch_size, num_heads, max_sequence_length, head_size)
// Type inference
- ONNX_NAMESPACE::propagateElemTypeFromInputToOutput(ctx, 2, 0);
- if (ctx.getNumOutputs() > 1) {
- ONNX_NAMESPACE::propagateElemTypeFromInputToOutput(ctx, 2, 1);
+ if (hasInputShape(ctx, 2)) {
+ ONNX_NAMESPACE::propagateElemTypeFromInputToOutput(ctx, 2, 0);
+ if (ctx.getNumOutputs() > 1) {
+ ONNX_NAMESPACE::propagateElemTypeFromInputToOutput(ctx, 2, 1);
+ }
+ } else {
+ ONNX_NAMESPACE::propagateElemTypeFromInputToOutput(ctx, 0, 0);
+ if (ctx.getNumOutputs() > 1) {
+ ONNX_NAMESPACE::propagateElemTypeFromInputToOutput(ctx, 0, 1);
+ }
}
// Shape inference
diff --git a/onnxruntime/test/contrib_ops/attention_op_test_helper.cc b/onnxruntime/test/contrib_ops/attention_op_test_helper.cc
index 802df2b1d4..5c7f9dfab0 100644
--- a/onnxruntime/test/contrib_ops/attention_op_test_helper.cc
+++ b/onnxruntime/test/contrib_ops/attention_op_test_helper.cc
@@ -2908,6 +2908,188 @@ void GetCrossAttentionData_HeadSize16(AttentionTestData& data) {
}
}
+void GetCrossAttentionDataWithPast(AttentionTestData& data) {
+ data.hidden_size = 8;
+ data.v_hidden_size = 8;
+ data.num_heads = 2;
+ data.batch_size = 1;
+ data.sequence_length = 2;
+ data.kv_sequence_length = 3;
+ data.mask_type = AttentionMaskType::MASK_2D_KEY_PADDING;
+ data.key_padding_mask_data = {1, 1, 1};
+
+ data.skip_kernel_types = {
+ AttentionKernelType::AttentionKernel_TrtFlashAttention,
+ AttentionKernelType::AttentionKernel_TrtFusedCrossAttention,
+ AttentionKernelType::AttentionKernel_TrtFusedAttention,
+ AttentionKernelType::AttentionKernel_CutlassMemoryEfficientAttention};
+
+ {
+ data.query_data = {
+ -0.10939738f, -0.11916742f, -0.23157823f, -0.12894472f,
+ -0.02661306f, 0.26251313f, 0.30725253f, -0.34759378f,
+ -0.11695808f, -0.13129434f, -0.17031054f, -0.14986445f,
+ -0.02826184f, 0.2797631f , 0.27337456f, -0.44312602f
+ };
+ }
+ // The past key and value data will be passed to the kernel as input 'key' and 'value'.
+ {
+ data.past_key_data = {
+ 0.5967375f , 0.5966938f , 0.48602432f, 0.5341031f,
+ 0.55797786f, 0.5663399f , 0.57087725f, 0.6240304f,
+ 0.5352563f , 0.5648297f , 0.4972945f , 0.56637144f,
+
+ 0.44123724f, 0.35872823f, 0.32176313f, 0.4490301f,
+ 0.3643952f , 0.51968557f, 0.50137347f, 0.5743993f,
+ 0.3905106f , 0.4741712f , 0.40881708f, 0.47243845f
+ };
+ }
+
+ {
+ data.past_value_data = {
+ 0.40251260f, 0.55487730f, 0.49565578f, 0.42683450f,
+ 0.44379145f, 0.58945787f, 0.54852820f, 0.43376005f,
+ 0.44116694f, 0.44007313f, 0.40293324f, 0.53202707f,
+
+ 0.35520583f, 0.47293650f, 0.45417705f, 0.33723440f,
+ 0.50175804f, 0.37620395f, 0.24103148f, 0.50958070f,
+ 0.56803876f, 0.37866923f, 0.32273075f, 0.44389135f
+ };
+ }
+
+ {
+ data.fp32_output_data = {
+ 0.4291f, 0.5275f, 0.4818f, 0.4645f, 0.4770f, 0.4082f, 0.3372f, 0.4319f,
+ 0.4291f, 0.5276f, 0.4818f, 0.4645f, 0.4768f, 0.4083f, 0.3377f, 0.4315f
+ };
+ }
+
+ {
+ data.fp16_output_data = data.fp32_output_data;
+ }
+}
+
+void GetSelfAttentionDataWithPast(AttentionTestData& data) {
+ data.hidden_size = 8;
+ data.v_hidden_size = 8;
+ data.num_heads = 2;
+ data.batch_size = 1;
+ data.sequence_length = 2;
+ data.kv_sequence_length = 3;
+ data.mask_type = AttentionMaskType::MASK_NONE;
+
+ data.skip_kernel_types = {
+ AttentionKernelType::AttentionKernel_TrtFlashAttention,
+ AttentionKernelType::AttentionKernel_TrtFusedCrossAttention,
+ AttentionKernelType::AttentionKernel_TrtFusedAttention,
+ AttentionKernelType::AttentionKernel_CutlassMemoryEfficientAttention};
+
+ {
+ data.query_data = {
+ 0.00403503f, 0.08716156f, -0.0358175f , -0.08171791f,
+ 0.48912194f, -0.22679007f, -0.09093101f, -0.5939322f,
+ 0.00878838f, 0.03355761f, -0.08080226f, -0.06677517f,
+ 0.55038965f, -0.2720567f , -0.12977877f, -0.634123f
+ };
+ }
+ {
+ data.key_data = {
+ 0.2808786f , 0.10041683f, 0.15880886f, 0.45283064f,
+ 0.39884242f, 0.12596075f, 0.4198916f , -0.0651141f,
+ 0.31678027f, 0.11010794f, 0.21594375f, 0.4975329f,
+ 0.436772f , 0.20940652f, 0.44072092f, -0.05601776f
+ };
+ }
+
+ {
+ data.value_data = {
+ 0.26421773f, -0.16541699f, -0.0599675f , 0.27200517f,
+ -0.1074627f , -0.4493224f , -0.03694462f, 0.17997989f,
+ 0.27960598f, -0.16643806f, -0.07019104f, 0.29006317f,
+ -0.11640988f, -0.47876123f, -0.01979145f, 0.11468418f
+ };
+ }
+
+ {
+ data.rel_pos_bias_data = {
+ 0.4781123f , 0.82420444f, 0.654424f , 0.3995186f , 0.5482078f,
+ 0.55570245f, 0.4216576f , 0.46001542f, 0.67183703f, 0.41973996f,
+
+ 0.28494194f, 0.60367906f, 0.3453173f , 0.44483483f, 0.6770777f,
+ 0.5460559f , 0.31994605f, 0.5470492f , 0.5433419f , 0.60349935f
+ };
+ }
+
+ {
+ data.past_key_data = {
+ 0.34734827f, 0.5592256f , 0.5333037f , 0.5122027f,
+ 0.5940516f , 0.44744077f, 0.43128848f, 0.55360645f,
+ 0.57874715f, 0.29512063f, 0.2780432f , 0.4693917f,
+
+ 0.4450266f , 0.530704f , 0.3124955f , 0.4273598f,
+ 0.44368753f, 0.5890438f , 0.5054336f , 0.46042535f,
+ 0.5352153f , 0.5157861f , 0.39744973f, 0.5441864f
+ };
+ }
+
+ {
+ data.past_value_data = {
+ 0.48998538f, 0.5493853f , 0.556647f , 0.7011929f,
+ 0.543909f , 0.5630743f , 0.5087797f , 0.3901024f,
+ 0.53116417f, 0.4086225f , 0.5320247f , 0.5145377f,
+
+ 0.4086198f , 0.6913348f , 0.50045484f, 0.5338214f,
+ 0.52980417f, 0.5243695f , 0.6046111f , 0.53555113f,
+ 0.44936907f, 0.6010697f , 0.38031512f, 0.427301f
+ };
+ }
+
+ {
+ data.fp32_output_data = {
+ 0.4358f, 0.2708f, 0.3201f, 0.4347f, 0.1886f, 0.0845f, 0.2479f, 0.3289f,
+ 0.4157f, 0.2247f, 0.2826f, 0.4321f, 0.1874f, 0.1021f, 0.2427f, 0.3305f
+ };
+ }
+
+ {
+ data.fp16_output_data = data.fp32_output_data;
+ }
+
+ {
+ data.present_key_data = {
+ 0.3473f, 0.5592f, 0.5333f, 0.5122f,
+ 0.5941f, 0.4474f, 0.4313f, 0.5536f,
+ 0.5787f, 0.2951f, 0.2780f, 0.4694f,
+ 0.2809f, 0.1004f, 0.1588f, 0.4528f,
+ 0.3168f, 0.1101f, 0.2159f, 0.4975f,
+
+ 0.4450f, 0.5307f, 0.3125f, 0.4274f,
+ 0.4437f, 0.5890f, 0.5054f, 0.4604f,
+ 0.5352f, 0.5158f, 0.3974f, 0.5442f,
+ 0.3988f, 0.1260f, 0.4199f, -0.0651f,
+ 0.4368f, 0.2094f, 0.4407f, -0.0560f
+ };
+ }
+
+ {
+ data.present_value_data = {
+ 0.4900f, 0.5494f, 0.5566f, 0.7012f,
+ 0.5439f, 0.5631f, 0.5088f, 0.3901f,
+ 0.5312f, 0.4086f, 0.5320f, 0.5145f,
+ 0.2642f, -0.1654f, -0.0600f, 0.2720f,
+ 0.2796f, -0.1664f, -0.0702f, 0.2901f,
+
+ 0.4086f, 0.6913f, 0.5005f, 0.5338f,
+ 0.5298f, 0.5244f, 0.6046f, 0.5356f,
+ 0.4494f, 0.6011f, 0.3803f, 0.4273f,
+ -0.1075f, -0.4493f, -0.0369f, 0.1800f,
+ -0.1164f, -0.4788f, -0.0198f, 0.1147f
+ };
+ }
+
+ data.is_static_kv = false;
+}
+
bool SkipAttentionKernel(AttentionTestData& data, AttentionKernelType kernel_type) {
return std::find(data.skip_kernel_types.begin(), data.skip_kernel_types.end(), kernel_type) != data.skip_kernel_types.end();
}
diff --git a/onnxruntime/test/contrib_ops/attention_op_test_helper.h b/onnxruntime/test/contrib_ops/attention_op_test_helper.h
index 3a00b08b45..807e1e207d 100644
--- a/onnxruntime/test/contrib_ops/attention_op_test_helper.h
+++ b/onnxruntime/test/contrib_ops/attention_op_test_helper.h
@@ -11,6 +11,7 @@ using contrib::AttentionMaskType;
namespace test {
struct AttentionTestData {
+ bool is_static_kv = true;
int hidden_size;
int v_hidden_size;
int num_heads;
@@ -27,8 +28,17 @@ struct AttentionTestData {
std::vector qkv_data;
std::vector bias_data;
+ std::vector rel_pos_bias_data;
+
+ std::vector past_key_data;
+ std::vector past_value_data;
+
std::vector fp32_output_data;
std::vector fp16_output_data;
+
+ std::vector present_key_data;
+ std::vector present_value_data;
+
std::vector skip_kernel_types; // skip some kernels if they do not supported this test case.
};
@@ -49,6 +59,9 @@ void GetSelfAttentionData_Batch2_HeadSize32_NoBias_NoMask_PackedQKV(AttentionTes
void GetCrossAttentionData_HeadSize16_8(AttentionTestData& data);
void GetCrossAttentionData_HeadSize16(AttentionTestData& data);
+void GetCrossAttentionDataWithPast(AttentionTestData& data);
+void GetSelfAttentionDataWithPast(AttentionTestData& data);
+
bool SkipAttentionKernel(AttentionTestData& data, AttentionKernelType kernel_type);
} // namespace test
} // namespace onnxruntime
diff --git a/onnxruntime/test/contrib_ops/multihead_attention_op_test.cc b/onnxruntime/test/contrib_ops/multihead_attention_op_test.cc
index 646f898ed0..415a1c6f8f 100644
--- a/onnxruntime/test/contrib_ops/multihead_attention_op_test.cc
+++ b/onnxruntime/test/contrib_ops/multihead_attention_op_test.cc
@@ -19,6 +19,11 @@ static void RunMultiHeadAttentionTest(
const std::vector& kv_data, // packed_kv: [batch_size, kv_sequence_length, num_heads, 2, head_size]
const std::vector& qkv_data, // packed_qkv: [batch_size, sequence_length, num_heads, 3, head_size]
const std::vector& bias_data, // bias: [hidden_size + hidden_size + v_hidden_size] or empty
+ const std::vector& rel_pos_bias_data, // relative_position_bias: [1, num_heads, sequence_length, total_sequence_length]
+ const std::vector& past_key_data, // past_key: [batch_size, num_heads, kv_sequence_length, head_size]
+ const std::vector& past_value_data, // past_value: [batch_size, num_heads, kv_sequence_length, head_size]
+ const std::vector& present_key_data, // present_key: [batch_size, num_heads, total_sequence_length, head_size]
+ const std::vector& present_value_data, // present_value: [batch_size, num_heads, total_sequence_length, head_size]
const std::vector& key_padding_mask_data, // key_padding_mask: see below
AttentionMaskType mask_type, // 1 for [batch_size], 2 for [batch_size, kv_sequence_length]
const std::vector& output_data, // output: [batch_size, sequence_length, v_hidden_size]
@@ -28,6 +33,7 @@ static void RunMultiHeadAttentionTest(
int kv_sequence_length,
int hidden_size,
int v_hidden_size,
+ bool is_static_kv = true,
bool use_float16 = false,
bool disable_cpu = true, // not supported in cpu right now.
bool disable_cuda = false,
@@ -46,10 +52,16 @@ static void RunMultiHeadAttentionTest(
tester.AddAttribute("mask_filter_value", static_cast(-10000.0f));
std::vector query_dims = {batch_size, sequence_length, hidden_size};
- std::vector key_dims = {batch_size, kv_sequence_length, hidden_size};
- std::vector value_dims = {batch_size, kv_sequence_length, v_hidden_size};
+ std::vector key_dims = {batch_size, is_static_kv ? kv_sequence_length : sequence_length, hidden_size};
+ std::vector value_dims = {batch_size, is_static_kv ? kv_sequence_length : sequence_length, v_hidden_size};
std::vector bias_dims = {hidden_size + hidden_size + v_hidden_size};
+ std::vector rel_pos_bias_dims = {1, num_heads, sequence_length, sequence_length + kv_sequence_length};
+ std::vector past_key_dims = {batch_size, num_heads, kv_sequence_length, hidden_size / num_heads};
+ std::vector past_value_dims = past_key_dims;
std::vector output_dims = {batch_size, sequence_length, v_hidden_size};
+ std::vector present_key_dims =
+ {batch_size, num_heads, is_static_kv ? kv_sequence_length : sequence_length + kv_sequence_length, hidden_size / num_heads};
+ std::vector present_value_dims = present_key_dims;
std::vector query = (qkv_data.size() > 0 ? qkv_data : query_data);
std::vector key;
@@ -79,12 +91,16 @@ static void RunMultiHeadAttentionTest(
if (key.size()) {
tester.AddInput("key", key_dims, ToFloat16(key));
+ } else if (past_key_data.size() && is_static_kv == true) {
+ tester.AddInput("key", past_key_dims, ToFloat16(past_key_data));
} else {
tester.AddOptionalInputEdge();
}
if (value.size()) {
tester.AddInput("value", value_dims, ToFloat16(value));
+ } else if (past_value_data.size() && is_static_kv == true) {
+ tester.AddInput("value", past_value_dims, ToFloat16(past_value_data));
} else {
tester.AddOptionalInputEdge();
}
@@ -101,22 +117,56 @@ static void RunMultiHeadAttentionTest(
tester.AddOptionalInputEdge();
}
+ if (rel_pos_bias_data.size()) {
+ tester.AddInput("relative_position_bias", rel_pos_bias_dims, ToFloat16(rel_pos_bias_data));
+ } else {
+ tester.AddOptionalInputEdge();
+ }
+
+ if (past_key_data.size() && is_static_kv == false) {
+ tester.AddInput("past_key", past_key_dims, ToFloat16(past_key_data));
+ } else {
+ tester.AddOptionalInputEdge();
+ }
+
+ if (past_value_data.size() && is_static_kv == false) {
+ tester.AddInput("past_value", past_value_dims, ToFloat16(past_value_data));
+ } else {
+ tester.AddOptionalInputEdge();
+ }
+
constexpr float rel_error = 0.0f;
constexpr float abs_error = 0.05f;
tester.AddOutput("output", output_dims, ToFloat16(output_data), /*sort*/ false, rel_error, abs_error);
+
+ if (present_key_data.size()) {
+ tester.AddOutput("present_key", present_key_dims, ToFloat16(present_key_data), /*sort*/ false, rel_error, abs_error);
+ } else {
+ tester.AddOptionalOutputEdge();
+ }
+
+ if (present_value_data.size()) {
+ tester.AddOutput("present_value", present_value_dims, ToFloat16(present_value_data), /*sort*/ false, rel_error, abs_error);
+ } else {
+ tester.AddOptionalOutputEdge();
+ }
} else {
tester.AddInput("query", query_dims, query);
if (key.size()) {
tester.AddInput("key", key_dims, key);
+ } else if (past_key_data.size() && is_static_kv == true) {
+ tester.AddInput("key", past_key_dims, past_key_data);
} else {
- tester.AddOptionalInputEdge();
+ tester.AddOptionalInputEdge();
}
if (value.size()) {
tester.AddInput("value", value_dims, value);
+ } else if (past_value_data.size() && is_static_kv == true) {
+ tester.AddInput("value", past_value_dims, past_value_data);
} else {
- tester.AddOptionalInputEdge();
+ tester.AddOptionalInputEdge();
}
if (bias_data.size()) {
@@ -131,9 +181,39 @@ static void RunMultiHeadAttentionTest(
tester.AddOptionalInputEdge();
}
+ if (rel_pos_bias_data.size()) {
+ tester.AddInput("relative_position_bias", rel_pos_bias_dims, rel_pos_bias_data);
+ } else {
+ tester.AddOptionalInputEdge();
+ }
+
+ if (past_key_data.size() && is_static_kv == false) {
+ tester.AddInput("past_key", past_key_dims, past_key_data);
+ } else {
+ tester.AddOptionalInputEdge();
+ }
+
+ if (past_value_data.size() && is_static_kv == false) {
+ tester.AddInput("past_value", past_value_dims, past_value_data);
+ } else {
+ tester.AddOptionalInputEdge();
+ }
+
constexpr float rel_error = 0.0f;
constexpr float abs_error = 0.02f;
tester.AddOutput("output", output_dims, output_data, /*sort*/ false, rel_error, abs_error);
+
+ if (present_key_data.size()) {
+ tester.AddOutput("present_key", present_key_dims, present_key_data, /*sort*/ false, rel_error, abs_error);
+ } else {
+ tester.AddOptionalOutputEdge();
+ }
+
+ if (present_value_data.size()) {
+ tester.AddOutput("present_value", present_value_dims, present_value_data, /*sort*/ false, rel_error, abs_error);
+ } else {
+ tester.AddOptionalOutputEdge();
+ }
}
if (enable_cuda) {
@@ -163,6 +243,11 @@ static void RunMultiHeadAttentionKernel(
const std::vector& kv_data, // packed_kv: [batch_size, kv_sequence_length, num_heads, 2, head_size]
const std::vector& qkv_data, // packed_qkv: [batch_size, sequence_length, num_heads, 3, head_size]
const std::vector& bias_data, // bias: [hidden_size + hidden_size + v_hidden_size]
+ const std::vector& rel_pos_bias_data, // relative_position_bias: [1, num_heads, sequence_length, total_sequence_length]
+ const std::vector& past_key_data, // past_key: [batch_size, num_heads, kv_sequence_length, head_size]
+ const std::vector& past_value_data, // past_value: [batch_size, num_heads, kv_sequence_length, head_size]
+ const std::vector& present_key_data, // present_key: [batch_size, num_heads, total_sequence_length, head_size]
+ const std::vector& present_value_data, // present_value: [batch_size, num_heads, total_sequence_length, head_size]
const std::vector& key_padding_mask_data, // key_padding_mask: see below
AttentionMaskType mask_type, // 1 for [batch_size], 2 for [batch_size, kv_sequence_length]
const std::vector& output_data, // output: [batch_size, sequence_length, v_hidden_size]
@@ -174,6 +259,7 @@ static void RunMultiHeadAttentionKernel(
int v_hidden_size,
AttentionKernelType kernel_type,
bool use_float16 = true,
+ bool is_static_kv = true,
bool disable_cpu = true, // not supported in cpu right now.
bool disable_cuda = false,
bool disable_rocm = true) {
@@ -185,9 +271,10 @@ static void RunMultiHeadAttentionKernel(
{onnxruntime::contrib::attention::kDisableFusedCrossAttention, "0"},
{onnxruntime::contrib::attention::kDisableMemoryEfficientAttention, "0"}}};
RunMultiHeadAttentionTest(
- query_data, key_data, value_data, kv_data, qkv_data, bias_data, key_padding_mask_data, mask_type, output_data,
- num_heads, batch_size, sequence_length, kv_sequence_length, hidden_size, v_hidden_size,
- use_float16, disable_cpu, disable_cuda, disable_rocm);
+ query_data, key_data, value_data, kv_data, qkv_data, bias_data, rel_pos_bias_data,
+ past_key_data, past_value_data, present_key_data, present_value_data, key_padding_mask_data,
+ mask_type, output_data, num_heads, batch_size, sequence_length, kv_sequence_length,
+ hidden_size, v_hidden_size, is_static_kv, use_float16, disable_cpu, disable_cuda, disable_rocm);
return;
}
@@ -199,9 +286,10 @@ static void RunMultiHeadAttentionKernel(
{onnxruntime::contrib::attention::kDisableFusedCrossAttention, "1"},
{onnxruntime::contrib::attention::kDisableMemoryEfficientAttention, "1"}}};
RunMultiHeadAttentionTest(
- query_data, key_data, value_data, kv_data, qkv_data, bias_data, key_padding_mask_data, mask_type, output_data,
- num_heads, batch_size, sequence_length, kv_sequence_length, hidden_size, v_hidden_size,
- use_float16, disable_cpu, disable_cuda, disable_rocm);
+ query_data, key_data, value_data, kv_data, qkv_data, bias_data, rel_pos_bias_data,
+ past_key_data, past_value_data, present_key_data, present_value_data, key_padding_mask_data,
+ mask_type, output_data, num_heads, batch_size, sequence_length, kv_sequence_length,
+ hidden_size, v_hidden_size, is_static_kv, use_float16, disable_cpu, disable_cuda, disable_rocm);
return;
}
@@ -213,9 +301,10 @@ static void RunMultiHeadAttentionKernel(
{onnxruntime::contrib::attention::kDisableFusedCrossAttention, "0"},
{onnxruntime::contrib::attention::kDisableMemoryEfficientAttention, "1"}}};
RunMultiHeadAttentionTest(
- query_data, key_data, value_data, kv_data, qkv_data, bias_data, key_padding_mask_data, mask_type, output_data,
- num_heads, batch_size, sequence_length, kv_sequence_length, hidden_size, v_hidden_size,
- use_float16, disable_cpu, disable_cuda, disable_rocm);
+ query_data, key_data, value_data, kv_data, qkv_data, bias_data, rel_pos_bias_data,
+ past_key_data, past_value_data, present_key_data, present_value_data, key_padding_mask_data,
+ mask_type, output_data, num_heads, batch_size, sequence_length, kv_sequence_length,
+ hidden_size, v_hidden_size, is_static_kv, use_float16, disable_cpu, disable_cuda, disable_rocm);
return;
}
@@ -228,9 +317,10 @@ static void RunMultiHeadAttentionKernel(
{onnxruntime::contrib::attention::kDisableFusedCrossAttention, "1"},
{onnxruntime::contrib::attention::kDisableMemoryEfficientAttention, "0"}}};
RunMultiHeadAttentionTest(
- query_data, key_data, value_data, kv_data, qkv_data, bias_data, key_padding_mask_data, mask_type, output_data,
- num_heads, batch_size, sequence_length, kv_sequence_length, hidden_size, v_hidden_size,
- use_float16, disable_cpu, disable_cuda, disable_rocm);
+ query_data, key_data, value_data, kv_data, qkv_data, bias_data, rel_pos_bias_data,
+ past_key_data, past_value_data, present_key_data, present_value_data, key_padding_mask_data,
+ mask_type, output_data, num_heads, batch_size, sequence_length, kv_sequence_length,
+ hidden_size, v_hidden_size, is_static_kv, use_float16, disable_cpu, disable_cuda, disable_rocm);
return;
}
#endif
@@ -243,9 +333,10 @@ static void RunMultiHeadAttentionKernel(
{onnxruntime::contrib::attention::kDisableFusedCrossAttention, "1"},
{onnxruntime::contrib::attention::kDisableMemoryEfficientAttention, "1"}}};
RunMultiHeadAttentionTest(
- query_data, key_data, value_data, kv_data, qkv_data, bias_data, key_padding_mask_data, mask_type, output_data,
- num_heads, batch_size, sequence_length, kv_sequence_length, hidden_size, v_hidden_size,
- use_float16, disable_cpu, disable_cuda, disable_rocm);
+ query_data, key_data, value_data, kv_data, qkv_data, bias_data, rel_pos_bias_data,
+ past_key_data, past_value_data, present_key_data, present_value_data, key_padding_mask_data,
+ mask_type, output_data, num_heads, batch_size, sequence_length, kv_sequence_length,
+ hidden_size, v_hidden_size, is_static_kv, use_float16, disable_cpu, disable_cuda, disable_rocm);
}
}
@@ -256,9 +347,11 @@ static void RunMultiHeadAttentionTests(AttentionTestData& data) {
AttentionKernelType kernel_type = AttentionKernelType::AttentionKernel_Unfused;
if (!SkipAttentionKernel(data, kernel_type)) {
RunMultiHeadAttentionKernel(
- data.query_data, data.key_data, data.value_data, data.kv_data, data.qkv_data, data.bias_data, data.key_padding_mask_data, data.mask_type,
- data.fp32_output_data, data.num_heads, data.batch_size, data.sequence_length, data.kv_sequence_length,
- data.hidden_size, data.v_hidden_size, kernel_type, use_float16);
+ data.query_data, data.key_data, data.value_data, data.kv_data, data.qkv_data, data.bias_data,
+ data.rel_pos_bias_data, data.past_key_data, data.past_value_data, data.present_key_data,
+ data.present_value_data, data.key_padding_mask_data, data.mask_type, data.fp32_output_data,
+ data.num_heads, data.batch_size, data.sequence_length, data.kv_sequence_length, data.hidden_size,
+ data.v_hidden_size, kernel_type, use_float16, data.is_static_kv);
}
#if USE_FLASH_ATTENTION
@@ -267,18 +360,22 @@ static void RunMultiHeadAttentionTests(AttentionTestData& data) {
kernel_type = AttentionKernelType::AttentionKernel_CutlassMemoryEfficientAttention;
if (!SkipAttentionKernel(data, kernel_type)) {
RunMultiHeadAttentionKernel(
- data.query_data, data.key_data, data.value_data, data.kv_data, data.qkv_data, data.bias_data, data.key_padding_mask_data, data.mask_type,
- data.fp32_output_data, data.num_heads, data.batch_size, data.sequence_length, data.kv_sequence_length,
- data.hidden_size, data.v_hidden_size, kernel_type, use_float16);
+ data.query_data, data.key_data, data.value_data, data.kv_data, data.qkv_data, data.bias_data,
+ data.rel_pos_bias_data, data.past_key_data, data.past_value_data, data.present_key_data,
+ data.present_value_data, data.key_padding_mask_data, data.mask_type, data.fp32_output_data,
+ data.num_heads, data.batch_size, data.sequence_length, data.kv_sequence_length, data.hidden_size,
+ data.v_hidden_size, kernel_type, use_float16, data.is_static_kv);
}
}
#endif
kernel_type = AttentionKernelType::AttentionKernel_Default;
RunMultiHeadAttentionKernel(
- data.query_data, data.key_data, data.value_data, data.kv_data, data.qkv_data, data.bias_data, data.key_padding_mask_data, data.mask_type,
- data.fp32_output_data, data.num_heads, data.batch_size, data.sequence_length, data.kv_sequence_length,
- data.hidden_size, data.v_hidden_size, kernel_type, use_float16);
+ data.query_data, data.key_data, data.value_data, data.kv_data, data.qkv_data, data.bias_data,
+ data.rel_pos_bias_data, data.past_key_data, data.past_value_data, data.present_key_data,
+ data.present_value_data, data.key_padding_mask_data, data.mask_type, data.fp32_output_data,
+ data.num_heads, data.batch_size, data.sequence_length, data.kv_sequence_length, data.hidden_size,
+ data.v_hidden_size, kernel_type, use_float16, data.is_static_kv);
}
if (data.fp16_output_data.size() > 0) {
@@ -286,34 +383,42 @@ static void RunMultiHeadAttentionTests(AttentionTestData& data) {
AttentionKernelType kernel_type = AttentionKernelType::AttentionKernel_TrtFusedCrossAttention;
if (!SkipAttentionKernel(data, kernel_type)) {
RunMultiHeadAttentionKernel(
- data.query_data, data.key_data, data.value_data, data.kv_data, data.qkv_data, data.bias_data, data.key_padding_mask_data, data.mask_type,
- data.fp16_output_data, data.num_heads, data.batch_size, data.sequence_length, data.kv_sequence_length,
- data.hidden_size, data.v_hidden_size, kernel_type, use_float16);
+ data.query_data, data.key_data, data.value_data, data.kv_data, data.qkv_data, data.bias_data,
+ data.rel_pos_bias_data, data.past_key_data, data.past_value_data, data.present_key_data,
+ data.present_value_data, data.key_padding_mask_data, data.mask_type, data.fp16_output_data,
+ data.num_heads, data.batch_size, data.sequence_length, data.kv_sequence_length, data.hidden_size,
+ data.v_hidden_size, kernel_type, use_float16, data.is_static_kv);
}
kernel_type = AttentionKernelType::AttentionKernel_TrtFusedAttention;
if (!SkipAttentionKernel(data, kernel_type)) {
RunMultiHeadAttentionKernel(
- data.query_data, data.key_data, data.value_data, data.kv_data, data.qkv_data, data.bias_data, data.key_padding_mask_data, data.mask_type,
- data.fp16_output_data, data.num_heads, data.batch_size, data.sequence_length, data.kv_sequence_length,
- data.hidden_size, data.v_hidden_size, kernel_type, use_float16);
+ data.query_data, data.key_data, data.value_data, data.kv_data, data.qkv_data, data.bias_data,
+ data.rel_pos_bias_data, data.past_key_data, data.past_value_data, data.present_key_data,
+ data.present_value_data, data.key_padding_mask_data, data.mask_type, data.fp16_output_data,
+ data.num_heads, data.batch_size, data.sequence_length, data.kv_sequence_length, data.hidden_size,
+ data.v_hidden_size, kernel_type, use_float16, data.is_static_kv);
}
#if USE_FLASH_ATTENTION
kernel_type = AttentionKernelType::AttentionKernel_CutlassMemoryEfficientAttention;
if (!SkipAttentionKernel(data, kernel_type)) {
RunMultiHeadAttentionKernel(
- data.query_data, data.key_data, data.value_data, data.kv_data, data.qkv_data, data.bias_data, data.key_padding_mask_data, data.mask_type,
- data.fp16_output_data, data.num_heads, data.batch_size, data.sequence_length, data.kv_sequence_length,
- data.hidden_size, data.v_hidden_size, kernel_type, use_float16);
+ data.query_data, data.key_data, data.value_data, data.kv_data, data.qkv_data, data.bias_data,
+ data.rel_pos_bias_data, data.past_key_data, data.past_value_data, data.present_key_data,
+ data.present_value_data, data.key_padding_mask_data, data.mask_type, data.fp16_output_data,
+ data.num_heads, data.batch_size, data.sequence_length, data.kv_sequence_length, data.hidden_size,
+ data.v_hidden_size, kernel_type, use_float16, data.is_static_kv);
}
#endif
kernel_type = AttentionKernelType::AttentionKernel_Default;
RunMultiHeadAttentionKernel(
- data.query_data, data.key_data, data.value_data, data.kv_data, data.qkv_data, data.bias_data, data.key_padding_mask_data, data.mask_type,
- data.fp16_output_data, data.num_heads, data.batch_size, data.sequence_length, data.kv_sequence_length,
- data.hidden_size, data.v_hidden_size, kernel_type, use_float16);
+ data.query_data, data.key_data, data.value_data, data.kv_data, data.qkv_data, data.bias_data,
+ data.rel_pos_bias_data, data.past_key_data, data.past_value_data, data.present_key_data,
+ data.present_value_data, data.key_padding_mask_data, data.mask_type, data.fp16_output_data,
+ data.num_heads, data.batch_size, data.sequence_length, data.kv_sequence_length, data.hidden_size,
+ data.v_hidden_size, kernel_type, use_float16, data.is_static_kv);
}
}
@@ -370,5 +475,17 @@ TEST(MultiHeadAttentionTest, CrossAttention_Batch1_HeadSize16) {
RunMultiHeadAttentionTests(data);
}
+TEST(MultiHeadAttentionTest, CrossAttentionWithPast) {
+ AttentionTestData data;
+ GetCrossAttentionDataWithPast(data);
+ RunMultiHeadAttentionTests(data);
+}
+
+TEST(MultiHeadAttentionTest, SelfAttentionWithPast) {
+ AttentionTestData data;
+ GetSelfAttentionDataWithPast(data);
+ RunMultiHeadAttentionTests(data);
+}
+
} // namespace test
} // namespace onnxruntime
diff --git a/onnxruntime/test/python/transformers/test_parity_t5_mha.py b/onnxruntime/test/python/transformers/test_parity_t5_mha.py
new file mode 100644
index 0000000000..22522c9183
--- /dev/null
+++ b/onnxruntime/test/python/transformers/test_parity_t5_mha.py
@@ -0,0 +1,580 @@
+# --------------------------------------------------------------------------
+# Copyright 2020 The HuggingFace Inc. team
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at http://www.apache.org/licenses/LICENSE-2.0
+# --------------------------------------------------------------------------
+# Copyright (c) Microsoft Corporation. All rights reserved.
+# Licensed under the MIT License. See License.txt in the project root for
+# license information.
+# -------------------------------------------------------------------------
+# Copyright (c) Microsoft Corporation. All rights reserved.
+# Licensed under the MIT License.
+
+import unittest
+
+import numpy as np
+import torch
+from torch import nn
+
+torch.manual_seed(0)
+
+
+def create_t5_mha_graph(
+ batch_size,
+ seq_len,
+ kv_sequence_length,
+ head_size,
+ num_heads,
+ use_past,
+ is_static_kv,
+):
+ from onnx import TensorProto, helper
+
+ use_present = not use_past
+ if not is_static_kv and use_past:
+ use_present = True
+ use_rpb = not is_static_kv
+ use_mask = not use_rpb
+
+ past_sequence_length = kv_sequence_length
+ total_sequence_length = kv_sequence_length if is_static_kv else seq_len
+
+ if not is_static_kv:
+ kv_sequence_length = seq_len
+
+ if not is_static_kv and use_past:
+ total_sequence_length += past_sequence_length
+
+ rpb_length = total_sequence_length if use_past else seq_len
+
+ nodes = [
+ helper.make_node(
+ "MultiHeadAttention",
+ [
+ "query",
+ "key" if use_present or is_static_kv else "",
+ "value" if use_present or is_static_kv else "",
+ "", # bias
+ "key_padding_mask" if use_mask else "",
+ "relative_position_bias" if use_rpb else "",
+ "past_key" if use_past and not is_static_kv else "",
+ "past_value" if use_past and not is_static_kv else "",
+ ],
+ [
+ "output",
+ "present_key" if use_present else "",
+ "present_value" if use_present else "",
+ ],
+ "MHA_0",
+ num_heads=num_heads,
+ mask_filter_value=-10000.0,
+ scale=1.0,
+ domain="com.microsoft",
+ ),
+ ]
+
+ initializers = []
+
+ hidden_size = head_size * num_heads
+
+ graph_inputs = [
+ helper.make_tensor_value_info("query", TensorProto.FLOAT, [batch_size, seq_len, hidden_size]),
+ ]
+
+ graph_outputs = [
+ helper.make_tensor_value_info("output", TensorProto.FLOAT, [batch_size, seq_len, hidden_size]),
+ ]
+
+ if use_mask:
+ graph_inputs.append(
+ helper.make_tensor_value_info("key_padding_mask", TensorProto.INT32, [batch_size, kv_sequence_length])
+ )
+
+ if use_rpb:
+ graph_inputs.append(
+ helper.make_tensor_value_info(
+ "relative_position_bias", TensorProto.FLOAT, [1, num_heads, seq_len, rpb_length]
+ )
+ )
+
+ if use_past and not is_static_kv:
+ graph_inputs.append(
+ helper.make_tensor_value_info(
+ "past_key", TensorProto.FLOAT, [batch_size, num_heads, past_sequence_length, head_size]
+ )
+ )
+ graph_inputs.append(
+ helper.make_tensor_value_info(
+ "past_value", TensorProto.FLOAT, [batch_size, num_heads, past_sequence_length, head_size]
+ )
+ )
+
+ if use_present:
+ graph_inputs.append(
+ helper.make_tensor_value_info("key", TensorProto.FLOAT, [batch_size, kv_sequence_length, hidden_size])
+ )
+ graph_inputs.append(
+ helper.make_tensor_value_info("value", TensorProto.FLOAT, [batch_size, kv_sequence_length, hidden_size])
+ )
+ elif is_static_kv and use_past:
+ graph_inputs.append(
+ helper.make_tensor_value_info(
+ "key", TensorProto.FLOAT, [batch_size, num_heads, past_sequence_length, head_size]
+ )
+ )
+ graph_inputs.append(
+ helper.make_tensor_value_info(
+ "value", TensorProto.FLOAT, [batch_size, num_heads, past_sequence_length, head_size]
+ )
+ )
+
+ if use_present:
+ graph_outputs.append(
+ helper.make_tensor_value_info(
+ "present_key", TensorProto.FLOAT, [batch_size, num_heads, total_sequence_length, head_size]
+ )
+ )
+ graph_outputs.append(
+ helper.make_tensor_value_info(
+ "present_value", TensorProto.FLOAT, [batch_size, num_heads, total_sequence_length, head_size]
+ )
+ )
+
+ graph = helper.make_graph(
+ nodes,
+ "T5_MHA_Graph",
+ graph_inputs,
+ graph_outputs,
+ initializers,
+ )
+
+ model = helper.make_model(graph)
+ return model.SerializeToString()
+
+
+class T5Config:
+ def __init__(self, is_decoder, batch_size, seq_len, kv_sequence_length, num_heads, head_size, use_past):
+ self.is_decoder = is_decoder
+ self.d_model = num_heads * head_size
+ self.key_value_proj_dim = head_size
+ self.n_heads = num_heads
+ self.inner_dim = self.n_heads * self.key_value_proj_dim
+
+ # ORT parameters
+ self.batch_size = batch_size
+ self.seq_len = seq_len
+ self.kv_sequence_length = kv_sequence_length
+ self.head_size = head_size
+ self.num_heads = num_heads
+ self.hidden_size = self.d_model
+ self.use_past = use_past
+
+
+class T5Attention(nn.Module):
+ def __init__(self, config: T5Config, is_static_kv):
+ super().__init__()
+ self.is_decoder = config.is_decoder
+ self.is_static_kv = is_static_kv
+ self.has_relative_attention_bias = not self.is_static_kv
+ self.d_model = config.d_model
+ self.key_value_proj_dim = config.head_size
+ self.n_heads = config.num_heads
+ self.inner_dim = self.n_heads * self.key_value_proj_dim
+
+ # Mesh TensorFlow initialization to avoid scaling before softmax
+ self.q = nn.Linear(self.d_model, self.inner_dim, bias=False)
+ self.k = nn.Linear(self.d_model, self.inner_dim, bias=False)
+ self.v = nn.Linear(self.d_model, self.inner_dim, bias=False)
+ # self.o = nn.Linear(self.inner_dim, self.d_model, bias=False)
+
+ self.pruned_heads = set()
+
+ # ORT parameters
+ self.batch_size = config.batch_size
+ self.seq_len = config.seq_len
+ self.kv_sequence_length = config.kv_sequence_length
+ self.head_size = config.head_size
+ self.num_heads = config.num_heads
+ self.hidden_size = self.d_model
+ self.use_past = config.use_past
+
+ # Create onnx graph
+ self.onnx_graph = create_t5_mha_graph(
+ self.batch_size,
+ self.seq_len,
+ self.kv_sequence_length,
+ self.head_size,
+ self.num_heads,
+ self.use_past,
+ is_static_kv,
+ )
+
+ def create_inputs(self):
+ hidden_states = torch.normal(mean=0.5, std=0.1, size=(self.batch_size, self.seq_len, self.hidden_size)).to(
+ torch.float32
+ )
+ key_value_states = torch.normal(
+ mean=0.5, std=0.1, size=(self.batch_size, self.kv_sequence_length, self.hidden_size)
+ ).to(torch.float32)
+ past_key = torch.normal(
+ mean=0.5, std=0.1, size=(self.batch_size, self.num_heads, self.kv_sequence_length, self.head_size)
+ ).to(torch.float32)
+ past_value = torch.normal(
+ mean=0.5, std=0.1, size=(self.batch_size, self.num_heads, self.kv_sequence_length, self.head_size)
+ ).to(torch.float32)
+ past_key_value = (past_key, past_value)
+ attention_mask = torch.ones((self.batch_size, self.kv_sequence_length)).to(torch.float32)
+ position_bias_length = self.seq_len if not self.use_past else self.kv_sequence_length + self.seq_len
+ position_bias = torch.normal(
+ mean=0.5, std=0.1, size=(1, self.num_heads, position_bias_length, position_bias_length)
+ ).to(torch.float32)
+ return hidden_states, key_value_states, past_key_value, attention_mask, position_bias
+
+ def torch_forward(
+ self,
+ hidden_states,
+ key_value_states=None,
+ past_key_value=None,
+ mask=None,
+ position_bias=None,
+ use_cache=False,
+ query_length=None,
+ ):
+ """
+ Self-attention (if key_value_states is None) or attention over source sentence (provided by key_value_states).
+ """
+ # Input is (batch_size, seq_length, dim)
+ # Mask is (batch_size, key_length) (non-causal) or (batch_size, key_length, key_length)
+ # past_key_value[0] is (batch_size, n_heads, q_len - 1, dim_per_head)
+ batch_size, seq_length = hidden_states.shape[:2]
+
+ real_seq_length = seq_length
+
+ if past_key_value is not None:
+ assert (
+ len(past_key_value) == 2
+ ), f"past_key_value should have 2 past states: keys and values. Got { len(past_key_value)} past states"
+ real_seq_length += past_key_value[0].shape[2] if query_length is None else query_length
+
+ key_length = real_seq_length if key_value_states is None else key_value_states.shape[1]
+
+ def shape(states):
+ """projection"""
+ return states.view(batch_size, -1, self.n_heads, self.key_value_proj_dim).transpose(1, 2)
+
+ def unshape(states):
+ """reshape"""
+ return states.transpose(1, 2).contiguous().view(batch_size, -1, self.inner_dim)
+
+ def project(hidden_states, proj_layer, key_value_states, past_key_value):
+ """projects hidden states correctly to key/query states"""
+ if key_value_states is None:
+ # self-attn
+ # (batch_size, n_heads, seq_length, dim_per_head)
+ hidden_states = shape(proj_layer(hidden_states))
+ elif past_key_value is None:
+ # cross-attn
+ # (batch_size, n_heads, seq_length, dim_per_head)
+ hidden_states = shape(proj_layer(key_value_states))
+
+ if past_key_value is not None:
+ if key_value_states is None:
+ # self-attn
+ # (batch_size, n_heads, key_length, dim_per_head)
+ hidden_states = torch.cat([past_key_value, hidden_states], dim=2)
+ elif past_key_value.shape[2] != key_value_states.shape[1]:
+ # checking that the `sequence_length` of the `past_key_value` is the same as
+ # the provided `key_value_states` to support prefix tuning
+ # cross-attn
+ # (batch_size, n_heads, seq_length, dim_per_head)
+ hidden_states = shape(proj_layer(key_value_states))
+ else:
+ # cross-attn
+ hidden_states = past_key_value
+ return hidden_states
+
+ # get query states
+ query_states = shape(self.q(hidden_states)) # (batch_size, n_heads, seq_length, dim_per_head)
+
+ # get key/value states
+ key_states = project(
+ hidden_states, self.k, key_value_states, past_key_value[0] if past_key_value is not None else None
+ )
+ value_states = project(
+ hidden_states, self.v, key_value_states, past_key_value[1] if past_key_value is not None else None
+ )
+
+ # compute scores
+ scores = torch.matmul(
+ query_states, key_states.transpose(3, 2)
+ ) # equivalent of torch.einsum("bnqd,bnkd->bnqk", query_states, key_states), compatible with onnx op>9
+
+ if position_bias is None:
+ if not self.has_relative_attention_bias:
+ position_bias = torch.zeros(
+ (1, self.n_heads, real_seq_length, key_length), device=scores.device, dtype=scores.dtype
+ )
+ else:
+ assert position_bias is not None
+
+ # if key and values are already calculated
+ # we want only the last query position bias
+ if past_key_value is not None and position_bias is not None:
+ position_bias = position_bias[:, :, -hidden_states.size(1) :, :]
+
+ if mask is not None:
+ # Adjust onnx mask shape
+ mask = (1 - mask.unsqueeze(1).unsqueeze(2)) * torch.finfo(torch.float32).min
+ position_bias = position_bias + mask # (batch_size, n_heads, seq_length, key_length)
+
+ if self.pruned_heads:
+ mask = torch.ones(position_bias.shape[1])
+ mask[list(self.pruned_heads)] = 0
+ position_bias_masked = position_bias[:, mask.bool()]
+ else:
+ position_bias_masked = position_bias
+
+ scores += position_bias_masked
+ attn_weights = nn.functional.softmax(scores.float(), dim=-1).type_as(
+ scores
+ ) # (batch_size, n_heads, seq_length, key_length)
+ attn_output = unshape(torch.matmul(attn_weights, value_states)) # (batch_size, seq_length, dim)
+ # attn_output = self.o(attn_output) # ORT places this matmul outside of MHA op
+
+ present_key_value_state = (key_states, value_states) if (self.is_decoder and use_cache) else None
+ outputs = (attn_output,) + (present_key_value_state,)
+
+ return outputs
+
+ def ort_forward(
+ self,
+ hidden_states,
+ key_value_states=None,
+ past_key_value=None,
+ mask=None,
+ position_bias=None,
+ use_cache=False,
+ query_length=None,
+ ):
+ import onnxruntime
+
+ sess_options = onnxruntime.SessionOptions()
+ cuda_providers = ["CUDAExecutionProvider"]
+ if cuda_providers[0] not in onnxruntime.get_available_providers():
+ return None
+ ort_session = onnxruntime.InferenceSession(self.onnx_graph, sess_options, providers=cuda_providers)
+
+ batch_size, seq_length = hidden_states.shape[:2]
+
+ real_seq_length = seq_length
+
+ if past_key_value is not None:
+ assert (
+ len(past_key_value) == 2
+ ), f"past_key_value should have 2 past states: keys and values. Got { len(past_key_value)} past states"
+ real_seq_length += past_key_value[0].shape[2] if query_length is None else query_length
+
+ def project(hidden_states, proj_layer, key_value_states, past_key_value):
+ if key_value_states is None:
+ # self-attn
+ hidden_states = proj_layer(hidden_states)
+ elif past_key_value is None:
+ # cross-attn
+ # (batch_size, n_heads, seq_length, dim_per_head)
+ hidden_states = proj_layer(key_value_states)
+
+ return hidden_states
+
+ # get query states
+ query_states = self.q(hidden_states) # (batch_size, n_heads, seq_length, dim_per_head)
+
+ # get key/value states
+ key_states = project(
+ hidden_states, self.k, key_value_states, past_key_value[0] if past_key_value is not None else None
+ )
+ value_states = project(
+ hidden_states, self.v, key_value_states, past_key_value[1] if past_key_value is not None else None
+ )
+
+ if past_key_value is not None and position_bias is not None:
+ position_bias = position_bias[:, :, -hidden_states.size(1) :, :]
+
+ torch_key_padding_mask = mask.to(torch.int32) if mask is not None else None
+ torch_position_bias = position_bias if position_bias is not None else None
+
+ ort_inputs = None
+ if past_key_value is None:
+ ort_inputs = {
+ "query": np.ascontiguousarray(query_states.detach().numpy()),
+ "key": np.ascontiguousarray(key_states.detach().numpy()),
+ "value": np.ascontiguousarray(value_states.detach().numpy()),
+ }
+ if torch_key_padding_mask is not None:
+ ort_inputs["key_padding_mask"] = np.ascontiguousarray(torch_key_padding_mask.detach().numpy())
+ if torch_position_bias is not None:
+ ort_inputs["relative_position_bias"] = np.ascontiguousarray(torch_position_bias.detach().numpy())
+ else:
+ torch_past_key = past_key_value[0]
+ torch_past_value = past_key_value[1]
+ ort_inputs = {
+ "query": np.ascontiguousarray(query_states.detach().numpy()),
+ }
+ if self.is_static_kv:
+ ort_inputs["key"] = np.ascontiguousarray(torch_past_key.detach().numpy())
+ ort_inputs["value"] = np.ascontiguousarray(torch_past_value.detach().numpy())
+ else:
+ ort_inputs["past_key"] = np.ascontiguousarray(torch_past_key.detach().numpy())
+ ort_inputs["past_value"] = np.ascontiguousarray(torch_past_value.detach().numpy())
+ ort_inputs["key"] = np.ascontiguousarray(key_states.detach().numpy())
+ ort_inputs["value"] = np.ascontiguousarray(value_states.detach().numpy())
+ if torch_key_padding_mask is not None:
+ ort_inputs["key_padding_mask"] = np.ascontiguousarray(torch_key_padding_mask.detach().numpy())
+ if torch_position_bias is not None:
+ ort_inputs["relative_position_bias"] = np.ascontiguousarray(torch_position_bias.detach().numpy())
+
+ ort_output = ort_session.run(None, ort_inputs)
+
+ output = None
+ if past_key_value is not None and self.is_static_kv:
+ output = torch.tensor(ort_output)
+ else:
+ output = (torch.tensor(ort_output[0]),) + ((torch.tensor(ort_output[1]), torch.tensor(ort_output[2])),)
+
+ return output
+
+
+def compare_t5_cross_attention_decoder(batch_size, seq_len, num_heads, head_size, kv_sequence_length):
+ config = T5Config(
+ is_decoder=True,
+ batch_size=batch_size,
+ seq_len=seq_len,
+ kv_sequence_length=kv_sequence_length,
+ num_heads=num_heads,
+ head_size=head_size,
+ use_past=True,
+ )
+ T5CrossAttention = T5Attention(config, is_static_kv=True)
+
+ hidden_states, key_value_states, past_key_value, attention_mask, _ = T5CrossAttention.create_inputs()
+ torch_output = T5CrossAttention.torch_forward(
+ hidden_states, key_value_states, past_key_value, attention_mask, position_bias=None, use_cache=False
+ )
+ ort_output = T5CrossAttention.ort_forward(
+ hidden_states, key_value_states, past_key_value, attention_mask, position_bias=None, use_cache=False
+ )
+
+ if ort_output is not None:
+ assert torch.allclose(torch_output[0], ort_output[0], atol=1e-4)
+
+
+def compare_t5_cross_attention_decoder_init(batch_size, seq_len, num_heads, head_size, kv_sequence_length):
+ config = T5Config(
+ is_decoder=True,
+ batch_size=batch_size,
+ seq_len=seq_len,
+ kv_sequence_length=kv_sequence_length,
+ num_heads=num_heads,
+ head_size=head_size,
+ use_past=False,
+ )
+ T5CrossAttention = T5Attention(config, is_static_kv=True)
+
+ hidden_states, key_value_states, _, attention_mask, _ = T5CrossAttention.create_inputs()
+ torch_output = T5CrossAttention.torch_forward(
+ hidden_states, key_value_states, None, attention_mask, position_bias=None, use_cache=True
+ )
+ ort_output = T5CrossAttention.ort_forward(
+ hidden_states, key_value_states, None, attention_mask, position_bias=None, use_cache=True
+ )
+
+ if ort_output is not None:
+ assert torch.allclose(torch_output[0], ort_output[0], atol=1e-4)
+ assert torch.allclose(torch_output[1][0], ort_output[1][0], atol=1e-4)
+ assert torch.allclose(torch_output[1][1], ort_output[1][1], atol=1e-4)
+
+
+def compare_t5_self_attention_decoder_init(batch_size, seq_len, num_heads, head_size, kv_sequence_length):
+ config = T5Config(
+ is_decoder=True,
+ batch_size=batch_size,
+ seq_len=seq_len,
+ kv_sequence_length=kv_sequence_length,
+ num_heads=num_heads,
+ head_size=head_size,
+ use_past=False,
+ )
+ T5CrossAttention = T5Attention(config, is_static_kv=False)
+
+ hidden_states, _, _, _, position_bias = T5CrossAttention.create_inputs()
+ torch_output = T5CrossAttention.torch_forward(
+ hidden_states, None, None, mask=None, position_bias=position_bias, use_cache=True
+ )
+ ort_output = T5CrossAttention.ort_forward(
+ hidden_states, None, None, mask=None, position_bias=position_bias, use_cache=True
+ )
+
+ if ort_output is not None:
+ assert torch.allclose(torch_output[0], ort_output[0], atol=1e-4)
+ assert torch.allclose(torch_output[1][0], ort_output[1][0], atol=1e-4)
+ assert torch.allclose(torch_output[1][1], ort_output[1][1], atol=1e-4)
+
+
+def compare_t5_self_attention_decoder(batch_size, seq_len, num_heads, head_size, kv_sequence_length):
+ config = T5Config(
+ is_decoder=True,
+ batch_size=batch_size,
+ seq_len=seq_len,
+ kv_sequence_length=kv_sequence_length,
+ num_heads=num_heads,
+ head_size=head_size,
+ use_past=True,
+ )
+ T5CrossAttention = T5Attention(config, is_static_kv=False)
+
+ hidden_states, _, past_key_value, _, position_bias = T5CrossAttention.create_inputs()
+ torch_output = T5CrossAttention.torch_forward(
+ hidden_states, None, past_key_value, mask=None, position_bias=position_bias, use_cache=True
+ )
+ ort_output = T5CrossAttention.ort_forward(
+ hidden_states, None, past_key_value, mask=None, position_bias=position_bias, use_cache=True
+ )
+
+ if ort_output is not None:
+ assert torch.allclose(torch_output[0], ort_output[0], atol=1e-4)
+ assert torch.allclose(torch_output[1][0], ort_output[1][0], atol=1e-4)
+ assert torch.allclose(torch_output[1][1], ort_output[1][1], atol=1e-4)
+
+
+class TestT5MHAParity(unittest.TestCase):
+ def setUp(self):
+ self.batch_size = 5
+ self.seq_len = 2
+ self.num_heads = 2
+ self.head_size = 4
+ self.kv_sequence_length = 3
+
+ def test_t5_cross_attention_decoder_init(self):
+ compare_t5_cross_attention_decoder_init(
+ self.batch_size, self.seq_len, self.num_heads, self.head_size, self.kv_sequence_length
+ )
+
+ def test_t5_self_attention_decoder_init(self):
+ compare_t5_self_attention_decoder_init(
+ self.batch_size, self.seq_len, self.num_heads, self.head_size, self.kv_sequence_length
+ )
+
+ def test_t5_cross_attention_decoder(self):
+ compare_t5_cross_attention_decoder(
+ self.batch_size, self.seq_len, self.num_heads, self.head_size, self.kv_sequence_length
+ )
+
+ def test_t5_self_attention_decoder(self):
+ compare_t5_self_attention_decoder(
+ self.batch_size, self.seq_len, self.num_heads, self.head_size, self.kv_sequence_length
+ )
+
+
+if __name__ == "__main__":
+ unittest.main()