diff --git a/docs/ContribOperators.md b/docs/ContribOperators.md
index 7e277c9aa8..03b1b24c49 100644
--- a/docs/ContribOperators.md
+++ b/docs/ContribOperators.md
@@ -61,6 +61,7 @@ Do not modify directly.*
* com.microsoft.Unique
* com.microsoft.WordConvEmbedding
* experimental com.microsoft.IsAllFinite
+ * experimental com.microsoft.QEmbedLayerNormalization
## com.microsoft
### **com.microsoft.Attention**
@@ -2905,3 +2906,83 @@ No versioning maintained for experimental ops.
+### experimental **com.microsoft.QEmbedLayerNormalization**
+
+ QEmbedLayerNormalization is the quantized fusion of embedding layer in BERT model, with optional mask processing.
+ The embedding layer takes input_ids (word IDs) and segment_ids (sentence IDs) to look up word_embedding, position_embedding,
+ and segment_emedding; the embeddings are added then applied layer normalization using gamma and beta tensors. The input_ids
+ and segment_ids remain int32. All embeddings, gamma, and beta tensors are converted to int8/uint8. The last input mask is optional.
+ If mask is provided, mask index (that is position of first 0 in mask, or number of words will be calculated.
+
+#### Version
+
+No versioning maintained for experimental ops.
+#### Attributes
+
+
+- epsilon : float
+- The epsilon value to use to avoid division by zero.
+
+
+#### Inputs
+
+
+- input_ids : T1
+- 2D words IDs with shape (batch_size, sequence_length)
+- segment_ids (optional) : T1
+- 2D segment IDs with shape (batch_size, sequence_length)
+- word_embedding_quant : T2
+- 2D with shape (,hidden_size)
+- position_embedding_quant : T2
+- 2D with shape (, hidden_size)
+- segment_embedding (optional) : T2
+- 2D with shape (, hidden_size)
+- gamma_quant : T2
+- 1D gamma tensor for layer normalization with shape (hidden_size)
+- beta_quant : T2
+- 1D beta tensor for layer normalization with shape (hidden_size)
+- mask (optional) : T1
+- Mask
+- word_embedding_scale : T
+- Scale for word embeddings
+- position_embedding_scale : T
+- Scale for position embeddings
+- segment_embedding_scale (optional) : T
+- Scale for segment embeddings
+- gamma_scale : T
+- Scale for 1D gamma tensor
+- beta_scale : T
+- Scale for 1D beta tensor
+- word_embedding_zero_point : T2
+- Zero point for word embeddings
+- position_embedding_zero_point : T2
+- Zero point for position embeddings
+- segment_embedding_zero_point (optional) : T2
+- Zero Point for segment embeddings
+- gamma_zero_point : T2
+- Zero Point for 1D gamma tensor
+- beta_zero_point : T2
+- Zero Point for 1D beta tensor
+
+
+#### Outputs
+
+
+- layernorm_out : T
+- LayerNorm Output
+- mask_index_out : T1
+- Mask Index Output
+
+
+#### Type Constraints
+
+
+- T1 : tensor(int32)
+- Constrain mask index to integer types
+- T2 : tensor(int8), tensor(uint8)
+- Constrain input and output types to int8 tensors.
+- T : tensor(float)
+- Constrain input and output types to float32 tensors.
+
+
+
diff --git a/docs/OperatorKernels.md b/docs/OperatorKernels.md
index 97935346e0..8968107c86 100644
--- a/docs/OperatorKernels.md
+++ b/docs/OperatorKernels.md
@@ -392,6 +392,7 @@ Do not modify directly.*
|NhwcMaxPool|*in* x:**T**
*out* y:**T**|1+|**T** = tensor(uint8)|
|Pad|*in* data:**T**
*in* pads:**tensor(int64)**
*in* value:**T**
*out* output:**T**|1+|**T** = tensor(float)|
|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(uint8)
**T2** = tensor(int8), tensor(uint8)
**T3** = tensor(float)
**T4** = tensor(int32)|
+|QEmbedLayerNormalization|*in* input_ids:**T1**
*in* segment_ids:**T1**
*in* word_embedding_quant:**T2**
*in* position_embedding_quant:**T2**
*in* segment_embedding:**T2**
*in* gamma_quant:**T2**
*in* beta_quant:**T2**
*in* mask:**T1**
*in* word_embedding_scale:**T**
*in* position_embedding_scale:**T**
*in* segment_embedding_scale:**T**
*in* gamma_scale:**T**
*in* beta_scale:**T**
*in* word_embedding_zero_point:**T2**
*in* position_embedding_zero_point:**T2**
*in* segment_embedding_zero_point:**T2**
*in* gamma_zero_point:**T2**
*in* beta_zero_point:**T2**
*out* layernorm_out:**T**
*out* mask_index_out:**T1**|1+|**T** = tensor(float)|
|QLinearAdd|*in* A:**T**
*in* A_scale:**tensor(float)**
*in* A_zero_point:**T**
*in* B:**T**
*in* B_scale:**tensor(float)**
*in* B_zero_point:**T**
*in* C_scale:**tensor(float)**
*in* C_zero_point:**T**
*out* C:**T**|1+|**T** = tensor(int8), tensor(uint8)|
|QLinearConv|*in* x:**T1**
*in* x_scale:**tensor(float)**
*in* x_zero_point:**T1**
*in* w:**T2**
*in* w_scale:**tensor(float)**
*in* w_zero_point:**T2**
*in* y_scale:**tensor(float)**
*in* y_zero_point:**T3**
*in* B:**T4**
*out* y:**T3**|1+|**T1** = tensor(uint8)
**T2** = tensor(int8), tensor(uint8)
**T3** = tensor(uint8)
**T4** = tensor(int32)|
|QLinearLeakyRelu|*in* X:**T**
*in* X_scale:**tensor(float)**
*in* X_zero_point:**T**
*in* Y_scale:**tensor(float)**
*in* Y_zero_point:**T**
*out* Y:**T**|1+|**T** = tensor(int8), tensor(uint8)|
diff --git a/onnxruntime/contrib_ops/cpu/bert/embed_layer_norm.cc b/onnxruntime/contrib_ops/cpu/bert/embed_layer_norm.cc
index 952ad156c1..91a341baab 100644
--- a/onnxruntime/contrib_ops/cpu/bert/embed_layer_norm.cc
+++ b/onnxruntime/contrib_ops/cpu/bert/embed_layer_norm.cc
@@ -24,12 +24,20 @@ namespace contrib {
REGISTER_KERNEL_TYPED(float)
-template
-EmbedLayerNorm::EmbedLayerNorm(const OpKernelInfo& op_kernel_info) : OpKernel(op_kernel_info) {
+EmbedLayerNormBase::EmbedLayerNormBase(const OpKernelInfo& op_kernel_info) : OpKernel(op_kernel_info) {
ORT_ENFORCE(op_kernel_info.GetAttr("epsilon", &epsilon_).IsOK());
ORT_ENFORCE(epsilon_ >= 0);
}
+float EmbedLayerNormBase::epsilon() const {
+ return epsilon_;
+}
+
+template
+EmbedLayerNorm::EmbedLayerNorm(const OpKernelInfo& op_kernel_info)
+ : EmbedLayerNormBase(op_kernel_info) {
+}
+
template
Status EmbedLayerNorm::Compute(OpKernelContext* context) const {
ORT_RETURN_IF_ERROR(embed_layer_norm::CheckInputs(context));
@@ -112,7 +120,7 @@ Status EmbedLayerNorm::Compute(OpKernelContext* context) const {
y[i] = a;
sum += a * a;
}
- T e = sqrt(sum / hidden_size + static_cast(epsilon_));
+ T e = sqrt(sum / hidden_size + static_cast(epsilon()));
for (int i = 0; i < hidden_size; i++) {
y[i] = y[i] / e * gamma_data[i] + beta_data[i];
}
@@ -125,11 +133,14 @@ Status EmbedLayerNorm::Compute(OpKernelContext* context) const {
// Calculate mask
if (nullptr != mask) {
+ // TODO: Consider summing the values in the mask and measure performance.
const int32_t* mask_data = mask->template Data();
+ int32_t* mask_index_data = mask_index->template MutableData();
for (int b = 0; b < batch_size; b++) {
- mask_index->template MutableData()[b] = static_cast(std::count_if(mask_data + (b * sequence_length),
- mask_data + (b * sequence_length) + sequence_length,
- [](int v) { return v == 1; }));
+ mask_index_data[b] =
+ static_cast(std::count_if(mask_data + (static_cast(b) * sequence_length),
+ mask_data + (static_cast(b) * sequence_length) + sequence_length,
+ [](int v) { return v == 1; }));
}
} else {
memset(mask_index->template MutableData(), 0, batch_size * sizeof(int32_t));
diff --git a/onnxruntime/contrib_ops/cpu/bert/embed_layer_norm.h b/onnxruntime/contrib_ops/cpu/bert/embed_layer_norm.h
index 5dfe6c833e..d124ed344e 100644
--- a/onnxruntime/contrib_ops/cpu/bert/embed_layer_norm.h
+++ b/onnxruntime/contrib_ops/cpu/bert/embed_layer_norm.h
@@ -8,13 +8,23 @@
namespace onnxruntime {
namespace contrib {
-template
-class EmbedLayerNorm : public OpKernel {
+
+class EmbedLayerNormBase : public OpKernel {
public:
- explicit EmbedLayerNorm(const OpKernelInfo& op_kernel_info);
- Status Compute(OpKernelContext* context) const override;
+ explicit EmbedLayerNormBase(const OpKernelInfo& op_kernel_info);
+
+ protected:
+ float epsilon() const;
+
private:
float epsilon_;
};
+
+template
+class EmbedLayerNorm : public EmbedLayerNormBase {
+ public:
+ explicit EmbedLayerNorm(const OpKernelInfo& op_kernel_info);
+ Status Compute(OpKernelContext* context) const override;
+};
} // namespace contrib
} // namespace onnxruntime
diff --git a/onnxruntime/contrib_ops/cpu/bert/embed_layer_norm_helper.cc b/onnxruntime/contrib_ops/cpu/bert/embed_layer_norm_helper.cc
index e991502745..de57333411 100644
--- a/onnxruntime/contrib_ops/cpu/bert/embed_layer_norm_helper.cc
+++ b/onnxruntime/contrib_ops/cpu/bert/embed_layer_norm_helper.cc
@@ -3,6 +3,7 @@
#include "embed_layer_norm_helper.h"
#include "core/framework/tensorprotoutils.h"
+#include "core/providers/common.h"
#include "onnx/defs/tensor_proto_util.h"
#include "longformer_attention_base.h"
@@ -54,6 +55,7 @@ Status CheckInputs(const OpKernelContext* context) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"word_embedding is expected to have 2 dimensions, got ", word_embedding_dims.size());
}
+ int64_t hidden_size = word_embedding->Shape()[1];
const auto& position_embedding_dims = position_embedding->Shape().GetDims();
if (position_embedding_dims.size() != 2) {
@@ -67,37 +69,37 @@ Status CheckInputs(const OpKernelContext* context) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"segment_embedding is expected to have 2 dimensions, got ", segment_embedding_dims.size());
}
- if (word_embedding_dims[1] != segment_embedding_dims[1]) {
+ if (segment_embedding_dims[1] != hidden_size) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"word_embedding and segment_embedding shall have same dimension 1");
}
}
- if (word_embedding_dims[1] != position_embedding_dims[1]) {
+ if (position_embedding_dims[1] != hidden_size) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"word_embedding and position_embedding shall have same dimension 1");
}
- const auto& beta_dims = beta->Shape().GetDims();
- if (beta_dims.size() != 1) {
- return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
- "beta is expected to have 1 dimensions, got ", beta_dims.size());
- }
-
- if (beta_dims[0] != word_embedding_dims[1]) {
- return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
- "beta is expected to have size of ", word_embedding_dims[1], ", got ", beta_dims[0]);
- }
-
const auto& gamma_dims = gamma->Shape().GetDims();
if (gamma_dims.size() != 1) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"gamma is expected to have 1 dimensions, got ", gamma_dims.size());
}
- if (gamma_dims[0] != word_embedding_dims[1]) {
+ if (gamma_dims[0] != hidden_size) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
- "gamma is expected to have size of ", word_embedding_dims[1], ", got ", gamma_dims[0]);
+ "gamma is expected to have size of ", hidden_size, ", got ", gamma_dims[0]);
+ }
+
+ const auto& beta_dims = beta->Shape().GetDims();
+ if (beta_dims.size() != 1) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "beta is expected to have 1 dimensions, got ", beta_dims.size());
+ }
+
+ if (beta_dims[0] != hidden_size) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "beta is expected to have size of ", hidden_size, ", got ", beta_dims[0]);
}
return Status::OK();
diff --git a/onnxruntime/contrib_ops/cpu/bert/qembed_layer_norm.cc b/onnxruntime/contrib_ops/cpu/bert/qembed_layer_norm.cc
new file mode 100644
index 0000000000..170c4353e0
--- /dev/null
+++ b/onnxruntime/contrib_ops/cpu/bert/qembed_layer_norm.cc
@@ -0,0 +1,306 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+#include "qembed_layer_norm.h"
+
+#include
+
+#include "embed_layer_norm_helper.h"
+#include "core/framework/op_kernel.h"
+#include "core/providers/common.h"
+
+namespace onnxruntime {
+namespace contrib {
+
+namespace {
+
+// TODO(kreeger): Drop this when ComputeInternal() is using a lookup table.
+template
+inline float Dequantize(T value, float scale, T zero_point) {
+ return static_cast(static_cast(value) - zero_point) * scale;
+}
+
+template
+Status ComputeInternal(OpKernelContext* context, float epsilon) {
+ const Tensor* input_ids = context->Input(0);
+ const Tensor* segment_ids = context->Input(1); // optional. nullptr if it's distill-bert
+ const Tensor* word_embedding = context->Input(2);
+ const Tensor* position_embedding = context->Input(3);
+ const Tensor* segment_embedding = context->Input(4); // optional. nullptr if it's distill-bert
+ const Tensor* gamma = context->Input(5);
+ const Tensor* beta = context->Input(6);
+ const Tensor* mask = context->Input(7); // optional. nullptr if not provided
+ const Tensor* word_embedding_scale = context->Input(8);
+ const Tensor* position_embedding_scale = context->Input(9);
+ const Tensor* segment_embedding_scale = context->Input(10);
+ const Tensor* gamma_scale = context->Input(11);
+ const Tensor* beta_scale = context->Input(12);
+ const Tensor* word_embedding_zero_point = context->Input(13);
+ const Tensor* position_embedding_zero_point = context->Input(14);
+ const Tensor* segment_embedding_zero_point = context->Input(15);
+ const Tensor* gamma_zero_point = context->Input(16);
+ const Tensor* beta_zero_point = context->Input(17);
+
+ const auto& input_dims = input_ids->Shape().GetDims();
+ int batch_size = static_cast(input_dims[0]);
+ int sequence_length = static_cast(input_dims[1]);
+ int64_t hidden_size = word_embedding->Shape()[1];
+
+ // Request outputs:
+ TensorShape output_shape({batch_size, sequence_length, hidden_size});
+ Tensor* output = context->Output(0, output_shape);
+
+ TensorShape mask_index_shape({batch_size});
+ Tensor* mask_index = context->Output(1, mask_index_shape);
+ bool has_segment_embedding = segment_ids != nullptr;
+
+ const int32_t* input_ids_data = input_ids->template Data();
+ const int32_t* segment_ids_data =
+ has_segment_embedding ? segment_ids->template Data() : nullptr;
+
+ int word_embedding_length = static_cast(word_embedding->Shape()[0]);
+ int position_embedding_length = static_cast(position_embedding->Shape()[0]);
+ int segment_embedding_length =
+ has_segment_embedding ? static_cast(segment_embedding->Shape()[0]) : 0;
+
+ // Grab quantization values:
+ float word_embedding_scale_data = *(word_embedding_scale->template Data());
+ T2 word_embedding_zero_point_data = *(word_embedding_zero_point->template Data());
+
+ float position_embedding_scale_data = *(position_embedding_scale->template Data());
+ T2 position_embedding_zero_point_data = *(position_embedding_zero_point->template Data());
+
+ float segment_embedding_scale_data =
+ has_segment_embedding ? *(segment_embedding_scale->template Data()) : 0.0f;
+ T2 segment_embedding_zero_point_data =
+ has_segment_embedding ? *(segment_embedding_zero_point->template Data()) : 0;
+
+ float gamma_scale_data = *(gamma_scale->template Data());
+ T2 gamma_zero_point_data = *(gamma_zero_point->template Data());
+
+ float beta_scale_data = *(beta_scale->template Data());
+ T2 beta_zero_point_data = *(beta_zero_point->template Data());
+
+ // Grab pointers to buffers each Tensor represents:
+ const T2* word_embedding_data = word_embedding->template Data();
+ const T2* position_embedding_data = position_embedding->template Data();
+ const T2* segment_embedding_data =
+ has_segment_embedding ? segment_embedding->template Data() : nullptr;
+ const T2* gamma_data = gamma->template Data();
+ const T2* beta_data = beta->template Data();
+
+ T* output_data = output->template MutableData();
+
+ // Perform the Op:
+ {
+ std::atomic_bool failed{false};
+
+ // TODO: Profile and tune this batch parallel execution based on input size.
+ // More info: https://github.com/microsoft/onnxruntime/pull/8124/files#r656629895
+ int n = batch_size * sequence_length;
+ concurrency::ThreadPool::TryBatchParallelFor(
+ context->GetOperatorThreadPool(), n, [=, &failed](ptrdiff_t index) {
+ int word_col_index = input_ids_data[index];
+ if (word_col_index < 0 || word_col_index >= word_embedding_length) {
+ failed.store(true, std::memory_order_release);
+ return;
+ }
+ int position_col_index = index % sequence_length;
+ if (position_col_index >= position_embedding_length) {
+ failed.store(true, std::memory_order_release);
+ return;
+ }
+ int segment_col_index = 0;
+ if (nullptr != segment_ids_data) {
+ segment_col_index = segment_ids_data[index];
+ if (segment_col_index < 0 || segment_col_index >= segment_embedding_length) {
+ failed.store(true, std::memory_order_release);
+ return;
+ }
+ }
+
+ // Grab inputs for the embeddings for the current batch index:
+ const T2* input_word_embedding = word_embedding_data + word_col_index * hidden_size;
+ const T2* input_position_embedding =
+ position_embedding_data + position_col_index * hidden_size;
+ const T2* input_segment_embedding = nullptr;
+ if (segment_embedding_data != nullptr) {
+ input_segment_embedding = segment_embedding_data + segment_col_index * hidden_size;
+ }
+
+ T* output = output_data + (index * hidden_size);
+
+ T sum = static_cast(0);
+ for (int i = 0; i < hidden_size; ++i) {
+ // TODO(kreeger): Use a table query to improve performance:
+ T subtotal = Dequantize(input_word_embedding[i],
+ word_embedding_scale_data,
+ word_embedding_zero_point_data) +
+ Dequantize(input_position_embedding[i],
+ position_embedding_scale_data,
+ position_embedding_zero_point_data);
+ if (segment_embedding_data != nullptr) {
+ subtotal += Dequantize(input_segment_embedding[i],
+ segment_embedding_scale_data,
+ segment_embedding_zero_point_data);
+ }
+ output[i] = subtotal;
+ sum += subtotal;
+ }
+
+ T mean = sum / hidden_size;
+ sum = 0;
+
+ for (int i = 0; i < hidden_size; i++) {
+ T a = output[i] - mean;
+ output[i] = a;
+ sum += a * a;
+ }
+
+ T e = sqrt(sum / hidden_size + epsilon);
+ for (int i = 0; i < hidden_size; i++) {
+ // TODO(kreeger): Consider keeping these as int8 or use PrePack()!
+ T cur_gamma = Dequantize(gamma_data[i], gamma_scale_data, gamma_zero_point_data);
+ T cur_beta = Dequantize(beta_data[i], beta_scale_data, beta_zero_point_data);
+ output[i] = output[i] / e * cur_gamma + cur_beta;
+ }
+ },
+ 0);
+
+ if (failed.load(std::memory_order_acquire)) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "input index out of range");
+ }
+ }
+
+ // Calculate mask
+ if (nullptr != mask) {
+ // TODO: Consider summing the values in the mask and measure performance.
+ const int32_t* mask_data = mask->template Data();
+ int32_t* mask_index_data = mask_index->template MutableData();
+ for (int b = 0; b < batch_size; b++) {
+ mask_index_data[b] =
+ static_cast(std::count_if(mask_data + (static_cast(b) * sequence_length),
+ mask_data + (static_cast(b) * sequence_length) + sequence_length,
+ [](int v) { return v == 1; }));
+ }
+ } else {
+ memset(mask_index->template MutableData(), 0, batch_size * sizeof(int32_t));
+ }
+ return Status::OK();
+}
+
+Status CheckQuantizedInputs(OpKernelContext* context, bool* is_signed_inputs) {
+ const Tensor* word_embedding_scale_tensor = context->Input(8);
+ const Tensor* position_embedding_scale_tensor = context->Input(9);
+ const Tensor* segment_embedding_scale_tensor = context->Input(10);
+ const Tensor* gamma_scale_tensor = context->Input(11);
+ const Tensor* beta_scale_tensor = context->Input(12);
+ const Tensor* word_embedding_zero_point_tensor = context->Input(13);
+ const Tensor* position_embedding_zero_point_tensor = context->Input(14);
+ const Tensor* segment_embedding_zero_point_tensor = context->Input(15);
+ const Tensor* gamma_zero_point_tensor = context->Input(16);
+ const Tensor* beta_zero_point_tensor = context->Input(17);
+
+ bool word_embedding_is_signed_inputs = word_embedding_scale_tensor->IsDataType();
+ bool has_segment_embedding = context->Input(1) != nullptr;
+
+ if (!IsScalarOr1ElementVector(word_embedding_scale_tensor)) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Word embedding scale must be a scalar or 1D tensor of size 1");
+ }
+
+ if (!IsScalarOr1ElementVector(position_embedding_scale_tensor) &&
+ position_embedding_scale_tensor->IsDataType() == word_embedding_is_signed_inputs) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Position embedding scale must be a scalar or 1D tensor of size 1");
+ }
+
+ if (has_segment_embedding && !IsScalarOr1ElementVector(segment_embedding_scale_tensor) &&
+ segment_embedding_scale_tensor->IsDataType() == word_embedding_is_signed_inputs) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Segment embedding scale must be a scalar or 1D tensor of size 1");
+ }
+
+ if (!IsScalarOr1ElementVector(gamma_scale_tensor) &&
+ gamma_scale_tensor->IsDataType() == word_embedding_is_signed_inputs) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Gamma scale must be a scalar or 1D tensor of size 1");
+ }
+
+ if (!IsScalarOr1ElementVector(beta_scale_tensor) &&
+ beta_scale_tensor->IsDataType() == word_embedding_is_signed_inputs) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Beta scale must be a scalar or 1D tensor of size 1");
+ }
+
+ if (!IsScalarOr1ElementVector(word_embedding_zero_point_tensor) &&
+ word_embedding_zero_point_tensor->IsDataType() == word_embedding_is_signed_inputs) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Word embedding zero point must be a scalar or 1D tensor of size 1");
+ }
+
+ if (!IsScalarOr1ElementVector(position_embedding_zero_point_tensor) &&
+ position_embedding_zero_point_tensor->IsDataType() == word_embedding_is_signed_inputs) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Position embedding zero point must be a scalar or 1D tensor of size 1");
+ }
+
+ if (has_segment_embedding && !IsScalarOr1ElementVector(segment_embedding_zero_point_tensor) &&
+ segment_embedding_zero_point_tensor->IsDataType() == word_embedding_is_signed_inputs) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Segment embedding zero point must be a scalar or 1D tensor of size 1");
+ }
+
+ if (!IsScalarOr1ElementVector(gamma_zero_point_tensor) &&
+ gamma_zero_point_tensor->IsDataType() == word_embedding_is_signed_inputs) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Gamma zero point must be a scalar or 1D tensor of size 1");
+ }
+
+ if (!IsScalarOr1ElementVector(beta_zero_point_tensor) &&
+ beta_zero_point_tensor->IsDataType() == word_embedding_is_signed_inputs) {
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Beta zero point must be a scalar or 1D tensor of size 1");
+ }
+
+ *is_signed_inputs = word_embedding_is_signed_inputs;
+ return Status::OK();
+}
+
+} // namespace
+
+// This op is internal-only, so register outside of onnx:
+#define REGISTER_KERNEL_TYPED(T) \
+ ONNX_OPERATOR_TYPED_KERNEL_EX( \
+ QEmbedLayerNormalization, \
+ kMSDomain, \
+ 1, \
+ T, \
+ kCpuExecutionProvider, \
+ KernelDefBuilder() \
+ .TypeConstraint("T", DataTypeImpl::GetTensorType()), \
+ QEmbedLayerNorm);
+
+REGISTER_KERNEL_TYPED(float)
+
+template
+QEmbedLayerNorm::QEmbedLayerNorm(const OpKernelInfo& op_kernel_info)
+ : EmbedLayerNormBase(op_kernel_info) {
+}
+
+template
+Status QEmbedLayerNorm::Compute(OpKernelContext* context) const {
+ ORT_RETURN_IF_ERROR(embed_layer_norm::CheckInputs(context));
+
+ bool is_signed_inputs = false;
+ ORT_RETURN_IF_ERROR(CheckQuantizedInputs(context, &is_signed_inputs));
+
+ if (is_signed_inputs) {
+ return ComputeInternal(context, epsilon());
+ } else {
+ return ComputeInternal(context, epsilon());
+ }
+}
+
+} // namespace contrib
+} // namespace onnxruntime
diff --git a/onnxruntime/contrib_ops/cpu/bert/qembed_layer_norm.h b/onnxruntime/contrib_ops/cpu/bert/qembed_layer_norm.h
new file mode 100644
index 0000000000..1788c2e2c6
--- /dev/null
+++ b/onnxruntime/contrib_ops/cpu/bert/qembed_layer_norm.h
@@ -0,0 +1,22 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+#pragma once
+
+#include "embed_layer_norm.h"
+#include "core/common/common.h"
+#include "core/framework/op_kernel.h"
+
+namespace onnxruntime {
+namespace contrib {
+
+// Quantized version of QEmbedLayerNorm.
+template
+class QEmbedLayerNorm final : public EmbedLayerNormBase {
+ public:
+ explicit QEmbedLayerNorm(const OpKernelInfo& op_kernel_info);
+ Status Compute(OpKernelContext* context) const override;
+};
+
+} // namespace contrib
+} // namespace onnxruntime
diff --git a/onnxruntime/contrib_ops/cpu/cpu_contrib_kernels.cc b/onnxruntime/contrib_ops/cpu/cpu_contrib_kernels.cc
index 0f9b2435a3..2d17b094d2 100644
--- a/onnxruntime/contrib_ops/cpu/cpu_contrib_kernels.cc
+++ b/onnxruntime/contrib_ops/cpu/cpu_contrib_kernels.cc
@@ -68,6 +68,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1,
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, float, DynamicQuantizeLSTM);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, QLinearConv);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, NhwcMaxPool);
+class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, float, QEmbedLayerNormalization);
// ******** End: Quantization ******************* //
// This section includes all op kernel declarations for former experimental ops which have now been removed from onnx.
@@ -154,6 +155,7 @@ Status RegisterQuantizationKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo,
BuildKernelCreateInfo,
BuildKernelCreateInfo,
+ BuildKernelCreateInfo,
};
for (auto& function_table_entry : function_table) {
diff --git a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc
index 452a245cb1..abad4abb75 100644
--- a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc
+++ b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc
@@ -166,6 +166,102 @@ void convTransposeWithDynamicPadsShapeInference(InferenceContext& ctx) {
return;
}
}
+
+void embedLayerNormalizationShapeInference(InferenceContext& ctx) {
+ propagateElemTypeFromInputToOutput(ctx, 2, 0);
+ propagateElemTypeFromInputToOutput(ctx, 0, 1);
+ if (!hasInputShape(ctx, 0)) {
+ // TODO(kreeger): In this case update the output to (?, ?, hidden_size).
+ return;
+ }
+
+ auto& input_ids_shape = getInputShape(ctx, 0);
+ auto& input_ids_dims = input_ids_shape.dim();
+
+ // Note that both batch size and sequence length could be symbolic.
+ // So we only check dimension size here.
+ if (input_ids_dims.size() != 2) {
+ fail_shape_inference("input_ids shall be 2 dimensions");
+ }
+
+ bool has_segment = hasInputShape(ctx, 1);
+ if (has_segment) {
+ // Ensure that segment_ids has the same shape.
+ auto& segment_ids_shape = getInputShape(ctx, 1);
+ auto& segment_ids_dims = segment_ids_shape.dim();
+ if (segment_ids_dims.size() != 2) {
+ fail_shape_inference("segment_ids input shall be 2 dimensions");
+ }
+ }
+
+ // get hidden_size from the last dimension of embedding
+ auto& word_embedding_shape = getInputShape(ctx, 2);
+ auto& word_embedding_dims = word_embedding_shape.dim();
+ if (word_embedding_dims.size() != 2 ||
+ !word_embedding_dims[1].has_dim_value() ||
+ word_embedding_shape.dim(1).dim_value() <= 0) {
+ fail_shape_inference("word_embedding should have 2 dimensions and dimension size is known.");
+ }
+ int64_t hidden_size = word_embedding_shape.dim(1).dim_value();
+
+ // Ensure that all embeddings + the gamma/beta tensors have the same hidden_size:
+ auto& position_embedding_shape = getInputShape(ctx, 3);
+ auto& position_embedding_dims = position_embedding_shape.dim();
+ if (position_embedding_dims.size() != 2 ||
+ !position_embedding_dims[1].has_dim_value() ||
+ position_embedding_shape.dim(1).dim_value() != hidden_size) {
+ fail_shape_inference(
+ "position_embedding should have 2 dimensions, dimension size known, "
+ "and same hidden size as word_embedding.");
+ }
+
+ if (has_segment) {
+ auto& segment_embedding_shape = getInputShape(ctx, 4);
+ auto& segment_embedding_dims = segment_embedding_shape.dim();
+ if (segment_embedding_dims.size() != 2 ||
+ !segment_embedding_dims[1].has_dim_value() ||
+ segment_embedding_shape.dim(1).dim_value() != hidden_size) {
+ fail_shape_inference(
+ "segment_embedding should have 2 dimensions, dimension size known, "
+ "and same hidden size as word_embedding.");
+ }
+ }
+
+ auto& gamma_shape = getInputShape(ctx, 5);
+ auto& gamma_dims = gamma_shape.dim();
+ if (gamma_dims.size() != 1 ||
+ !gamma_dims[0].has_dim_value() ||
+ gamma_shape.dim(0).dim_value() != hidden_size) {
+ fail_shape_inference(
+ "gamma should have 2 dimension, dimension size known, "
+ "and same hidden size as word_embedding.");
+ }
+
+ auto& beta_shape = getInputShape(ctx, 6);
+ auto& beta_dims = gamma_shape.dim();
+ if (beta_dims.size() != 1 ||
+ !beta_dims[0].has_dim_value() ||
+ beta_shape.dim(0).dim_value() != hidden_size) {
+ fail_shape_inference(
+ "beta should have 1 dimension, dimension size known, "
+ "and same hidden size as word_embedding.");
+ }
+
+ // input shape is (batch_size, sequence_length), output shape is (batch_size, sequence_length, hidden_size)
+ ONNX_NAMESPACE::TensorShapeProto output_shape;
+ *output_shape.add_dim() = input_ids_dims[0];
+ *output_shape.add_dim() = input_ids_dims[1];
+
+ output_shape.add_dim();
+ output_shape.mutable_dim(2)->set_dim_value(hidden_size);
+
+ updateOutputShape(ctx, 0, output_shape);
+
+ // mask_index shape is (batch_size)
+ ONNX_NAMESPACE::TensorShapeProto mask_index_shape;
+ *mask_index_shape.add_dim() = input_ids_dims[0];
+ updateOutputShape(ctx, 1, mask_index_shape);
+}
} // namespace ONNX_NAMESPACE
namespace onnxruntime {
@@ -505,46 +601,45 @@ will be calculated.)DOC";
.Output(1, "mask_index", "1D mask_index tensor with shape (batch_size)", "T1")
.TypeConstraint("T1", {"tensor(int32)"}, "Constrain input and output integer tensors types")
.TypeConstraint("T", {"tensor(float)", "tensor(float16)"}, "Constrain input and output float tensors types.")
- .TypeAndShapeInferenceFunction([](ONNX_NAMESPACE::InferenceContext& ctx) {
- propagateElemTypeFromInputToOutput(ctx, 2, 0);
- propagateElemTypeFromInputToOutput(ctx, 0, 1);
- if (!hasInputShape(ctx, 0))
- return;
+ .TypeAndShapeInferenceFunction(ONNX_NAMESPACE::embedLayerNormalizationShapeInference);
- auto& input_ids_shape = getInputShape(ctx, 0);
- auto& input_ids_dims = input_ids_shape.dim();
+ static const char* QEmbedLayerNormalization_ver1_doc = R"DOC(
+QEmbedLayerNormalization is the quantized fusion of embedding layer in BERT model, with optional mask processing.
+The embedding layer takes input_ids (word IDs) and segment_ids (sentence IDs) to look up word_embedding, position_embedding,
+and segment_emedding; the embeddings are added then applied layer normalization using gamma and beta tensors. The input_ids
+and segment_ids remain int32. All embeddings, gamma, and beta tensors are converted to int8/uint8. The last input mask is optional.
+If mask is provided, mask index (that is position of first 0 in mask, or number of words will be calculated.)DOC";
- // Note that both batch size and sequence length could be symbolic.
- // So we only check dimension size here.
- if (input_ids_dims.size() != 2) {
- fail_shape_inference("Inputs 0 shall be 2 dimensions");
- }
-
- // get hidden_size from the last dimension of embedding
- auto& word_embedding_shape = getInputShape(ctx, 3);
- auto& word_embedding_dims = word_embedding_shape.dim();
- if (word_embedding_dims.size() != 2 ||
- !word_embedding_dims[1].has_dim_value() ||
- word_embedding_shape.dim(1).dim_value() <= 0) {
- fail_shape_inference("word_embedding should have 2 dimensions and dimension size is known.");
- }
- int64_t hidden_size = word_embedding_shape.dim(1).dim_value();
-
- // input shape is (batch_size, sequence_length), output shape is (batch_size, sequence_length, hidden_size)
- ONNX_NAMESPACE::TensorShapeProto output_shape;
- for (auto& dim : input_ids_dims) {
- *output_shape.add_dim() = dim;
- }
- output_shape.add_dim();
- output_shape.mutable_dim(2)->set_dim_value(hidden_size);
-
- updateOutputShape(ctx, 0, output_shape);
-
- // mask_index shape is (batch_size)
- ONNX_NAMESPACE::TensorShapeProto mask_index_shape;
- *mask_index_shape.add_dim() = input_ids_dims[0];
- updateOutputShape(ctx, 1, mask_index_shape);
- });
+ ONNX_CONTRIB_OPERATOR_SCHEMA(QEmbedLayerNormalization)
+ .SetDomain(kMSDomain)
+ .SinceVersion(1)
+ .SetSupportLevel(OpSchema::SupportType::EXPERIMENTAL)
+ .SetDoc(QEmbedLayerNormalization_ver1_doc)
+ .Attr("epsilon", "The epsilon value to use to avoid division by zero.", AttributeProto::FLOAT, kDefaultEmbedLayerNormEpsilon)
+ .Input(0, "input_ids", "2D words IDs with shape (batch_size, sequence_length)", "T1")
+ .Input(1, "segment_ids", "2D segment IDs with shape (batch_size, sequence_length)", "T1", OpSchema::Optional)
+ .Input(2, "word_embedding_quant", "2D with shape (,hidden_size)", "T2")
+ .Input(3, "position_embedding_quant", "2D with shape (, hidden_size)", "T2")
+ .Input(4, "segment_embedding", "2D with shape (, hidden_size)", "T2", OpSchema::Optional)
+ .Input(5, "gamma_quant", "1D gamma tensor for layer normalization with shape (hidden_size)", "T2")
+ .Input(6, "beta_quant", "1D beta tensor for layer normalization with shape (hidden_size)", "T2")
+ .Input(7, "mask", "Mask", "T1", OpSchema::Optional)
+ .Input(8, "word_embedding_scale", "Scale for word embeddings", "T")
+ .Input(9, "position_embedding_scale", "Scale for position embeddings", "T")
+ .Input(10, "segment_embedding_scale", "Scale for segment embeddings", "T", OpSchema::Optional)
+ .Input(11, "gamma_scale", "Scale for 1D gamma tensor", "T")
+ .Input(12, "beta_scale", "Scale for 1D beta tensor", "T")
+ .Input(13, "word_embedding_zero_point", "Zero point for word embeddings", "T2")
+ .Input(14, "position_embedding_zero_point", "Zero point for position embeddings", "T2")
+ .Input(15, "segment_embedding_zero_point", "Zero Point for segment embeddings", "T2", OpSchema::Optional)
+ .Input(16, "gamma_zero_point", "Zero Point for 1D gamma tensor", "T2")
+ .Input(17, "beta_zero_point", "Zero Point for 1D beta tensor", "T2")
+ .Output(0, "layernorm_out", "LayerNorm Output", "T")
+ .Output(1, "mask_index_out", "Mask Index Output", "T1")
+ .TypeConstraint("T1", {"tensor(int32)"}, "Constrain mask index to integer types")
+ .TypeConstraint("T2", {"tensor(int8)", "tensor(uint8)"}, "Constrain input and output types to int8 tensors.")
+ .TypeConstraint("T", {"tensor(float)"}, "Constrain input and output types to float32 tensors.")
+ .TypeAndShapeInferenceFunction(ONNX_NAMESPACE::embedLayerNormalizationShapeInference);
static const char* FastGelu_ver1_doc = R"DOC(
GELU (Gaussian Error Linear Unit) approximation: Y=0.5*X*(1+tanh(0.797885*X+0.035677*X*X*X)) with an optional input of bias that will be added to X before GELU.)DOC";
diff --git a/onnxruntime/python/tools/quantization/operators/embed_layernorm.py b/onnxruntime/python/tools/quantization/operators/embed_layernorm.py
index 2e88a1fa00..0a6c2e57dc 100644
--- a/onnxruntime/python/tools/quantization/operators/embed_layernorm.py
+++ b/onnxruntime/python/tools/quantization/operators/embed_layernorm.py
@@ -1,11 +1,14 @@
import onnx
from .base_operator import QuantOperatorBase
+from ..quant_utils import attribute_to_kwarg, ms_domain
from onnx import onnx_pb as onnx_proto
-'''
-Quantize EmbedLayerNormalization
-'''
+'''
+Quantizes the EmbedLayerNorm fused ONNXRuntime Op.
+This Quant operator keeps the input and segment IDs at int32 but will quantize all initializer and
+weight inputs associated with the node to uint8.
+'''
class EmbedLayerNormalizationQuant(QuantOperatorBase):
def __init__(self, onnx_quantizer, onnx_node):
super().__init__(onnx_quantizer, onnx_node)
@@ -14,8 +17,82 @@ class EmbedLayerNormalizationQuant(QuantOperatorBase):
node = self.node
assert (node.op_type == "EmbedLayerNormalization")
+ '''
+ Pre-quantization EmbedLayerNorm inputs:
+ [0] input_ids (int32)
+ [1] segment_ids (int32)
+ [2] word_embedding (float32)
+ [3] position_embedding (float32)
+ [4] segment_embedding (float32)
+ [5] gamma (float32)
+ [6] beta (float32)
+ [7] mask (int32) (optional)
+ '''
(quantized_input_names, zero_point_names, scale_names, nodes) = \
- self.quantizer.quantize_inputs(node, [2, 3, 4])
+ self.quantizer.quantize_inputs(node, [2, 3, 4, 5, 6])
+
+ qembed_layer_norm_name = "" if node.name == "" else node.name + "_quant"
+
+ '''
+ Quantized Input Tensor List
+ [0] input_ids (int32)
+ [1] segment_ids (int32)
+ [2] word_embedding (uint8)
+ [3] position_embedding (uint8)
+ [4] segment_embedding (uint8)
+ [5] gamma (uint8)
+ [6] beta (uint8)
+ [7] mask (int32) (optional)
+ [8] word_embedding_scale (float)
+ [9] position_embedding_scale (float)
+ [10] segment_embedding_scale (float)
+ [11] gamma_scale (float)
+ [12] beta_scale (float)
+ [13] word_embedding_zero_point (uint8)
+ [14] position_embedding_zero_point (uint8)
+ [15] segment_embedding_zero_point (uint8)
+ [16] gamma_zero_point (uint8)
+ [17] beta_zero_point (uint8)
+ '''
+ inputs = []
+ # 'input_ids'
+ inputs.extend([node.input[0]])
+ # 'segment_ids'
+ inputs.extend([node.input[1]])
+ # 'word_embedding_quant'
+ inputs.extend([quantized_input_names[0]])
+ # 'position_embedding_quant'
+ inputs.extend([quantized_input_names[1]])
+ # 'segment_embedding_quant'
+ inputs.extend([quantized_input_names[2]])
+ # 'gamma_quant'
+ inputs.extend([quantized_input_names[3]])
+ # 'beta_quant'
+ inputs.extend([quantized_input_names[4]])
+ # 'mask' (optional)
+ inputs.extend([node.input[7] if len(node.input) > 7 else ""])
+
+ # Add all scales:
+ inputs.extend([scale_names[0]])
+ inputs.extend([scale_names[1]])
+ inputs.extend([scale_names[2]])
+ inputs.extend([scale_names[3]])
+ inputs.extend([scale_names[4]])
+
+ # Add all zero points:
+ inputs.extend([zero_point_names[0]])
+ inputs.extend([zero_point_names[1]])
+ inputs.extend([zero_point_names[2]])
+ inputs.extend([zero_point_names[3]])
+ inputs.extend([zero_point_names[4]])
+
+ kwargs = {}
+ for attribute in node.attribute:
+ kwargs.update(attribute_to_kwarg(attribute))
+ kwargs["domain"] = ms_domain
+
+ qembed_layer_norm_node = onnx.helper.make_node("QEmbedLayerNormalization", inputs, node.output,
+ qembed_layer_norm_name, **kwargs)
+ nodes.append(qembed_layer_norm_node)
- super().quantize()
self.quantizer.new_nodes += nodes
diff --git a/onnxruntime/test/contrib_ops/embed_layer_norm_test_vectors.h b/onnxruntime/test/contrib_ops/embed_layer_norm_test_vectors.h
new file mode 100644
index 0000000000..98f42cefe0
--- /dev/null
+++ b/onnxruntime/test/contrib_ops/embed_layer_norm_test_vectors.h
@@ -0,0 +1,314 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+#pragma once
+
+#include
+#include
+
+namespace onnxruntime {
+namespace test {
+namespace embedlayernorm {
+
+constexpr float kEpsilon = 1e-12f;
+
+// EmbedLayerNorm and QEmbedLayerNorm contain many inputs and outputs. This
+// utility class helps readability of Op unit tests by wrapping Op data.
+class OpData {
+ public:
+ explicit OpData(
+ int batch_size,
+ int sequence_size,
+ int hidden_size,
+ const std::vector& input_ids_data,
+ const std::vector& segment_ids_data,
+ const std::vector& mask_data,
+ const std::vector& word_embedding_data,
+ const std::vector& position_embedding_data,
+ const std::vector& segment_embedding_data,
+ const std::vector& gamma_data,
+ const std::vector& beta_data,
+ const std::vector& output_data,
+ const std::vector& mask_index_data,
+ float epsilon = kEpsilon,
+ bool has_mask = true,
+ bool has_segment = true)
+ : batch_size(batch_size)
+ , sequence_size(sequence_size)
+ , hidden_size(hidden_size)
+ , input_ids_data(input_ids_data)
+ , segment_ids_data(segment_ids_data)
+ , mask_data(mask_data)
+ , word_embedding_data(word_embedding_data)
+ , position_embedding_data(position_embedding_data)
+ , segment_embedding_data(segment_embedding_data)
+ , gamma_data(gamma_data)
+ , beta_data(beta_data)
+ , output_data(output_data)
+ , mask_index_data(mask_index_data)
+ , epsilon(epsilon)
+ , has_mask(has_mask)
+ , has_segment(has_segment)
+ {}
+
+ const int batch_size;
+ const int sequence_size;
+ const int hidden_size;
+ const std::vector input_ids_data;
+ const std::vector segment_ids_data;
+ const std::vector mask_data;
+ const std::vector word_embedding_data;
+ const std::vector position_embedding_data;
+ const std::vector segment_embedding_data;
+ const std::vector gamma_data;
+ const std::vector beta_data;
+ const std::vector output_data;
+ const std::vector mask_index_data;
+ const float epsilon;
+ const bool has_mask = true;
+ const bool has_segment = true;
+};
+
+inline OpData EmbedLayerNormBatch1() {
+ int batch_size = 1;
+ int sequence_size = 2;
+ int hidden_size = 4;
+
+ std::vector input_ids_data = {
+ 1, 3};
+
+ std::vector segment_ids_data = {
+ 0, 1};
+
+ std::vector mask_data = {
+ 1, 1};
+
+ std::vector word_embedding_data = {
+ 0.2f, 0.1f, 0.4f, -0.6f,
+ 0.3f, 0.2f, 0.5f, 0.6f,
+ 0.6f, 0.7f, 0.0f, -0.1f,
+ 0.8f, 0.6f, 0.9f, 1.2f,
+ 0.1f, 0.3f, 0.5f, 0.9f,
+ 1.0f, -2.0f, 1.1f, 0.8f};
+
+ std::vector position_embedding_data = {
+ 0.1f, 0.1f, 0.4f, 0.6f,
+ 0.6f, 0.0f, 0.8f, 0.6f,
+ 0.3f, 0.9f, -2.0f, 0.8f};
+
+ std::vector segment_embedding_data = {
+ 0.3f, 0.4f, 0.9f, 0.1f,
+ 0.7f, 0.3f, 0.5f, 0.2f};
+
+ std::vector gamma_data = {
+ 0.25f, 0.15f, 0.45f, -0.66f};
+
+ std::vector beta_data = {
+ 0.6f, 0.2f, 0.5f, -0.6f};
+
+ std::vector output_data = {
+ 0.36917170882225037, 0.061503000557422638, 1.1598974466323853, -0.85092413425445557,
+ 0.74301940202713013, -0.057434864342212677, 0.84324657917022705, -0.85171419382095337};
+
+ std::vector mask_index_data = {
+ 2};
+
+ return OpData(batch_size, sequence_size, hidden_size, input_ids_data, segment_ids_data,
+ mask_data, word_embedding_data, position_embedding_data, segment_embedding_data,
+ gamma_data, beta_data, output_data, mask_index_data);
+}
+
+inline OpData EmbedLayerNormBatch2(bool has_mask=true) {
+ int batch_size = 3;
+ int sequence_size = 2;
+ int hidden_size = 4;
+
+ std::vector input_ids_data = {
+ 1, 3,
+ 1, 3,
+ 2, 0};
+
+ std::vector segment_ids_data = {
+ 0, 1,
+ 0, 1,
+ 0, 0};
+
+ std::vector mask_data = {};
+ if (has_mask) {
+ mask_data = {
+ 1, 1,
+ 1, 1,
+ 1, 0};
+ }
+
+ std::vector word_embedding_data = {
+ 0.2f, 0.1f, 0.4f, -0.6f,
+ 0.3f, 0.2f, 0.5f, 0.6f,
+ 0.6f, 0.7f, 0.0f, -0.1f,
+ 0.8f, 0.6f, 0.9f, 1.2f,
+ 0.1f, 0.3f, 0.5f, 0.9f,
+ 1.0f, -2.0f, 1.1f, 0.8f};
+
+ std::vector position_embedding_data = {
+ 0.1f, 0.1f, 0.4f, 0.6f,
+ 0.6f, 0.0f, 0.8f, 0.6f,
+ 0.3f, 0.9f, -2.0f, 0.8f};
+
+ std::vector segment_embedding_data = {
+ 0.3f, 0.4f, 0.9f, 0.1f,
+ 0.7f, 0.3f, 0.5f, 0.2f};
+
+ std::vector gamma_data = {
+ 0.25f, 0.15f, 0.45f, -0.66f};
+
+ std::vector beta_data = {
+ 0.6f, 0.2f, 0.5f, -0.6f};
+
+ std::vector output_data = {
+ 0.36917170882225037, 0.061503000557422638, 1.1598974466323853, -0.85092413425445557,
+ 0.74301940202713013, -0.057434864342212677, 0.84324657917022705, -0.85171419382095337,
+ 0.36917170882225037, 0.061503000557422638, 1.1598974466323853, -0.85092413425445557,
+ 0.74301940202713013, -0.057434864342212677, 0.84324657917022705, -0.85171419382095337,
+ 0.57668739557266235, 0.2979130744934082, 0.96158987283706665, 0.44627034664154053,
+ 0.64977931976318359, 0.11039737612009048, 1.1869535446166992, 0.14469735324382782};
+
+ std::vector mask_index_data;
+ if (has_mask) {
+ mask_index_data = {2, 2, 1};
+ } else {
+ mask_index_data = {0, 0, 0};
+ }
+
+ return OpData(batch_size, sequence_size, hidden_size, input_ids_data, segment_ids_data,
+ mask_data, word_embedding_data, position_embedding_data, segment_embedding_data,
+ gamma_data, beta_data, output_data, mask_index_data, kEpsilon, has_mask);
+}
+
+inline OpData EmbedLayerNormLargeBatchSmallHiddenSize() {
+ int batch_size = 5;
+ int sequence_size = 2;
+ int hidden_size = 4;
+
+ std::vector input_ids_data = {
+ 1, 3,
+ 1, 3,
+ 2, 0,
+ 1, 3,
+ 2, 0};
+
+ std::vector segment_ids_data = {
+ 0, 1,
+ 0, 1,
+ 0, 0,
+ 0, 1,
+ 0, 0};
+
+ std::vector mask_data = {
+ 1, 1,
+ 1, 1,
+ 1, 0,
+ 1, 1,
+ 1, 0};
+
+ std::vector word_embedding_data = {
+ 0.2f, 0.1f, 0.4f, -0.6f,
+ 0.3f, 0.2f, 0.5f, 0.6f,
+ 0.6f, 0.7f, 0.0f, -0.1f,
+ 0.8f, 0.6f, 0.9f, 1.2f,
+ 0.1f, 0.3f, 0.5f, 0.9f,
+ 1.0f, -2.0f, 1.1f, 0.8f};
+
+ std::vector position_embedding_data = {
+ 0.1f, 0.1f, 0.4f, 0.6f,
+ 0.6f, 0.0f, 0.8f, 0.6f,
+ 0.3f, 0.9f, -2.0f, 0.8f};
+
+ std::vector segment_embedding_data = {
+ 0.3f, 0.4f, 0.9f, 0.1f,
+ 0.7f, 0.3f, 0.5f, 0.2f};
+
+ std::vector gamma_data = {
+ 0.25f, 0.15f, 0.45f, -0.66f};
+
+ std::vector beta_data = {
+ 0.6f, 0.2f, 0.5f, -0.6f};
+
+ std::vector output_data = {
+ 0.36917170882225037, 0.061503000557422638, 1.1598974466323853, -0.85092413425445557,
+ 0.74301940202713013, -0.057434864342212677, 0.84324657917022705, -0.85171419382095337,
+ 0.36917170882225037, 0.061503000557422638, 1.1598974466323853, -0.85092413425445557,
+ 0.74301940202713013, -0.057434864342212677, 0.84324657917022705, -0.85171419382095337,
+ 0.57668739557266235, 0.2979130744934082, 0.96158987283706665, 0.44627034664154053,
+ 0.64977931976318359, 0.11039737612009048, 1.1869535446166992, 0.14469735324382782,
+ 0.36917170882225037, 0.061503000557422638, 1.1598974466323853, -0.85092413425445557,
+ 0.74301940202713013, -0.057434864342212677, 0.84324657917022705, -0.85171419382095337,
+ 0.57668739557266235, 0.2979130744934082, 0.96158987283706665, 0.44627034664154053,
+ 0.64977931976318359, 0.11039737612009048, 1.1869535446166992, 0.14469735324382782};
+
+ std::vector mask_index_data = {
+ 2, 2, 1, 2, 1};
+
+ return OpData(batch_size, sequence_size, hidden_size, input_ids_data, segment_ids_data,
+ mask_data, word_embedding_data, position_embedding_data, segment_embedding_data,
+ gamma_data, beta_data, output_data, mask_index_data);
+}
+
+inline OpData EmbedLayerNormBatch_Distill() {
+ int batch_size = 3;
+ int sequence_size = 2;
+ int hidden_size = 4;
+
+ std::vector input_ids_data = {
+ 1, 3,
+ 1, 3,
+ 2, 0};
+
+ std::vector segment_ids_data = {};
+
+ std::vector mask_data = {
+ 1, 1,
+ 1, 1,
+ 1, 0};
+
+ std::vector word_embedding_data = {
+ 0.2f, 0.1f, 0.4f, -0.6f,
+ 0.3f, 0.2f, 0.5f, 0.6f,
+ 0.6f, 0.7f, 0.0f, -0.1f,
+ 0.8f, 0.6f, 0.9f, 1.2f,
+ 0.1f, 0.3f, 0.5f, 0.9f,
+ 1.0f, -2.0f, 1.1f, 0.8f};
+
+ std::vector position_embedding_data = {
+ 0.1f, 0.1f, 0.4f, 0.6f,
+ 0.6f, 0.0f, 0.8f, 0.6f,
+ 0.3f, 0.9f, -2.0f, 0.8f};
+
+ std::vector segment_embedding_data = {};
+
+ std::vector gamma_data = {
+ 0.25f, 0.15f, 0.45f, -0.66f};
+
+ std::vector beta_data = {
+ 0.6f, 0.2f, 0.5f, -0.6f};
+
+ std::vector output_data = {
+ 0.39587587118148804, 0.03670068085193634, 0.7449488639831543, -1.4981462955474854,
+ 0.61326867341995239, -0.046796366572380066, 0.81048583984375, -1.1954958438873291,
+ 0.39587587118148804, 0.03670068085193634, 0.7449488639831543, -1.4981462955474854,
+ 0.61326867341995239, -0.046796366572380066, 0.81048583984375, -1.1954958438873291,
+ 0.75811392068862915, 0.38973665237426758, -0.069209933280944824, -0.18257927894592285,
+ 0.73836749792098999, 0.071695566177368164, 1.111332893371582, 0.097372293472290039};
+
+ std::vector mask_index_data = {
+ 2, 2, 1};
+
+ return OpData(batch_size, sequence_size, hidden_size, input_ids_data, segment_ids_data,
+ mask_data, word_embedding_data, position_embedding_data, segment_embedding_data,
+ gamma_data, beta_data, output_data, mask_index_data, kEpsilon,
+ /*has_mask=*/true,
+ /*has_segment=*/false);
+}
+
+} // namespace embedlayernorm
+} // namespace test
+} // namespace onnxruntime
diff --git a/onnxruntime/test/contrib_ops/embedlayernorm_op_test.cc b/onnxruntime/test/contrib_ops/embedlayernorm_op_test.cc
index 177d8d745b..5ca8185040 100644
--- a/onnxruntime/test/contrib_ops/embedlayernorm_op_test.cc
+++ b/onnxruntime/test/contrib_ops/embedlayernorm_op_test.cc
@@ -1,6 +1,8 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
+// TODO(kreeger): rename this file "embed_layer_norm_op_test.cc" to match existing styles.
+#include "embed_layer_norm_test_vectors.h"
#include "gtest/gtest.h"
#include "test/common/tensor_op_test_utils.h"
#include "test/common/cuda_op_test_utils.h"
@@ -8,26 +10,9 @@
namespace onnxruntime {
namespace test {
-constexpr float epsilon_ = 1e-12f;
-static void RunTest(
- const std::vector& input_ids_data,
- const std::vector& segment_ids_data,
- const std::vector& mask_data,
- const std::vector& word_embedding_data,
- const std::vector& position_embedding_data,
- const std::vector& segment_embedding_data,
- const std::vector& gamma_data,
- const std::vector& beta_data,
- const std::vector& output_data,
- const std::vector& mask_index_data,
- float epsilon,
- int batch_size,
- int sequence_length,
- int hidden_size,
- bool use_float16 = false,
- bool has_mask = true,
- bool has_segment = true) {
+static void RunTest(const embedlayernorm::OpData& data,
+ bool use_float16 = false) {
int min_cuda_architecture = use_float16 ? 530 : 0;
bool enable_cuda = HasCudaEnvironment(min_cuda_architecture);
@@ -35,484 +20,120 @@ static void RunTest(
if (enable_cpu || enable_cuda) {
// Input and output shapes
- // Input 0 - input_ids : (batch_size, sequence_length)
- // Input 1 - segment_ids : (batch_size, sequence_length)
+ // Input 0 - input_ids : (batch_size, sequence_size)
+ // Input 1 - segment_ids : (batch_size, sequence_size)
// Input 2 - word_embedding : (,hidden_size)
// Input 3 - position_embedding : (,hidden_size)
// Input 4 - segment_embedding : (,hidden_size)
// Input 5 - gamma : (hidden_size)
// Input 6 - beta : (hidden_size)
- // Input 7 - mask : (batch_size, sequence_length)
- // Output 0 - output : (batch_size, sequence_length, hidden_size)
+ // Input 7 - mask : (batch_size, sequence_size)
+ // Output 0 - output : (batch_size, sequence_size, hidden_size)
// Output 1 - mask_index : (batch_size)
- std::vector input_ids_dims = {batch_size, sequence_length};
- std::vector segment_ids_dims = {batch_size, sequence_length};
- std::vector mask_dims = {batch_size, sequence_length};
+ std::vector input_ids_dims = {data.batch_size, data.sequence_size};
+ std::vector segment_ids_dims = {data.batch_size, data.sequence_size};
+ std::vector mask_dims = {data.batch_size, data.sequence_size};
- ASSERT_TRUE(word_embedding_data.size() % hidden_size == 0);
- std::vector word_embedding_dims = {static_cast(word_embedding_data.size() / hidden_size), hidden_size};
+ ASSERT_TRUE(data.word_embedding_data.size() % data.hidden_size == 0);
+ std::vector word_embedding_dims = {
+ static_cast(data.word_embedding_data.size() / data.hidden_size),
+ data.hidden_size};
- ASSERT_TRUE(position_embedding_data.size() % hidden_size == 0);
- std::vector position_embedding_dims = {static_cast(position_embedding_data.size() / hidden_size), hidden_size};
+ ASSERT_TRUE(data.position_embedding_data.size() % data.hidden_size == 0);
+ std::vector position_embedding_dims = {
+ static_cast(data.position_embedding_data.size() / data.hidden_size),
+ data.hidden_size};
- ASSERT_TRUE(segment_embedding_data.size() % hidden_size == 0);
- std::vector segment_embedding_dims = {static_cast(segment_embedding_data.size() / hidden_size), hidden_size};
+ ASSERT_TRUE(data.segment_embedding_data.size() % data.hidden_size == 0);
+ std::vector segment_embedding_dims = {
+ static_cast(data.segment_embedding_data.size() / data.hidden_size),
+ data.hidden_size};
- std::vector gamma_dims = {hidden_size};
+ std::vector gamma_dims = {data.hidden_size};
std::vector beta_dims = gamma_dims;
- std::vector output_dims = {batch_size, sequence_length, hidden_size};
- std::vector mask_index_dims = {batch_size};
+ std::vector output_dims = {data.batch_size, data.sequence_size, data.hidden_size};
+ std::vector mask_index_dims = {data.batch_size};
OpTester tester("EmbedLayerNormalization", 1, onnxruntime::kMSDomain);
- tester.AddInput("input_ids", input_ids_dims, input_ids_data);
- if (!has_segment) {
+ tester.AddInput("input_ids", input_ids_dims, data.input_ids_data);
+ if (!data.has_segment) {
tester.AddMissingOptionalInput();
} else {
- tester.AddInput("segment_ids", segment_ids_dims, segment_ids_data);
+ tester.AddInput("segment_ids", segment_ids_dims, data.segment_ids_data);
}
if (use_float16) {
- tester.AddInput("word_embedding", word_embedding_dims, ToFloat16(word_embedding_data));
- tester.AddInput("position_embedding", position_embedding_dims, ToFloat16(position_embedding_data));
- if (!has_segment) {
+ tester.AddInput("word_embedding",
+ word_embedding_dims,
+ ToFloat16(data.word_embedding_data));
+ tester.AddInput("position_embedding",
+ position_embedding_dims,
+ ToFloat16(data.position_embedding_data));
+ if (!data.has_segment) {
tester.AddMissingOptionalInput();
} else {
- tester.AddInput("segment_embedding", segment_embedding_dims, ToFloat16(segment_embedding_data));
+ tester.AddInput("segment_embedding",
+ segment_embedding_dims,
+ ToFloat16(data.segment_embedding_data));
}
- tester.AddInput("gamma", gamma_dims, ToFloat16(gamma_data));
- tester.AddInput("beta", beta_dims, ToFloat16(beta_data));
- tester.AddAttribute("epsilon", epsilon);
- if (has_mask) {
- tester.AddInput("mask", mask_dims, mask_data);
+ tester.AddInput("gamma", gamma_dims, ToFloat16(data.gamma_data));
+ tester.AddInput("beta", beta_dims, ToFloat16(data.beta_data));
+ tester.AddAttribute("epsilon", data.epsilon);
+ if (data.has_mask) {
+ tester.AddInput("mask", mask_dims, data.mask_data);
}
- tester.AddOutput("output", output_dims, ToFloat16(output_data));
+ tester.AddOutput("output", output_dims, ToFloat16(data.output_data));
} else {
- tester.AddInput("word_embedding", word_embedding_dims, word_embedding_data);
- tester.AddInput("position_embedding", position_embedding_dims, position_embedding_data);
- if (!has_segment) {
+ tester.AddInput("word_embedding",
+ word_embedding_dims,
+ data.word_embedding_data);
+ tester.AddInput("position_embedding",
+ position_embedding_dims,
+ data.position_embedding_data);
+ if (!data.has_segment) {
tester.AddMissingOptionalInput();
} else {
- tester.AddInput("segment_embedding", segment_embedding_dims, segment_embedding_data);
+ tester.AddInput("segment_embedding",
+ segment_embedding_dims,
+ data.segment_embedding_data);
}
- tester.AddInput("gamma", gamma_dims, gamma_data);
- tester.AddInput("beta", beta_dims, beta_data);
- tester.AddAttribute("epsilon", epsilon);
- if (has_mask) {
- tester.AddInput("mask", mask_dims, mask_data);
+ tester.AddInput("gamma", gamma_dims, data.gamma_data);
+ tester.AddInput("beta", beta_dims, data.beta_data);
+ tester.AddAttribute("epsilon", data.epsilon);
+ if (data.has_mask) {
+ tester.AddInput("mask", mask_dims, data.mask_data);
}
- tester.AddOutput("output", output_dims, output_data);
+ tester.AddOutput("output", output_dims, data.output_data);
}
- tester.AddOutput("mask_index", mask_index_dims, mask_index_data);
+ tester.AddOutput("mask_index", mask_index_dims, data.mask_index_data);
tester.Run();
}
}
TEST(EmbedLayerNormTest, EmbedLayerNormBatch1) {
- int batch_size = 1;
- int sequence_length = 2;
- int hidden_size = 4;
-
- std::vector input_ids_data = {
- 1, 3};
-
- std::vector segment_ids_data = {
- 0, 1};
-
- std::vector mask_data = {
- 1, 1};
-
- std::vector word_embedding_data = {
- 0.2f, 0.1f, 0.4f, -0.6f,
- 0.3f, 0.2f, 0.5f, 0.6f,
- 0.6f, 0.7f, 0.0f, -0.1f,
- 0.8f, 0.6f, 0.9f, 1.2f,
- 0.1f, 0.3f, 0.5f, 0.9f,
- 1.0f, -2.0f, 1.1f, 0.8f};
-
- std::vector position_embedding_data = {
- 0.1f, 0.1f, 0.4f, 0.6f,
- 0.6f, 0.0f, 0.8f, 0.6f,
- 0.3f, 0.9f, -2.0f, 0.8f};
-
- std::vector segment_embedding_data = {
- 0.3f, 0.4f, 0.9f, 0.1f,
- 0.7f, 0.3f, 0.5f, 0.2f};
-
- std::vector