diff --git a/tensorflow/lite/micro/kernels/decode.cc b/tensorflow/lite/micro/kernels/decode.cc index 941f2b3b05c..b8ac3f006f9 100644 --- a/tensorflow/lite/micro/kernels/decode.cc +++ b/tensorflow/lite/micro/kernels/decode.cc @@ -13,6 +13,8 @@ See the License for the specific language governing permissions and limitations under the License. ==============================================================================*/ +#include + #include "tensorflow/lite/micro/c/common.h" #include "tensorflow/lite/micro/kernels/decode_state.h" #include "tensorflow/lite/micro/kernels/internal/compatibility.h" @@ -49,6 +51,27 @@ TfLiteStatus SetOutputTensorData(TfLiteContext* context, const TfLiteNode* node, return kTfLiteOk; } +DecodeState* GetDecodeStateFromCustomRegistration(const TfLiteContext* context, + uint8_t type) { + const MicroContext* mc = GetMicroContext(context); + const MicroContext::CustomDecodeRegistration* registrations; + size_t registrations_count; + std::tie(registrations, registrations_count) = + mc->GetCustomDecodeRegistrations(); + if (registrations == nullptr) { + return nullptr; + } + + for (size_t i = 0; i < registrations_count; i++) { + auto& reg = registrations[i]; + if (reg.type == type && reg.create_state != nullptr) { + return reg.create_state(reg, *context, mc->GetAlternateProfiler()); + } + } + + return nullptr; +} + TfLiteStatus Prepare(TfLiteContext* context, TfLiteNode* node) { const size_t num_inputs = NumInputs(node); const size_t num_outputs = NumOutputs(node); @@ -113,21 +136,22 @@ TfLiteStatus Prepare(TfLiteContext* context, TfLiteNode* node) { dsp = DecodeState::CreateDecodeStateHuffman( context, micro_context->GetAlternateProfiler()); break; - case DecodeState::kDcmTypeCustom: - MicroPrintf("Custom decode type not yet supported"); - break; default: - MicroPrintf("unsupported decode type %u", - DecodeState::Type(*ancillary)); + uint32_t type = DecodeState::Type(*ancillary); + if (type >= DecodeState::kDcmTypeCustomFirst && + type <= DecodeState::kDcmTypeCustomLast) { + dsp = GetDecodeStateFromCustomRegistration(context, type); + } else { + MicroPrintf("unsupported decode type %u", type); + } break; } - status = SetOutputTensorData(context, node, i / 2, output); - if (status != kTfLiteOk) { - break; - } - if (dsp != nullptr) { + status = SetOutputTensorData(context, node, i / 2, output); + if (status != kTfLiteOk) { + break; + } status = dsp->Setup(*input, *ancillary, *output); if (status != kTfLiteOk) { break; diff --git a/tensorflow/lite/micro/kernels/decode_state.h b/tensorflow/lite/micro/kernels/decode_state.h index 48d706583da..d28ad742d45 100644 --- a/tensorflow/lite/micro/kernels/decode_state.h +++ b/tensorflow/lite/micro/kernels/decode_state.h @@ -26,6 +26,7 @@ limitations under the License. namespace tflite { namespace micro { + class DecodeState { public: DecodeState() = delete; @@ -71,7 +72,8 @@ class DecodeState { static constexpr uint8_t kDcmTypeLUT = 0; static constexpr uint8_t kDcmTypeHuffman = 1; static constexpr uint8_t kDcmTypePrune = 2; - static constexpr uint8_t kDcmTypeCustom = 127; + static constexpr uint8_t kDcmTypeCustomFirst = 128; + static constexpr uint8_t kDcmTypeCustomLast = 255; static constexpr size_t kDcmSizeInBytes = 16; diff --git a/tensorflow/lite/micro/kernels/decode_state_huffman.h b/tensorflow/lite/micro/kernels/decode_state_huffman.h index 18862ad49dd..2428286e4ad 100644 --- a/tensorflow/lite/micro/kernels/decode_state_huffman.h +++ b/tensorflow/lite/micro/kernels/decode_state_huffman.h @@ -24,6 +24,7 @@ limitations under the License. namespace tflite { namespace micro { + class DecodeStateHuffman : public DecodeState { public: DecodeStateHuffman() = delete; diff --git a/tensorflow/lite/micro/kernels/decode_state_huffman_test.cc b/tensorflow/lite/micro/kernels/decode_state_huffman_test.cc index c32f346ec44..5eeffb582e5 100644 --- a/tensorflow/lite/micro/kernels/decode_state_huffman_test.cc +++ b/tensorflow/lite/micro/kernels/decode_state_huffman_test.cc @@ -269,7 +269,7 @@ TEST(DecodeStateHuffmanTest, DecodeHuffmanTable16BitsInt16Fail) { tflite::testing::TestDecode( encodes, ancillaries, outputs, expected, tflite::Register_DECODE(), - nullptr, kTfLiteError); + nullptr, nullptr, kTfLiteError); } TEST(DecodeStateHuffmanTest, DecodeHuffmanTable32BitsInt8) { diff --git a/tensorflow/lite/micro/kernels/decode_state_lut.h b/tensorflow/lite/micro/kernels/decode_state_lut.h index 016a0bf607f..107ab5d6819 100644 --- a/tensorflow/lite/micro/kernels/decode_state_lut.h +++ b/tensorflow/lite/micro/kernels/decode_state_lut.h @@ -23,6 +23,7 @@ limitations under the License. namespace tflite { namespace micro { + class DecodeStateLut : public DecodeState { public: DecodeStateLut() = delete; diff --git a/tensorflow/lite/micro/kernels/decode_state_prune.h b/tensorflow/lite/micro/kernels/decode_state_prune.h index 3012a5e3d7b..a717e751ec4 100644 --- a/tensorflow/lite/micro/kernels/decode_state_prune.h +++ b/tensorflow/lite/micro/kernels/decode_state_prune.h @@ -23,6 +23,7 @@ limitations under the License. namespace tflite { namespace micro { + class DecodeStatePrune : public DecodeState { public: DecodeStatePrune() = delete; diff --git a/tensorflow/lite/micro/kernels/decode_state_prune_test.cc b/tensorflow/lite/micro/kernels/decode_state_prune_test.cc index bbb71f5cf43..0065d56e60f 100644 --- a/tensorflow/lite/micro/kernels/decode_state_prune_test.cc +++ b/tensorflow/lite/micro/kernels/decode_state_prune_test.cc @@ -573,7 +573,7 @@ TEST(DecodeStatePruneTest, DecodePruneQuantizedInvalidZeroPointInt16) { tflite::testing::TestDecode( kEncodes, kAncillaries, kOutputs, kExpected, tflite::Register_DECODE(), - nullptr, kTfLiteError); + nullptr, nullptr, kTfLiteError); } TF_LITE_MICRO_TESTS_MAIN diff --git a/tensorflow/lite/micro/kernels/decode_test.cc b/tensorflow/lite/micro/kernels/decode_test.cc index fb3851bcee5..a1e897a723c 100644 --- a/tensorflow/lite/micro/kernels/decode_test.cc +++ b/tensorflow/lite/micro/kernels/decode_test.cc @@ -66,6 +66,76 @@ constexpr int kEncodedShapeLUT[] = {1, sizeof(kEncodedLUT)}; constexpr int8_t kExpectLUT0[] = {1, 2, 3, 4, 4, 3, 2, 1}; constexpr int16_t kExpectLUT1[] = {5, 6, 7, 8, 8, 7, 6, 5}; +// +// Custom DECODE test data +// +constexpr int kDecodeTypeCustom = 200; + +constexpr int8_t kAncillaryDataCustom[] = {0x42}; + +constexpr uint8_t kDcmCustom[tflite::DecodeState::kDcmSizeInBytes] = { + kDecodeTypeCustom, // type: custom + 1, // DCM version: 1 +}; + +// Align the tensor data the same as a Buffer in the TfLite schema +alignas(16) const uint8_t kEncodedCustom[] = {0x42, 0x43, 0x40, 0x46, + 0x4A, 0x52, 0x62, 0x02}; + +// Tensor shapes as TfLiteIntArray +constexpr int kOutputShapeCustom[] = {1, 8}; +constexpr int kEncodedShapeCustom[] = {1, sizeof(kEncodedCustom)}; + +constexpr int8_t kExpectCustom[] = {0x00, 0x01, 0x02, 0x04, + 0x08, 0x10, 0x20, 0x40}; + +class DecodeStateCustom : public tflite::DecodeState { + public: + DecodeStateCustom() = delete; + + DecodeStateCustom(const TfLiteContext* context, + tflite::MicroProfilerInterface* profiler) + : DecodeState(context, profiler) {} + + TfLiteStatus Setup(const TfLiteTensor& input, const TfLiteTensor& ancillary, + const TfLiteTensor& output) override { + return kTfLiteOk; + } + + TfLiteStatus Decode(const TfLiteEvalTensor& input, + const TfLiteEvalTensor& ancillary, + const TfLiteEvalTensor& output) override { + const uint8_t* inp = tflite::micro::GetTensorData(&input); + TF_LITE_ENSURE(const_cast(context_), inp != nullptr); + uint8_t* outp = tflite::micro::GetTensorData( + const_cast(&output)); + TF_LITE_ENSURE(const_cast(context_), outp != nullptr); + const uint8_t* vp = tflite::micro::GetTensorData(&ancillary); + TF_LITE_ENSURE(const_cast(context_), vp != nullptr); + vp += kDcmSizeInBytes; + + // simple XOR de-obfuscation + std::transform(inp, inp + input.dims->data[0], outp, + [vp](uint8_t i) { return i ^ *vp; }); + + return kTfLiteOk; + } + + static DecodeState* CreateDecodeStateCustom( + const tflite::MicroContext::CustomDecodeRegistration& registration, + const TfLiteContext& context, tflite::MicroProfilerInterface* profiler) { + alignas(DecodeStateCustom) static uint8_t buffer[sizeof(DecodeStateCustom)]; + DecodeState* instance = new (buffer) DecodeStateCustom(&context, profiler); + return instance; + } + + protected: + ~DecodeStateCustom() override = default; + + private: + TF_LITE_REMOVE_VIRTUAL_DELETE +}; + } // namespace using tflite::testing::AncillaryData; @@ -244,4 +314,112 @@ TEST(DecodeTest, DecodeWithAltDecompressionMemory) { encodes, ancillaries, outputs, expected, tflite::Register_DECODE(), &amr); } +TEST(DecodeTest, DecodeWithCustomRegistration) { + // Align the tensor data the same as a Buffer in the TfLite schema + alignas(16) int8_t output_data[std::size(kExpectCustom)] = {}; + alignas(16) const AncillaryData + kAncillaryData = {{kDcmCustom}, {kAncillaryDataCustom}}; + + constexpr int kAncillaryShapeCustom[] = {1, sizeof(kAncillaryData)}; + + const TfLiteIntArray* const encoded_dims = + tflite::testing::IntArrayFromInts(kEncodedShapeCustom); + static const TensorInDatum tid_encode = { + kEncodedCustom, + *encoded_dims, + }; + static constexpr std::initializer_list encodes = { + &tid_encode, + }; + + const TfLiteIntArray* const ancillary_dims = + tflite::testing::IntArrayFromInts(kAncillaryShapeCustom); + static const TensorInDatum tid_ancillary = { + &kAncillaryData, + *ancillary_dims, + }; + static constexpr std::initializer_list ancillaries = { + &tid_ancillary}; + + const TfLiteIntArray* const output_dims = + tflite::testing::IntArrayFromInts(kOutputShapeCustom); + constexpr int kOutputZeroPointsData[] = {0}; + const TfLiteIntArray* const kOutputZeroPoints = + tflite::testing::IntArrayFromInts(kOutputZeroPointsData); + const TfLiteFloatArray kOutputScales = {kOutputZeroPoints->size}; + static const TensorOutDatum tod = { + output_data, *output_dims, kTfLiteInt8, kOutputScales, *kOutputZeroPoints, + 0, {}, + }; + static constexpr std::initializer_list outputs = { + &tod}; + + const std::initializer_list expected = {kExpectCustom}; + + const std::initializer_list + cdr = {{ + DecodeStateCustom::CreateDecodeStateCustom, + kDecodeTypeCustom, + }}; + + tflite::testing::TestDecode( + encodes, ancillaries, outputs, expected, tflite::Register_DECODE(), + nullptr, &cdr); +} + +TEST(DecodeTest, DecodeWithCustomMismatchedRegistration) { + // Align the tensor data the same as a Buffer in the TfLite schema + alignas(16) int8_t output_data[std::size(kExpectCustom)] = {}; + alignas(16) const AncillaryData + kAncillaryData = {{kDcmCustom}, {kAncillaryDataCustom}}; + + constexpr int kAncillaryShapeCustom[] = {1, sizeof(kAncillaryData)}; + + const TfLiteIntArray* const encoded_dims = + tflite::testing::IntArrayFromInts(kEncodedShapeCustom); + static const TensorInDatum tid_encode = { + kEncodedCustom, + *encoded_dims, + }; + static constexpr std::initializer_list encodes = { + &tid_encode, + }; + + const TfLiteIntArray* const ancillary_dims = + tflite::testing::IntArrayFromInts(kAncillaryShapeCustom); + static const TensorInDatum tid_ancillary = { + &kAncillaryData, + *ancillary_dims, + }; + static constexpr std::initializer_list ancillaries = { + &tid_ancillary}; + + const TfLiteIntArray* const output_dims = + tflite::testing::IntArrayFromInts(kOutputShapeCustom); + constexpr int kOutputZeroPointsData[] = {0}; + const TfLiteIntArray* const kOutputZeroPoints = + tflite::testing::IntArrayFromInts(kOutputZeroPointsData); + const TfLiteFloatArray kOutputScales = {kOutputZeroPoints->size}; + static const TensorOutDatum tod = { + output_data, *output_dims, kTfLiteInt8, kOutputScales, *kOutputZeroPoints, + 0, {}, + }; + static constexpr std::initializer_list outputs = { + &tod}; + + const std::initializer_list expected = {kExpectCustom}; + + const std::initializer_list + cdr = {{ + DecodeStateCustom::CreateDecodeStateCustom, + kDecodeTypeCustom + 1, + }}; + + tflite::testing::TestDecode( + encodes, ancillaries, outputs, expected, tflite::Register_DECODE(), + nullptr, &cdr, kTfLiteError); +} + TF_LITE_MICRO_TESTS_MAIN diff --git a/tensorflow/lite/micro/kernels/decode_test_helpers.h b/tensorflow/lite/micro/kernels/decode_test_helpers.h index 4f76e6f0c11..57752d3f657 100644 --- a/tensorflow/lite/micro/kernels/decode_test_helpers.h +++ b/tensorflow/lite/micro/kernels/decode_test_helpers.h @@ -81,6 +81,8 @@ void ExecuteDecodeTest( const std::initializer_list& expected, TfLiteStatus expected_status, const std::initializer_list* amr = + nullptr, + const std::initializer_list* cdr = nullptr) { int kInputArrayData[kNumInputs + 1] = {kNumInputs}; for (size_t i = 0; i < kNumInputs; i++) { @@ -102,6 +104,11 @@ void ExecuteDecodeTest( amr->size()); } + if (cdr != nullptr) { + runner.GetFakeMicroContext()->SetCustomDecodeRegistrations(cdr->begin(), + cdr->size()); + } + TfLiteStatus status = runner.InitAndPrepare(); if (status == kTfLiteOk) { status = runner.Invoke(); @@ -152,6 +159,8 @@ void TestDecode( const TFLMRegistration& registration, const std::initializer_list* amr = nullptr, + const std::initializer_list* cdr = + nullptr, const TfLiteStatus expected_status = kTfLiteOk) { TfLiteTensor tensors[kNumInputs + kNumOutputs] = {}; @@ -185,7 +194,7 @@ void TestDecode( } ExecuteDecodeTest(tensors, registration, expected, - expected_status, amr); + expected_status, amr, cdr); } } // namespace testing diff --git a/tensorflow/lite/micro/micro_context.cc b/tensorflow/lite/micro/micro_context.cc index b5e12953aec..a25d62bc1ed 100644 --- a/tensorflow/lite/micro/micro_context.cc +++ b/tensorflow/lite/micro/micro_context.cc @@ -174,4 +174,14 @@ void MicroContext::ResetDecompressionMemoryAllocations() { std::fill_n(decompress_regions_allocations_, decompress_regions_size_, 0); } +TfLiteStatus MicroContext::SetCustomDecodeRegistrations( + const CustomDecodeRegistration* registrations, size_t count) { + if (custom_decode_registrations_ != nullptr) { + return kTfLiteError; + } + custom_decode_registrations_ = registrations; + custom_decode_registrations_size_ = count; + return kTfLiteOk; +} + } // namespace tflite diff --git a/tensorflow/lite/micro/micro_context.h b/tensorflow/lite/micro/micro_context.h index a9b6254bb3e..39e2330c045 100644 --- a/tensorflow/lite/micro/micro_context.h +++ b/tensorflow/lite/micro/micro_context.h @@ -17,7 +17,7 @@ limitations under the License. #define TENSORFLOW_LITE_MICRO_MICRO_CONTEXT_H_ #include -#include +#include #include "tensorflow/lite/micro/c/common.h" #include "tensorflow/lite/micro/micro_graph.h" @@ -33,6 +33,10 @@ namespace tflite { // TODO(b/149795762): kTfLiteAbort cannot be part of the tflite TfLiteStatus. const TfLiteStatus kTfLiteAbort = static_cast(15); +namespace micro { +class DecodeState; // can't use decode_state.h due to circular include +} // namespace micro + // MicroContext is eventually going to become the API between TFLM and the // kernels, replacing all the functions in TfLiteContext. The end state is code // kernels to have code like: @@ -136,7 +140,7 @@ class MicroContext { }; // Set the alternate decompression memory regions. - // Can only be called during the MicroInterpreter kInit state. + // Can only be called during the kInit state. virtual TfLiteStatus SetDecompressionMemory( const AlternateMemoryRegion* regions, size_t count); @@ -169,12 +173,34 @@ class MicroContext { return nullptr; } + struct CustomDecodeRegistration { + tflite::micro::DecodeState* (*create_state)(const CustomDecodeRegistration&, + const TfLiteContext&, + MicroProfilerInterface*); + uint8_t type; // custom decode type + }; + + // Set the DECODE operator custom registrations. + // Can only be called during the kInit state. + virtual TfLiteStatus SetCustomDecodeRegistrations( + const CustomDecodeRegistration* registrations, size_t count); + + // Get the custom DECODE operator registrations. + std::pair + GetCustomDecodeRegistrations() const { + return std::make_pair(custom_decode_registrations_, + custom_decode_registrations_size_); + } + private: const AlternateMemoryRegion* decompress_regions_ = nullptr; size_t decompress_regions_size_ = 0; // array of size_t elements with length equal to decompress_regions_size_ size_t* decompress_regions_allocations_ = nullptr; + const CustomDecodeRegistration* custom_decode_registrations_ = nullptr; + size_t custom_decode_registrations_size_ = 0; + TF_LITE_REMOVE_VIRTUAL_DELETE }; diff --git a/tensorflow/lite/micro/micro_interpreter.cc b/tensorflow/lite/micro/micro_interpreter.cc index d37883a2c58..73921297b17 100644 --- a/tensorflow/lite/micro/micro_interpreter.cc +++ b/tensorflow/lite/micro/micro_interpreter.cc @@ -349,4 +349,9 @@ TfLiteStatus MicroInterpreter::SetDecompressionMemory( return micro_context_.SetDecompressionMemory(regions, count); } +TfLiteStatus MicroInterpreter::SetCustomDecodeRegistrations( + const MicroContext::CustomDecodeRegistration* registrations, size_t count) { + return micro_context_.SetCustomDecodeRegistrations(registrations, count); +} + } // namespace tflite diff --git a/tensorflow/lite/micro/micro_interpreter.h b/tensorflow/lite/micro/micro_interpreter.h index e08e073e55c..58b63c67f9b 100644 --- a/tensorflow/lite/micro/micro_interpreter.h +++ b/tensorflow/lite/micro/micro_interpreter.h @@ -174,6 +174,18 @@ class MicroInterpreter { TfLiteStatus SetDecompressionMemory( const MicroContext::AlternateMemoryRegion* regions, size_t count); + // Set the DECODE operator custom registrations. + // Can only be called during the MicroInterpreter kInit state (i.e. must + // be called before MicroInterpreter::AllocateTensors). + // The registrations pointer argument is the start of a + // MicroContext::CustomDecodeRegistration array where the length of the array + // is given by the count argument. The lifetime of the + // MicroContext::CustomDecodeRegistration array must be at least that of the + // MicroInterpreter. + TfLiteStatus SetCustomDecodeRegistrations( + const MicroContext::CustomDecodeRegistration* registrations, + size_t count); + protected: const MicroAllocator& allocator() const { return allocator_; } const TfLiteContext& context() const { return context_; } diff --git a/tensorflow/lite/micro/micro_interpreter_context.cc b/tensorflow/lite/micro/micro_interpreter_context.cc index 0d61bed24a2..0ff4f03d553 100644 --- a/tensorflow/lite/micro/micro_interpreter_context.cc +++ b/tensorflow/lite/micro/micro_interpreter_context.cc @@ -247,4 +247,12 @@ MicroProfilerInterface* MicroInterpreterContext::GetAlternateProfiler() const { return alt_profiler_; } +TfLiteStatus MicroInterpreterContext::SetCustomDecodeRegistrations( + const CustomDecodeRegistration* registrations, size_t count) { + if (state_ != InterpreterState::kInit) { + return kTfLiteError; + } + return MicroContext::SetCustomDecodeRegistrations(registrations, count); +} + } // namespace tflite diff --git a/tensorflow/lite/micro/micro_interpreter_context.h b/tensorflow/lite/micro/micro_interpreter_context.h index 8293e4dcd72..806b881e3dd 100644 --- a/tensorflow/lite/micro/micro_interpreter_context.h +++ b/tensorflow/lite/micro/micro_interpreter_context.h @@ -159,6 +159,11 @@ class MicroInterpreterContext : public MicroContext { // decompression subsystem. MicroProfilerInterface* GetAlternateProfiler() const override; + // Set the DECODE operator custom registrations. + // Can only be called during the kInit state. + TfLiteStatus SetCustomDecodeRegistrations( + const CustomDecodeRegistration* registrations, size_t count) override; + private: MicroAllocator& allocator_; MicroInterpreterGraph& graph_; diff --git a/tensorflow/lite/micro/micro_interpreter_context_test.cc b/tensorflow/lite/micro/micro_interpreter_context_test.cc index 84e6f8077a9..c65b12a86cc 100644 --- a/tensorflow/lite/micro/micro_interpreter_context_test.cc +++ b/tensorflow/lite/micro/micro_interpreter_context_test.cc @@ -15,7 +15,9 @@ limitations under the License. #include "tensorflow/lite/micro/micro_interpreter_context.h" #include +#include +#include "tensorflow/lite/micro/kernels/decode_state.h" #include "tensorflow/lite/micro/micro_allocator.h" #include "tensorflow/lite/micro/micro_arena_constants.h" #include "tensorflow/lite/micro/micro_interpreter_graph.h" @@ -310,4 +312,54 @@ TEST(MicroInterpreterContextTest, TestResetDecompressionMemory) { EXPECT_EQ(p, &g_alt_memory[0]); } +TEST(MicroInterpreterContextTest, TestSetCustomDecode) { + tflite::MicroInterpreterContext micro_context = + tflite::CreateMicroInterpreterContext(); + + const std::initializer_list + cdr = { + { + nullptr, // the test won't instantiate tflite::DecodeState + tflite::DecodeState::kDcmTypeCustomFirst, + }, + }; + TfLiteStatus status; + + // Test that all of the MicroInterpreterContext fences are correct, by + // forcing the MicroInterpreterContext state. The SetCustomDecodeRegistrations + // method should only be allowed during the kInit state, and can only be + // set once. + + // fail during Prepare state + micro_context.SetInterpreterState( + tflite::MicroInterpreterContext::InterpreterState::kPrepare); + status = micro_context.SetCustomDecodeRegistrations(cdr.begin(), cdr.size()); + EXPECT_EQ(status, kTfLiteError); + + // fail during Invoke state + micro_context.SetInterpreterState( + tflite::MicroInterpreterContext::InterpreterState::kInvoke); + status = micro_context.SetCustomDecodeRegistrations(cdr.begin(), cdr.size()); + EXPECT_EQ(status, kTfLiteError); + + // succeed during Init state + micro_context.SetInterpreterState( + tflite::MicroInterpreterContext::InterpreterState::kInit); + status = micro_context.SetCustomDecodeRegistrations(cdr.begin(), cdr.size()); + EXPECT_EQ(status, kTfLiteOk); + + // fail on second Init state attempt + micro_context.SetInterpreterState( + tflite::MicroInterpreterContext::InterpreterState::kInit); + status = micro_context.SetCustomDecodeRegistrations(cdr.begin(), cdr.size()); + EXPECT_EQ(status, kTfLiteError); + + // check registered info. matches + const tflite::MicroContext::CustomDecodeRegistration* registration; + size_t count; + std::tie(registration, count) = micro_context.GetCustomDecodeRegistrations(); + EXPECT_EQ(registration, cdr.begin()); + EXPECT_EQ(count, 1U); +} + TF_LITE_MICRO_TESTS_MAIN diff --git a/tensorflow/lite/micro/tools/make/Makefile b/tensorflow/lite/micro/tools/make/Makefile index 171be5c2511..94435c3dc67 100644 --- a/tensorflow/lite/micro/tools/make/Makefile +++ b/tensorflow/lite/micro/tools/make/Makefile @@ -161,6 +161,7 @@ endif CC_WARNINGS := \ + -Werror \ -Wsign-compare \ -Wdouble-promotion \ -Wunused-variable \