From f81173ac6fa173a4006fda5b9e1f605904f0e9ef Mon Sep 17 00:00:00 2001 From: ddavis-2015 Date: Mon, 31 Aug 2026 03:18:17 -0700 Subject: [PATCH 1/7] Application custom DECODE operator support (part 1) @tensorflow/micro This is part 1 of 2 to add support for custom DECODE operators and their registration. Part 2 will add this support to the Python MicroInterpreter wrapper. Add unit tests for custom DECODE operator registration. Re-enable -Werror in the Makefile. Fix minor comment typos. bug=fixes #3215 --- tensorflow/lite/micro/kernels/decode.cc | 44 ++++- tensorflow/lite/micro/kernels/decode_state.h | 3 +- .../kernels/decode_state_huffman_test.cc | 2 +- .../micro/kernels/decode_state_prune_test.cc | 2 +- tensorflow/lite/micro/kernels/decode_test.cc | 183 ++++++++++++++++++ .../lite/micro/kernels/decode_test_helpers.h | 11 +- tensorflow/lite/micro/micro_context.cc | 10 + tensorflow/lite/micro/micro_context.h | 28 ++- tensorflow/lite/micro/micro_interpreter.cc | 5 + tensorflow/lite/micro/micro_interpreter.h | 12 ++ .../lite/micro/micro_interpreter_context.cc | 8 + .../lite/micro/micro_interpreter_context.h | 5 + .../micro/micro_interpreter_context_test.cc | 52 +++++ tensorflow/lite/micro/tools/make/Makefile | 1 + 14 files changed, 350 insertions(+), 16 deletions(-) diff --git a/tensorflow/lite/micro/kernels/decode.cc b/tensorflow/lite/micro/kernels/decode.cc index 9f4d34cff15..92391c6d24e 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/c/common.h" #include "tensorflow/lite/kernels/internal/compatibility.h" #include "tensorflow/lite/kernels/kernel_util.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 06f821dbc3c..9be36e32de3 100644 --- a/tensorflow/lite/micro/kernels/decode_state.h +++ b/tensorflow/lite/micro/kernels/decode_state.h @@ -72,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_test.cc b/tensorflow/lite/micro/kernels/decode_state_huffman_test.cc index 089414e2184..feebca24e48 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_prune_test.cc b/tensorflow/lite/micro/kernels/decode_state_prune_test.cc index 1144913fa64..d268bd6369e 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 376fa9c586d..ba24452c439 100644 --- a/tensorflow/lite/micro/kernels/decode_test.cc +++ b/tensorflow/lite/micro/kernels/decode_test.cc @@ -66,6 +66,77 @@ 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) {} + + virtual TfLiteStatus Setup(const TfLiteTensor& input, + const TfLiteTensor& ancillary, + const TfLiteTensor& output) override { + return kTfLiteOk; + } + + virtual 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: + virtual ~DecodeStateCustom() = default; + + private: + TF_LITE_REMOVE_VIRTUAL_DELETE +}; + } // namespace using tflite::testing::AncillaryData; @@ -244,4 +315,116 @@ 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 47c42972415..c59b65f44b9 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 fb03ac08351..acb21688170 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 0b8120abe1d..3bd337591a4 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/c/common.h" #include "tensorflow/lite/micro/micro_graph.h" @@ -33,6 +33,8 @@ namespace tflite { // TODO(b/149795762): kTfLiteAbort cannot be part of the tflite TfLiteStatus. const TfLiteStatus kTfLiteAbort = static_cast(15); +class DecodeState; // can't use decode_state.h due to circular include + // 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 +138,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 +171,34 @@ class MicroContext { return nullptr; } + struct CustomDecodeRegistration { + tflite::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 decompression registrations. + virtual const 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 1e96d50519a..111dabc1c46 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 f1de91962e2..cc788ef939a 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 regions 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 e454509362c..841b7e7e00e 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 72a5262f30c..8f18359d7af 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. + virtual 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 db053388c6d..2a0f09ef3c9 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 \ From 2966f2686f31bd36ec3b9da099c817b6ee6bebdb Mon Sep 17 00:00:00 2001 From: ddavis-2015 Date: Tue, 1 Sep 2026 16:06:38 -0700 Subject: [PATCH 2/7] Re-enable -Werror in Makefile @tensorflow/micro Put the -Werror flag back into the primary makefile, to prevent creaping warning accumulation. Fix various warning generators in codebase. bug=fixes #3689 --- tensorflow/lite/micro/kernels/decompress_common.cc | 4 ++-- tensorflow/lite/micro/kernels/xtensa/depthwise_conv.cc | 2 ++ tensorflow/lite/micro/tools/make/Makefile | 1 + 3 files changed, 5 insertions(+), 2 deletions(-) diff --git a/tensorflow/lite/micro/kernels/decompress_common.cc b/tensorflow/lite/micro/kernels/decompress_common.cc index 5c40af83bf3..602e87e64b8 100644 --- a/tensorflow/lite/micro/kernels/decompress_common.cc +++ b/tensorflow/lite/micro/kernels/decompress_common.cc @@ -326,7 +326,7 @@ void DecompressionState::DecompressToBufferWidthAny(T* buffer) { const T* value_table = static_cast(comp_data_.data.lut_data->value_table); for (size_t channel = 0; channel < num_channels_; channel++) { - size_t index; + size_t index = 0; switch (compressed_bit_width_) { case 1: index = GetNextTableIndexWidth1(current_offset); @@ -367,7 +367,7 @@ void DecompressionState::DecompressToBufferWidthAny(T* buffer) { size_t count = max_count; while (count-- > 0) { - size_t index; + size_t index = 0; switch (compressed_bit_width_) { case 1: index = GetNextTableIndexWidth1(current_offset); diff --git a/tensorflow/lite/micro/kernels/xtensa/depthwise_conv.cc b/tensorflow/lite/micro/kernels/xtensa/depthwise_conv.cc index e5d7f8f00fc..b59ff7d85a0 100644 --- a/tensorflow/lite/micro/kernels/xtensa/depthwise_conv.cc +++ b/tensorflow/lite/micro/kernels/xtensa/depthwise_conv.cc @@ -98,9 +98,11 @@ TfLiteStatus Eval(TfLiteContext* context, TfLiteNode* node) { MicroContext* micro_context = GetMicroContext(context); + [[maybe_unused]] const CompressionTensorData* filter_comp_td = micro_context->GetTensorCompressionData(node, kDepthwiseConvWeightsTensor); + [[maybe_unused]] const CompressionTensorData* bias_comp_td = micro_context->GetTensorCompressionData(node, kDepthwiseConvBiasTensor); diff --git a/tensorflow/lite/micro/tools/make/Makefile b/tensorflow/lite/micro/tools/make/Makefile index db053388c6d..2a0f09ef3c9 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 \ From be35234aef696ffdb1e4191353c944167bf61241 Mon Sep 17 00:00:00 2001 From: ddavis-2015 Date: Tue, 1 Sep 2026 17:32:53 -0700 Subject: [PATCH 3/7] suppress warning. --- tensorflow/lite/micro/kernels/cast_test.cc | 1 + 1 file changed, 1 insertion(+) diff --git a/tensorflow/lite/micro/kernels/cast_test.cc b/tensorflow/lite/micro/kernels/cast_test.cc index 3a2299095f5..d69f335ba67 100644 --- a/tensorflow/lite/micro/kernels/cast_test.cc +++ b/tensorflow/lite/micro/kernels/cast_test.cc @@ -59,6 +59,7 @@ void TestCast(const InputT (&input)[N], const OutputT (&golden)[N]) { template void TestFloatToInt(float pos_overflow, float neg_overflow) { + [[maybe_unused]] constexpr bool is_signed = std::numeric_limits::is_signed; const float input[] = {100.f, 1.0f, From 77d23958e5d5d550ce57d1aea24b8e91000b96c0 Mon Sep 17 00:00:00 2001 From: ddavis-2015 Date: Mon, 28 Sep 2026 16:35:31 -0700 Subject: [PATCH 4/7] Review fixes. Remove virtual specifiers where not needed. --- tensorflow/lite/micro/micro_context.h | 6 +++--- tensorflow/lite/micro/micro_interpreter.h | 2 +- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/tensorflow/lite/micro/micro_context.h b/tensorflow/lite/micro/micro_context.h index 3bd337591a4..92adaf93f8b 100644 --- a/tensorflow/lite/micro/micro_context.h +++ b/tensorflow/lite/micro/micro_context.h @@ -180,11 +180,11 @@ class MicroContext { // Set the DECODE operator custom registrations. // Can only be called during the kInit state. - virtual TfLiteStatus SetCustomDecodeRegistrations( + TfLiteStatus SetCustomDecodeRegistrations( const CustomDecodeRegistration* registrations, size_t count); - // Get the custom decompression registrations. - virtual const std::pair + // Get the custom DECODE operator registrations. + std::pair GetCustomDecodeRegistrations() const { return std::make_pair(custom_decode_registrations_, custom_decode_registrations_size_); diff --git a/tensorflow/lite/micro/micro_interpreter.h b/tensorflow/lite/micro/micro_interpreter.h index 775e16ceb05..72d74bf55ab 100644 --- a/tensorflow/lite/micro/micro_interpreter.h +++ b/tensorflow/lite/micro/micro_interpreter.h @@ -177,7 +177,7 @@ class MicroInterpreter { // Set the DECODE operator custom registrations. // Can only be called during the MicroInterpreter kInit state (i.e. must // be called before MicroInterpreter::AllocateTensors). - // The regions pointer argument is the start of a + // 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 From d98b8d9068146b13b878df4d1365efe737525115 Mon Sep 17 00:00:00 2001 From: ddavis-2015 Date: Mon, 28 Sep 2026 18:07:32 -0700 Subject: [PATCH 5/7] post merge fixup --- tensorflow/lite/micro/kernels/decode_state.h | 1 + tensorflow/lite/micro/micro_context.h | 10 ++++++---- 2 files changed, 7 insertions(+), 4 deletions(-) diff --git a/tensorflow/lite/micro/kernels/decode_state.h b/tensorflow/lite/micro/kernels/decode_state.h index 939fc4814d7..787f7d8cd00 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; diff --git a/tensorflow/lite/micro/micro_context.h b/tensorflow/lite/micro/micro_context.h index 92adaf93f8b..a10e38a30ad 100644 --- a/tensorflow/lite/micro/micro_context.h +++ b/tensorflow/lite/micro/micro_context.h @@ -33,7 +33,9 @@ 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 @@ -172,15 +174,15 @@ class MicroContext { } struct CustomDecodeRegistration { - tflite::DecodeState* (*create_state)(const CustomDecodeRegistration&, - const TfLiteContext&, - MicroProfilerInterface*); + 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. - TfLiteStatus SetCustomDecodeRegistrations( + virtual TfLiteStatus SetCustomDecodeRegistrations( const CustomDecodeRegistration* registrations, size_t count); // Get the custom DECODE operator registrations. From 35fe6f84105000c67083a75dd97dfbbdb5bfe2e1 Mon Sep 17 00:00:00 2001 From: ddavis-2015 Date: Mon, 28 Sep 2026 18:36:48 -0700 Subject: [PATCH 6/7] more post-merge fixups --- .../lite/micro/kernels/decode_state_huffman.h | 1 + .../lite/micro/kernels/decode_state_lut.h | 1 + .../lite/micro/kernels/decode_state_prune.h | 1 + tensorflow/lite/micro/kernels/decode_test.cc | 17 ++++++++--------- .../lite/micro/micro_interpreter_context.h | 2 +- 5 files changed, 12 insertions(+), 10 deletions(-) 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_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_test.cc b/tensorflow/lite/micro/kernels/decode_test.cc index af71540f75e..77e9e902cc2 100644 --- a/tensorflow/lite/micro/kernels/decode_test.cc +++ b/tensorflow/lite/micro/kernels/decode_test.cc @@ -97,15 +97,14 @@ class DecodeStateCustom : public tflite::DecodeState { tflite::MicroProfilerInterface* profiler) : DecodeState(context, profiler) {} - virtual TfLiteStatus Setup(const TfLiteTensor& input, - const TfLiteTensor& ancillary, - const TfLiteTensor& output) override { + TfLiteStatus Setup(const TfLiteTensor& input, const TfLiteTensor& ancillary, + const TfLiteTensor& output) override { return kTfLiteOk; } - virtual TfLiteStatus Decode(const TfLiteEvalTensor& input, - const TfLiteEvalTensor& ancillary, - const TfLiteEvalTensor& output) override { + 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( @@ -131,7 +130,7 @@ class DecodeStateCustom : public tflite::DecodeState { } protected: - virtual ~DecodeStateCustom() = default; + ~DecodeStateCustom() override = default; private: TF_LITE_REMOVE_VIRTUAL_DELETE @@ -363,7 +362,7 @@ TEST(DecodeTest, DecodeWithCustomRegistration) { DecodeStateCustom::CreateDecodeStateCustom, kDecodeTypeCustom, }, - }; + }; tflite::testing::TestDecode( @@ -419,7 +418,7 @@ TEST(DecodeTest, DecodeWithCustomMismatchedRegistration) { DecodeStateCustom::CreateDecodeStateCustom, kDecodeTypeCustom + 1, }, - }; + }; tflite::testing::TestDecode( diff --git a/tensorflow/lite/micro/micro_interpreter_context.h b/tensorflow/lite/micro/micro_interpreter_context.h index 8f18359d7af..abecccad430 100644 --- a/tensorflow/lite/micro/micro_interpreter_context.h +++ b/tensorflow/lite/micro/micro_interpreter_context.h @@ -161,7 +161,7 @@ class MicroInterpreterContext : public MicroContext { // Set the DECODE operator custom registrations. // Can only be called during the kInit state. - virtual TfLiteStatus SetCustomDecodeRegistrations( + TfLiteStatus SetCustomDecodeRegistrations( const CustomDecodeRegistration* registrations, size_t count) override; private: From 79af2114cbd0fdb86f42d3afd0159f0726dea96f Mon Sep 17 00:00:00 2001 From: ddavis-2015 Date: Mon, 28 Sep 2026 19:46:34 -0700 Subject: [PATCH 7/7] trailing commas in initializer list blocks are no longer allowed --- tensorflow/lite/micro/kernels/decode_test.cc | 20 ++++++++------------ 1 file changed, 8 insertions(+), 12 deletions(-) diff --git a/tensorflow/lite/micro/kernels/decode_test.cc b/tensorflow/lite/micro/kernels/decode_test.cc index 77e9e902cc2..b75596f1a20 100644 --- a/tensorflow/lite/micro/kernels/decode_test.cc +++ b/tensorflow/lite/micro/kernels/decode_test.cc @@ -357,12 +357,10 @@ TEST(DecodeTest, DecodeWithCustomRegistration) { const std::initializer_list expected = {kExpectCustom}; const std::initializer_list - cdr = { - { - DecodeStateCustom::CreateDecodeStateCustom, - kDecodeTypeCustom, - }, - }; + cdr = {{ + DecodeStateCustom::CreateDecodeStateCustom, + kDecodeTypeCustom, + }}; tflite::testing::TestDecode( @@ -413,12 +411,10 @@ TEST(DecodeTest, DecodeWithCustomMismatchedRegistration) { const std::initializer_list expected = {kExpectCustom}; const std::initializer_list - cdr = { - { - DecodeStateCustom::CreateDecodeStateCustom, - kDecodeTypeCustom + 1, - }, - }; + cdr = {{ + DecodeStateCustom::CreateDecodeStateCustom, + kDecodeTypeCustom + 1, + }}; tflite::testing::TestDecode(