From d8ed662f7f2e1de66a98e1f8d34ccda73c092317 Mon Sep 17 00:00:00 2001 From: Robert Ogden Date: Fri, 14 Jun 2024 12:56:17 -0700 Subject: [PATCH] add int8 quantization support to seq_flow_lite ops --- .../tflite_ops/sequence_string_projection.cc | 32 +++++- .../tflite_ops/tflite_qrnn_pooling.cc | 99 +++++++++++++++---- 2 files changed, 107 insertions(+), 24 deletions(-) diff --git a/third_party/tensorflow_models/src/research/seq_flow_lite/tflite_ops/sequence_string_projection.cc b/third_party/tensorflow_models/src/research/seq_flow_lite/tflite_ops/sequence_string_projection.cc index f53cf650f16f8..22f4c609e86c9 100644 --- a/third_party/tensorflow_models/src/research/seq_flow_lite/tflite_ops/sequence_string_projection.cc +++ b/third_party/tensorflow_models/src/research/seq_flow_lite/tflite_ops/sequence_string_projection.cc @@ -28,8 +28,8 @@ limitations under the License. #include "flatbuffers/flexbuffers.h" // flatbuffer #include "tensorflow/lite/string_util.h" #include "tf_ops/projection_normalizer_util.h" // seq_flow_lite -#include "tf_ops/projection_util.h" // seq_flow_lite -#include "tflite_ops/quantization_util.h" // seq_flow_lite +#include "tf_ops/projection_util.h" // seq_flow_lite +#include "tflite_ops/quantization_util.h" // seq_flow_lite namespace seq_flow_lite { namespace ops { @@ -144,8 +144,17 @@ class ProjectionParams { void WordNoveltyFeature(uint8_t* data, int word_count) const { float word_novelty_feature; WordNoveltyFeature(&word_novelty_feature, word_count); - *data = PodQuantize(word_novelty_feature, 127.0f, 127); + *data = PodQuantize(word_novelty_feature, /*zero_point=*/127, + /*inverse_scale=*/127.0f); } + + void WordNoveltyFeature(int8_t* data, int word_count) const { + float word_novelty_feature; + WordNoveltyFeature(&word_novelty_feature, word_count); + *data = PodQuantize(word_novelty_feature, /*zero_point=*/0, + /*inverse_scale=*/127.5f); + } + bool DocSizeFeatureEnabled() const { return (doc_size_levels_ != 0); } bool FirstCap() const { return add_first_cap_feature_; } bool AllCaps() const { return add_all_caps_feature_; } @@ -161,8 +170,17 @@ class ProjectionParams { void DocSizeFeature(uint8_t* data, int num_tokens) { float doc_size_feature; DocSizeFeature(&doc_size_feature, num_tokens); - *data = PodQuantize(doc_size_feature, 127.0f, 127); + *data = PodQuantize(doc_size_feature, /*zero_point=*/127, + /*inverse_scale=*/127.0f); } + + void DocSizeFeature(int8_t* data, int num_tokens) { + float doc_size_feature; + DocSizeFeature(&doc_size_feature, num_tokens); + *data = PodQuantize(doc_size_feature, /*zero_point=*/0, + /*inverse_scale=*/127.5f); + } + void Hash(const std::string& word, std::vector& hash_codes) { hasher_->GetHashCodes(word, hash_codes); } @@ -473,11 +491,15 @@ TfLiteStatus Eval(TfLiteContext* context, TfLiteNode* node) { if (output->type == kTfLiteUInt8) { const uint8_t kMappingTable[1 << kMapBits] = {127, 255, 0, 127}; TypedEval(kMappingTable, params, output->data.uint8); + } else if (output->type == kTfLiteInt8) { + const int8_t kMappingTable[1 << kMapBits] = {0, 127, -128, 0}; + TypedEval(kMappingTable, params, output->data.int8); } else if (output->type == kTfLiteFloat32) { const float kMappingTable[1 << kMapBits] = {0.0, 1.0, -1.0, 0.0}; TypedEval(kMappingTable, params, output->data.f); } else { - context->ReportError(context, "Output type must be UInt8 or Float32."); + context->ReportError(context, + "Output type must be Int8, UInt8, or Float32."); return kTfLiteError; } diff --git a/third_party/tensorflow_models/src/research/seq_flow_lite/tflite_ops/tflite_qrnn_pooling.cc b/third_party/tensorflow_models/src/research/seq_flow_lite/tflite_ops/tflite_qrnn_pooling.cc index 6641234c53d4b..205e8a619e5c2 100644 --- a/third_party/tensorflow_models/src/research/seq_flow_lite/tflite_ops/tflite_qrnn_pooling.cc +++ b/third_party/tensorflow_models/src/research/seq_flow_lite/tflite_ops/tflite_qrnn_pooling.cc @@ -22,7 +22,9 @@ namespace custom { namespace { -const uint8_t kPoolingForward = 255; +const uint8_t kPoolingUInt8Forward = 255; +const int8_t kPoolingInt8Forward = 127; +const float kPoolingFloatForward = 1.0; TfLiteStatus Prepare(TfLiteContext* context, TfLiteNode* node) { TF_LITE_ENSURE_EQ(context, node->inputs->size, 3); @@ -34,9 +36,11 @@ TfLiteStatus Prepare(TfLiteContext* context, TfLiteNode* node) { TfLiteTensor* constant = &context->tensors[node->inputs->data[1]]; TfLiteTensor* direction = &context->tensors[node->inputs->data[2]]; - TF_LITE_ENSURE_EQ(context, multiplier->type, kTfLiteUInt8); - TF_LITE_ENSURE_EQ(context, constant->type, kTfLiteUInt8); - TF_LITE_ENSURE_EQ(context, direction->type, kTfLiteUInt8); + TF_LITE_ENSURE(context, (constant->type == kTfLiteUInt8) || + (constant->type == kTfLiteInt8) || + (constant->type == kTfLiteFloat32)); + TF_LITE_ENSURE_TYPES_EQ(context, multiplier->type, constant->type); + TF_LITE_ENSURE_TYPES_EQ(context, direction->type, constant->type); TF_LITE_ENSURE_EQ(context, multiplier->dims->size, 3); TF_LITE_ENSURE_EQ(context, multiplier->dims->data[0], 1); @@ -71,9 +75,13 @@ TfLiteStatus Prepare(TfLiteContext* context, TfLiteNode* node) { return kTfLiteOk; } -TfLiteStatus QRNNPooling(TfLiteContext* context, TfLiteTensor* multiplier, - TfLiteTensor* constant, TfLiteTensor* outputs, - TfLiteTensor* final_state, bool forward) { +template +TfLiteStatus QRNNPooling(TfLiteContext* context, + TfLiteTensor* multiplier, + TfLiteTensor* constant, + TfLiteTensor* outputs, + TfLiteTensor* final_state, + bool forward) { const int time_steps = multiplier->dims->data[1]; const int state_size = multiplier->dims->data[2]; @@ -82,28 +90,67 @@ TfLiteStatus QRNNPooling(TfLiteContext* context, TfLiteTensor* multiplier, const int32_t out_zero_point = outputs ? outputs->params.zero_point : 0; const float out_inverse_scale = outputs ? 1.0f / outputs->params.scale : 1.0f; - uint8_t* out_ptr = outputs ? outputs->data.uint8 : nullptr; + T* outputs_ptr = outputs ? tflite::GetTensorData(outputs) : nullptr; for (int i = 0; i < time_steps; ++i) { for (int j = 0; j < state_size; ++j) { const int time_index = forward ? i : time_steps - (i + 1); const int index = time_index * state_size + j; - float multiplier_value = PodDequantize(*multiplier, index); - float constant_vale = PodDequantize(*constant, index); - state[j] = state[j] * multiplier_value + constant_vale; + float multiplier_value = PodDequantize(*multiplier, index); + float constant_value = PodDequantize(*constant, index); + state[j] = state[j] * multiplier_value + constant_value; if (outputs) { - out_ptr[index] = - PodQuantize(state[j], out_zero_point, out_inverse_scale); + outputs_ptr[index] = + PodQuantize(state[j], out_zero_point, out_inverse_scale); } } } if (final_state) { - uint8_t* final_state_ptr = final_state->data.uint8; + T* final_state_ptr = tflite::GetTensorData(final_state); const int32_t zero_point = final_state->params.zero_point; const float inverse_scale = 1.0f / final_state->params.scale; for (int j = 0; j < state_size; ++j) { - final_state_ptr[j] = - PodQuantize(state[j], zero_point, inverse_scale); + final_state_ptr[j] = PodQuantize(state[j], zero_point, inverse_scale); + } + } + + return kTfLiteOk; +} + +template <> +TfLiteStatus QRNNPooling(TfLiteContext* context, + TfLiteTensor* multiplier, + TfLiteTensor* constant, + TfLiteTensor* outputs, + TfLiteTensor* final_state, + bool forward) { + const int time_steps = multiplier->dims->data[1]; + const int state_size = multiplier->dims->data[2]; + + auto state = std::make_unique(state_size); + memset(state.get(), 0, sizeof(float) * state_size); + + float* multiplier_ptr = tflite::GetTensorData(multiplier); + float* constant_ptr = tflite::GetTensorData(constant); + float* outputs_ptr = + outputs ? tflite::GetTensorData(outputs) : nullptr; + for (int i = 0; i < time_steps; ++i) { + for (int j = 0; j < state_size; ++j) { + const int time_index = forward ? i : time_steps - (i + 1); + const int index = time_index * state_size + j; + float multiplier_value = multiplier_ptr[index]; + float constant_value = constant_ptr[index]; + state[j] = state[j] * multiplier_value + constant_value; + if (outputs) { + outputs_ptr[index] = state[j]; + } + } + } + + if (final_state) { + float* final_state_ptr = tflite::GetTensorData(final_state); + for (int j = 0; j < state_size; ++j) { + final_state_ptr[j] = state[j]; } } @@ -126,10 +173,24 @@ TfLiteStatus Invoke(TfLiteContext* context, TfLiteNode* node) { // When pooling forward the direction parameter is expected to be // kPoolingForward. - return QRNNPooling(context, multiplier, constant, outputs, final_state, - (direction->data.uint8[0] == kPoolingForward)); + switch (multiplier->type) { + case kTfLiteUInt8: + return QRNNPooling( + context, multiplier, constant, outputs, final_state, + (tflite::GetTensorData(direction)[0] == + kPoolingUInt8Forward)); + case kTfLiteInt8: + return QRNNPooling( + context, multiplier, constant, outputs, final_state, + (tflite::GetTensorData(direction)[0] == kPoolingInt8Forward)); + case kTfLiteFloat32: + return QRNNPooling( + context, multiplier, constant, outputs, final_state, + (tflite::GetTensorData(direction)[0] == kPoolingFloatForward)); + default: + return kTfLiteError; + } } - } // namespace const char kPoolingOp[] = "PoolingOp"; -- 2.45.2.627.g7a2c4fd464-goog