Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions tensorflow/lite/micro/kernels/cast_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,7 @@ void TestCast(const InputT (&input)[N], const OutputT (&golden)[N]) {

template <typename IntT>
void TestFloatToInt(float pos_overflow, float neg_overflow) {
[[maybe_unused]]
constexpr bool is_signed = std::numeric_limits<IntT>::is_signed;
const float input[] = {100.f,
1.0f,
Expand Down
44 changes: 34 additions & 10 deletions tensorflow/lite/micro/kernels/decode.cc
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,8 @@ See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/

#include <utility>

#include "tensorflow/lite/c/common.h"
#include "tensorflow/lite/kernels/internal/compatibility.h"
#include "tensorflow/lite/kernels/kernel_util.h"
Expand Down Expand Up @@ -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);
Expand Down Expand Up @@ -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;
Expand Down
3 changes: 2 additions & 1 deletion tensorflow/lite/micro/kernels/decode_state.h
Original file line number Diff line number Diff line change
Expand Up @@ -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;

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -269,7 +269,7 @@ TEST(DecodeStateHuffmanTest, DecodeHuffmanTable16BitsInt16Fail) {
tflite::testing::TestDecode<encodes.size() + ancillaries.size(),
outputs.size()>(
encodes, ancillaries, outputs, expected, tflite::Register_DECODE(),
nullptr, kTfLiteError);
nullptr, nullptr, kTfLiteError);
}

TEST(DecodeStateHuffmanTest, DecodeHuffmanTable32BitsInt8) {
Expand Down
41 changes: 28 additions & 13 deletions tensorflow/lite/micro/kernels/decode_state_lut.cc
Original file line number Diff line number Diff line change
Expand Up @@ -30,25 +30,40 @@ TfLiteStatus DecodeStateLut::Setup(const TfLiteTensor& input,
const TfLiteTensor& ancillary,
const TfLiteTensor& output) {
const uint8_t* const ancillary_data = GetTensorData<uint8_t>(&ancillary);
if (ancillary_data[kDcmVersionOffset] != 1) {
MicroPrintf("unsupported version %u", ancillary_data[kDcmVersionOffset]);
return kTfLiteError;
TF_LITE_ENSURE_MSG(const_cast<TfLiteContext*>(context_),
ancillary_data[kDcmVersionOffset] == 1,
"unsupported version %u",
ancillary_data[kDcmVersionOffset]);

// Resolve num_channels_, use_alternate_axis_.
// Axis parameter of the DCM is used to extract the channel count dimension
// from the output tensor shape.
// Constructor defaults num_channels_ to 1 (one).
// If the axis has all masked bits set, the output tensor has a single
// channel.
const uint8_t axis_mask =
ancillary_data[kDcmParamsOffset] & kDcmParamsAxisMask;
if (axis_mask != kDcmParamsAxisMask) {
const uint8_t axis = axis_mask >> kDcmParamsAxisMaskShift;
TFLITE_DCHECK(axis < NumDimensions(&output));
num_channels_ = SizeOfDimension(&output, axis);

if ((axis == NumDimensions(&output) - 1)) {
if (num_channels_ > 1) {
use_alternate_axis_ = true;
}
} else {
TF_LITE_ENSURE_MSG(const_cast<TfLiteContext*>(context_), axis == 0,
"unsupported channel axis %u", axis);
}
}

// resolve num_channels_ and use_alternate_axis_
if (output.quantization.type == kTfLiteAffineQuantization &&
output.quantization.params != nullptr) {
const TfLiteAffineQuantization* quantization =
reinterpret_cast<TfLiteAffineQuantization*>(output.quantization.params);
num_channels_ = quantization->scale->size;
if ((quantization->quantized_dimension == output.dims->size - 1) &&
num_channels_ > 1) {
use_alternate_axis_ = true;
} else if (quantization->quantized_dimension != 0) {
MicroPrintf("unsupported quantization axis %u",
quantization->quantized_dimension);
return kTfLiteError;
}
TFLITE_DCHECK(num_channels_ ==
static_cast<size_t>(quantization->scale->size));
}

compressed_indices_ = GetTensorData<uint8_t>(&input);
Expand Down
4 changes: 3 additions & 1 deletion tensorflow/lite/micro/kernels/decode_state_lut.h
Original file line number Diff line number Diff line change
Expand Up @@ -42,11 +42,13 @@ class DecodeStateLut : public DecodeState {
static constexpr size_t kMaxBitWidth = 7;
static constexpr size_t kMaxValueTableChannelStride = 128;

private:
public:
// LUT Decode Common Metadata constants
static constexpr size_t kDcmVersionOffset = 4;
static constexpr size_t kDcmParamsOffset = 5;
static constexpr uint8_t kDcmParamsBitWidthMask = 0x07;
static constexpr uint8_t kDcmParamsAxisMask = 0xF0;
static constexpr uint8_t kDcmParamsAxisMaskShift = 4;
static constexpr size_t kDcmValueTableStrideOffset = 6;

protected:
Expand Down
82 changes: 50 additions & 32 deletions tensorflow/lite/micro/kernels/decode_state_lut_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,8 @@ See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/

#include "tensorflow/lite/micro/kernels/decode_state_lut.h"

#include <algorithm>
#include <array>
#include <cstdint>
Expand Down Expand Up @@ -265,19 +267,60 @@ void TestDataSetup(TestingInfo<T>* info, TestingData<T>* data) {
info->decode_common_metadata = &data->decode_common_metadata;
}

uint8_t ComputeAxis(bool use_alt_axis, bool quantized,
const TfLiteIntArray& dims) {
const uint8_t axis = use_alt_axis ? dims.size - 1 : 0;
const int channel_count = dims.data[axis];
if (quantized || channel_count > 1) {
return axis << tflite::DecodeStateLut::kDcmParamsAxisMaskShift;
}

return tflite::DecodeStateLut::kDcmParamsAxisMask;
}

template <typename T>
void TestDecompression(TestingInfo<T>* info) {
GenerateData(*info);

const int first_dim = info->use_alt_axis
? info->total_elements / info->channel_count
: info->channel_count;
const int last_dim = info->use_alt_axis
? info->channel_count
: info->total_elements / info->channel_count;
const int output_dims_array[] = {2, first_dim, last_dim};
const TfLiteIntArray* const output_dims =
tflite::testing::IntArrayFromInts(output_dims_array);
// The actual zero-point and scale data are never used,
// so only supply the sizes.
const TfLiteIntArray kOutputZeroPoints = {
std::is_same<T, bool>::value || std::is_same<T, float>::value
? 0
: static_cast<int>(info->channel_count)};
const TfLiteFloatArray kOutputScales = {kOutputZeroPoints.size};
const TensorOutDatum tod = {
info->output,
*output_dims,
typeToTfLiteType<T>(),
kOutputScales,
kOutputZeroPoints,
info->use_alt_axis ? output_dims->size - 1 : 0,
{},
};
const std::initializer_list<const TensorOutDatum*> outputs = {&tod};

const size_t stride = info->total_value_table_elements / info->channel_count;
const uint8_t axis_param = ComputeAxis(
info->use_alt_axis, kOutputZeroPoints.size != 0, *output_dims);
const uint8_t lut_params = axis_param | static_cast<uint8_t>(info->bit_width);
*info->decode_common_metadata = {
tflite::DecodeState::kDcmTypeLUT, // type: LUT
1, // DCM version: 1
0, // reserved
0, // reserved
1, // LUT version: 1
static_cast<uint8_t>(info->bit_width), // Parameters: bit-width
static_cast<uint8_t>(stride), // value table channel stride
tflite::DecodeState::kDcmTypeLUT, // type: LUT
1, // DCM version: 1
0, // reserved
0, // reserved
1, // LUT version: 1
lut_params, // Parameters: axis, bit-width
static_cast<uint8_t>(stride), // value table channel stride
};

const int encoded_dims_array[] = {
Expand All @@ -304,31 +347,6 @@ void TestDecompression(TestingInfo<T>* info) {
const std::initializer_list<const TensorInDatum*> ancillaries = {
&tid_ancillary};

const int first_dim = info->use_alt_axis
? info->total_elements / info->channel_count
: info->channel_count;
const int last_dim = info->use_alt_axis
? info->channel_count
: info->total_elements / info->channel_count;
const int output_dims_array[] = {2, first_dim, last_dim};
const TfLiteIntArray* const output_dims =
tflite::testing::IntArrayFromInts(output_dims_array);
// The actual zero-point and scale data are never used,
// so only supply the sizes.
const TfLiteIntArray kOutputZeroPoints = {
static_cast<int>(info->channel_count)};
const TfLiteFloatArray kOutputScales = {kOutputZeroPoints.size};
const TensorOutDatum tod = {
info->output,
*output_dims,
typeToTfLiteType<T>(),
kOutputScales,
kOutputZeroPoints,
info->use_alt_axis ? kOutputScales.size - 1 : 0,
{},
};
const std::initializer_list<const TensorOutDatum*> outputs = {&tod};

const std::initializer_list<const void*> expected = {info->goldens};

std::fill_n(info->output, info->total_elements, static_cast<T>(~0ULL));
Expand Down
46 changes: 28 additions & 18 deletions tensorflow/lite/micro/kernels/decode_state_prune.cc
Original file line number Diff line number Diff line change
Expand Up @@ -31,26 +31,38 @@ TfLiteStatus DecodeStatePrune::Setup(const TfLiteTensor& input,
const TfLiteTensor& ancillary,
const TfLiteTensor& output) {
const uint8_t* const ancillary_data = GetTensorData<uint8_t>(&ancillary);
if (ancillary_data[kDcmVersionOffset] != 1) {
MicroPrintf("unsupported version %u", ancillary_data[kDcmVersionOffset]);
return kTfLiteError;
TF_LITE_ENSURE_MSG(const_cast<TfLiteContext*>(context_),
ancillary_data[kDcmVersionOffset] == 1,
"unsupported version %u",
ancillary_data[kDcmVersionOffset]);

// Resolve num_channels_, use_alternate_axis_, and zero points.
// Axis parameter of the DCM is used to extract the channel count dimension
// from the output tensor shape.
// Constructor defaults num_channels_ to 1 (one).
// If the axis has all masked bits set, the output tensor has a single
// channel.
const uint8_t axis_mask =
ancillary_data[kDcmParamsOffset] & kDcmParamsAxisMask;
if (axis_mask != kDcmParamsAxisMask) {
const uint8_t axis = axis_mask >> kDcmParamsAxisMaskShift;
TFLITE_DCHECK(axis < NumDimensions(&output));
num_channels_ = SizeOfDimension(&output, axis);

if ((axis == NumDimensions(&output) - 1)) {
if (num_channels_ > 1) {
use_alternate_axis_ = true;
}
} else {
TF_LITE_ENSURE_MSG(const_cast<TfLiteContext*>(context_), axis == 0,
"unsupported channel axis %u", axis);
}
}

// resolve num_channels_, use_alternate_axis_, and zero points
if (output.quantization.type == kTfLiteAffineQuantization &&
output.quantization.params != nullptr) {
const TfLiteAffineQuantization* quantization =
reinterpret_cast<TfLiteAffineQuantization*>(output.quantization.params);
num_channels_ = quantization->scale->size;
if ((quantization->quantized_dimension == output.dims->size - 1) &&
num_channels_ > 1) {
use_alternate_axis_ = true;
} else if (quantization->quantized_dimension != 0) {
MicroPrintf("unsupported quantization axis %u",
quantization->quantized_dimension);
return kTfLiteError;
}

TFLITE_DCHECK(num_channels_ ==
static_cast<size_t>(quantization->zero_point->size));
bool has_non_zero_zp =
Expand All @@ -71,10 +83,8 @@ TfLiteStatus DecodeStatePrune::Setup(const TfLiteTensor& input,
const size_t bufsize = num_channels_ * sizeof(*zero_points_);
zero_points_ = static_cast<decltype(zero_points_)>(
micro_context->AllocatePersistentBuffer(bufsize));
if (zero_points_ == nullptr) {
MicroPrintf("unable to allocate zero_points_");
return kTfLiteError;
}
TF_LITE_ENSURE(const_cast<TfLiteContext*>(context_),
zero_points_ != nullptr);
std::copy_n(quantization->zero_point->data, num_channels_, zero_points_);
} else {
single_zero_point_ = quantization->zero_point->data[0];
Expand Down
5 changes: 4 additions & 1 deletion tensorflow/lite/micro/kernels/decode_state_prune.h
Original file line number Diff line number Diff line change
Expand Up @@ -38,9 +38,12 @@ class DecodeStatePrune : public DecodeState {
const TfLiteEvalTensor& ancillary,
const TfLiteEvalTensor& output) override;

private:
public:
// Prune Decode Common Metadata constants
static constexpr size_t kDcmVersionOffset = 4;
static constexpr size_t kDcmParamsOffset = 5;
static constexpr uint8_t kDcmParamsAxisMask = 0xF0;
static constexpr uint8_t kDcmParamsAxisMaskShift = 4;

protected:
virtual ~DecodeStatePrune() = default;
Expand Down
Loading
Loading