From 800b62a139733f6ec51903136b8ba86a6d7ca80e Mon Sep 17 00:00:00 2001 From: Nick Kreeger Date: Fri, 25 Jun 2021 17:51:43 -0500 Subject: [PATCH] Create a quantized EmbedLayerNorm for ORT. (#8124) Create a quantized EmbedLayerNorm Op for ORT --- docs/ContribOperators.md | 81 +++ docs/OperatorKernels.md | 1 + .../contrib_ops/cpu/bert/embed_layer_norm.cc | 23 +- .../contrib_ops/cpu/bert/embed_layer_norm.h | 18 +- .../cpu/bert/embed_layer_norm_helper.cc | 32 +- .../contrib_ops/cpu/bert/qembed_layer_norm.cc | 306 +++++++++++ .../contrib_ops/cpu/bert/qembed_layer_norm.h | 22 + .../contrib_ops/cpu/cpu_contrib_kernels.cc | 2 + .../core/graph/contrib_ops/contrib_defs.cc | 171 ++++-- .../quantization/operators/embed_layernorm.py | 87 ++- .../embed_layer_norm_test_vectors.h | 314 +++++++++++ .../contrib_ops/embedlayernorm_op_test.cc | 517 +++--------------- .../contrib_ops/qembed_layer_norm_op_test.cc | 191 +++++++ .../quantization/test_op_embed_layernorm.py | 23 +- .../kernel_def_hashes/contrib.cpu.json | 6 +- 15 files changed, 1266 insertions(+), 528 deletions(-) create mode 100644 onnxruntime/contrib_ops/cpu/bert/qembed_layer_norm.cc create mode 100644 onnxruntime/contrib_ops/cpu/bert/qembed_layer_norm.h create mode 100644 onnxruntime/test/contrib_ops/embed_layer_norm_test_vectors.h create mode 100644 onnxruntime/test/contrib_ops/qembed_layer_norm_op_test.cc 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 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}; - - RunTest(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, - epsilon_, - batch_size, - sequence_length, - hidden_size); + RunTest(embedlayernorm::EmbedLayerNormBatch1()); } TEST(EmbedLayerNormTest, EmbedLayerNormBatch1_Float16) { - 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 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.369873046875, 0.061676025390625, 1.1591796875, -0.8515625, - 0.7431640625, -0.057586669921875, 0.84326171875, -0.8525390625}; - - std::vector mask_index_data = { - 2}; - - RunTest(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, - epsilon_, - batch_size, - sequence_length, - hidden_size, - true); + RunTest(embedlayernorm::EmbedLayerNormBatch1(), /*use_float16=*/true); } TEST(EmbedLayerNormTest, EmbedLayerNormBatch2) { - int batch_size = 3; - int sequence_length = 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 = { - 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 = { - 2, 2, 1}; - - RunTest(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, - epsilon_, - batch_size, - sequence_length, - hidden_size); + RunTest(embedlayernorm::EmbedLayerNormBatch2()); } TEST(EmbedLayerNormTest, EmbedLayerNormBatch2_NoMask) { - int batch_size = 3; - int sequence_length = 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 = {}; - - 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 = {0, 0, 0}; - - RunTest(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, - epsilon_, - batch_size, - sequence_length, - hidden_size, - false, - false); // no mask + RunTest(embedlayernorm::EmbedLayerNormBatch2(/*has_mask=*/false)); } // BatchSize > HiddenSize to reproduce mask processing bug TEST(EmbedLayerNormTest, EmbedLayerNormLargeBatchSmallHiddenSize) { - int batch_size = 5; - int sequence_length = 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}; - - RunTest(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, - epsilon_, - batch_size, - sequence_length, - hidden_size); + RunTest(embedlayernorm::EmbedLayerNormLargeBatchSmallHiddenSize()); } TEST(EmbedLayerNormTest, EmbedLayerNormBatch_Distill) { - int batch_size = 3; - int sequence_length = 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}; - - RunTest(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, - epsilon_, - batch_size, - sequence_length, - hidden_size, - false, - true, - false); + RunTest(embedlayernorm::EmbedLayerNormBatch_Distill()); } + } // namespace test } // namespace onnxruntime diff --git a/onnxruntime/test/contrib_ops/qembed_layer_norm_op_test.cc b/onnxruntime/test/contrib_ops/qembed_layer_norm_op_test.cc new file mode 100644 index 0000000000..6d0b08fbc2 --- /dev/null +++ b/onnxruntime/test/contrib_ops/qembed_layer_norm_op_test.cc @@ -0,0 +1,191 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "embed_layer_norm_test_vectors.h" +#include "gtest/gtest.h" +#include "test/common/quantization_test_utils.h" +#include "test/common/tensor_op_test_utils.h" +#include "test/providers/provider_test_utils.h" + +namespace onnxruntime { +namespace test { + +namespace { + +static void RunTest(const embedlayernorm::OpData& data, + float accuracy_threshold = 0.25f) { + ASSERT_TRUE(data.word_embedding_data.size() % data.hidden_size == 0); + ASSERT_TRUE(data.position_embedding_data.size() % data.hidden_size == 0); + ASSERT_TRUE(data.segment_embedding_data.size() % data.hidden_size == 0); + + std::vector input_ids_dims = {data.batch_size, data.sequence_size}; + std::vector segment_ids_dims = {data.batch_size, data.sequence_size}; + std::vector word_embedding_dims = { + static_cast(data.word_embedding_data.size() / data.hidden_size), + data.hidden_size}; + std::vector position_embedding_dims = { + static_cast(data.position_embedding_data.size() / data.hidden_size), + data.hidden_size}; + std::vector segment_embedding_dims = { + static_cast(data.segment_embedding_data.size() / data.hidden_size), + data.hidden_size}; + std::vector gamma_dims = {data.hidden_size}; + std::vector beta_dims = {data.hidden_size}; + std::vector output_dims = {data.batch_size, data.sequence_size, data.hidden_size}; + std::vector mask_index_dims = {data.batch_size}; + + float word_embedding_scale = 0.0f; + uint8_t word_embedding_zero_point = 0; + std::vector word_embedding_data_quant = + QuantizeLinear( + data.word_embedding_data, word_embedding_scale, word_embedding_zero_point); + + float position_embedding_scale = 0.0f; + uint8_t position_embedding_zero_point = 0; + std::vector position_embedding_data_quant = + QuantizeLinear( + data.position_embedding_data, position_embedding_scale, position_embedding_zero_point); + + float segment_embedding_scale = 0.0f; + uint8_t segment_embedding_zero_point = 0; + std::vector segment_embedding_data_quant; + if (data.has_segment) { + segment_embedding_data_quant = + QuantizeLinear( + data.segment_embedding_data, segment_embedding_scale, segment_embedding_zero_point); + } + + float gamma_scale = 0.0f; + uint8_t gamma_zero_point = 0; + std::vector gamma_data_quant = + QuantizeLinear( + data.gamma_data, gamma_scale, gamma_zero_point); + + float beta_scale = 0.0f; + uint8_t beta_zero_point = 0; + std::vector beta_data_quant = + QuantizeLinear( + data.beta_data, beta_scale, beta_zero_point); + + OpTester tester("QEmbedLayerNormalization", 1, onnxruntime::kMSDomain); + + // Operator inputs passed in at int32_t: + tester.AddInput("input_ids", input_ids_dims, data.input_ids_data); + if (data.has_segment) { + tester.AddInput("segment_ids", segment_ids_dims, data.segment_ids_data); + } else { + tester.AddMissingOptionalInput(); + } + + // Quantized initializer inputs: + tester.AddInput("word_embedding_data", + word_embedding_dims, + word_embedding_data_quant); + tester.AddInput("position_embedding_data", + position_embedding_dims, + position_embedding_data_quant); + if (data.has_segment) { + tester.AddInput("segment_embedding_data", + segment_embedding_dims, + segment_embedding_data_quant); + } else { + tester.AddMissingOptionalInput(); + } + tester.AddInput("gamma", + gamma_dims, + gamma_data_quant); + tester.AddInput("beta", + beta_dims, + beta_data_quant); + if (data.has_mask) { + std::vector mask_dims = {data.batch_size, data.sequence_size}; + tester.AddInput("mask", mask_dims, data.mask_data); + } else { + tester.AddMissingOptionalInput(); + } + + // Quantized scales: + tester.AddInput("word_embedding_scale", + /*dims=*/{}, + {word_embedding_scale}); + tester.AddInput("position_embedding_scale", + /*dims=*/{}, + {position_embedding_scale}); + if (data.has_segment) { + tester.AddInput("segment_embedding_scale", + /*dims=*/{}, + {segment_embedding_scale}); + } else { + tester.AddMissingOptionalInput(); + } + tester.AddInput("gamma_scale", + /*dims=*/{}, + {gamma_scale}); + tester.AddInput("beta_scale", + /*dims=*/{}, + {beta_scale}); + + // Quantized zero points: + tester.AddInput("word_embedding_zero_point", + /*dims=*/{}, + {word_embedding_zero_point}); + tester.AddInput("position_embedding_zero_point", + /*dims=*/{}, + {position_embedding_zero_point}); + if (data.has_segment) { + tester.AddInput("segment_embedding_zero_point", + /*dims=*/{}, + {segment_embedding_zero_point}); + } else { + tester.AddMissingOptionalInput(); + } + tester.AddInput("gamma_zero_point", + /*dims=*/{}, + {gamma_zero_point}); + tester.AddInput("beta_zero_point", + /*dims=*/{}, + {beta_zero_point}); + // Outputs: + tester.AddOutput("output", output_dims, data.output_data); + tester.AddOutput("mask_index", mask_index_dims, data.mask_index_data); + + // Floating point test vectors are quantized, passed through the operator, + // and dequantized. This dance will result in some loss in precision, ensure + // the test framework accounts for this loss: + tester.SetOutputAbsErr("output", accuracy_threshold); + + // Attributes: + tester.AddAttribute("epsilon", embedlayernorm::kEpsilon); + + tester.Run(); +} + +} // namespace + +TEST(QEmbedLayerNormTest, EmbedLayerNormBatch1) { + RunTest(embedlayernorm::EmbedLayerNormBatch1()); +} + +TEST(QEmbedLayerNormTest, EmbedLayerNormBatch1_Float16) { + RunTest(embedlayernorm::EmbedLayerNormBatch1(), /*use_float16=*/true); +} + +TEST(QEmbedLayerNormTest, EmbedLayerNormBatch2) { + RunTest(embedlayernorm::EmbedLayerNormBatch2()); +} + +TEST(QEmbedLayerNormTest, EmbedLayerNormBatch2_NoMask) { + RunTest(embedlayernorm::EmbedLayerNormBatch2(/*has_mask=*/false)); +} + +// BatchSize > HiddenSize to reproduce mask processing bug +TEST(QEmbedLayerNormTest, EmbedLayerNormLargeBatchSmallHiddenSize) { + RunTest(embedlayernorm::EmbedLayerNormLargeBatchSmallHiddenSize()); +} + +TEST(QEmbedLayerNormTest, EmbedLayerNormBatch_Distill) { + RunTest(embedlayernorm::EmbedLayerNormBatch_Distill()); +} + +} // namespace test +} // namespace onnxruntime diff --git a/onnxruntime/test/python/quantization/test_op_embed_layernorm.py b/onnxruntime/test/python/quantization/test_op_embed_layernorm.py index 5f67c7e9ee..75c351648d 100644 --- a/onnxruntime/test/python/quantization/test_op_embed_layernorm.py +++ b/onnxruntime/test/python/quantization/test_op_embed_layernorm.py @@ -53,13 +53,13 @@ class TestOpEmbedLayerNormalization(unittest.TestCase): seg_embed_weights = np.random.random_sample(seg_embed_shape).astype(dtype='float32') seg_embed_initializer = onnx.numpy_helper.from_array(seg_embed_weights, name='seg_embed') - layer_norm_weight_shape = [hidden_size] - layer_norm_weights = np.random.random_sample(layer_norm_weight_shape).astype(dtype='float32') - layer_norm_weights_initializer = onnx.numpy_helper.from_array(layer_norm_weights, name='layer_norm_weight') + gamma_shape = [hidden_size] + gamma = np.random.random_sample(gamma_shape).astype(dtype='float32') + gamma_initializer = onnx.numpy_helper.from_array(gamma, name='gamma') - layer_norm_bias_shape = [hidden_size] - layer_norm_bias_weights = np.random.random_sample(layer_norm_bias_shape).astype(dtype='float32') - layer_norm_bias_initializer = onnx.numpy_helper.from_array(layer_norm_bias_weights, name='layer_norm_bias') + beta_shape = [hidden_size] + beta = np.random.random_sample(beta_shape).astype(dtype='float32') + beta_initializer = onnx.numpy_helper.from_array(beta, name='beta') # EmbedLayerNormalization Outputs: layernorm_out_shape = [batch, sequence_length, hidden_size] @@ -70,7 +70,7 @@ class TestOpEmbedLayerNormalization(unittest.TestCase): # EmbedLayerNormalization Node: embed_layer_norm_inputs = [ - 'input_ids', 'segment_ids', 'word_embed', 'pos_embed', 'seg_embed', 'layer_norm_weight', 'layer_norm_bias' + 'input_ids', 'segment_ids', 'word_embed', 'pos_embed', 'seg_embed', 'gamma', 'beta' ] embed_layer_norm_outputs = ['layernorm_out', 'mask_index_out'] embed_layer_norm_node = helper.make_node('EmbedLayerNormalization', @@ -84,8 +84,7 @@ class TestOpEmbedLayerNormalization(unittest.TestCase): inputs = [input_ids_tensor, segment_ids_tensor] outputs = [layernorm_out_tensor, mask_index_out_tensor] initializers = [ - word_embed_initializer, pos_embed_initializer, seg_embed_initializer, layer_norm_weights_initializer, - layer_norm_bias_initializer + word_embed_initializer, pos_embed_initializer, seg_embed_initializer, gamma_initializer, beta_initializer ] graph = helper.make_graph(nodes, graph_name, inputs, outputs, initializer=initializers) @@ -110,7 +109,8 @@ class TestOpEmbedLayerNormalization(unittest.TestCase): quantize_dynamic(model_f32_path, model_uint8_path) - qnode_counts = {'DequantizeLinear': 3} + # Quantization should not have any DequantizeLinear nodes: + qnode_counts = {'DequantizeLinear': 0, 'QEmbedLayerNormalization': 1} check_op_type_count(self, model_uint8_path, **qnode_counts) data_reader.rewind() @@ -133,7 +133,8 @@ class TestOpEmbedLayerNormalization(unittest.TestCase): quantize_dynamic(model_f32_path, model_uint8_path) - qnode_counts = {'DequantizeLinear': 3} + # Quantization should not have any DequantizeLinear nodes: + qnode_counts = {'DequantizeLinear': 0, 'QEmbedLayerNormalization': 1} check_op_type_count(self, model_uint8_path, **qnode_counts) data_reader.rewind() diff --git a/onnxruntime/test/testdata/kernel_def_hashes/contrib.cpu.json b/onnxruntime/test/testdata/kernel_def_hashes/contrib.cpu.json index 6d79f1b5cc..869996aa7b 100644 --- a/onnxruntime/test/testdata/kernel_def_hashes/contrib.cpu.json +++ b/onnxruntime/test/testdata/kernel_def_hashes/contrib.cpu.json @@ -254,5 +254,9 @@ [ "QLinearConcat com.microsoft CPUExecutionProvider", 1734858160766311432 + ], + [ + "QEmbedLayerNormalization com.microsoft CPUExecutionProvider", + 9235385557940152248 ] -] \ No newline at end of file +]