diff --git a/litert/tensor/backends/xnnpack/BUILD b/litert/tensor/backends/xnnpack/BUILD index e91d996bd5b..c5dd17e4c6d 100644 --- a/litert/tensor/backends/xnnpack/BUILD +++ b/litert/tensor/backends/xnnpack/BUILD @@ -91,6 +91,7 @@ cc_library( hdrs = ["utils.h"], deps = [ "//:XNNPACK", + "//litert/tensor/internal:graph", "//litert/tensor/utils:macros", "@abseil-cpp//absl/status", "@abseil-cpp//absl/strings", diff --git a/litert/tensor/backends/xnnpack/arithmetic.cc b/litert/tensor/backends/xnnpack/arithmetic.cc index b5c91dd26b9..b7dddb05efa 100644 --- a/litert/tensor/backends/xnnpack/arithmetic.cc +++ b/litert/tensor/backends/xnnpack/arithmetic.cc @@ -59,6 +59,78 @@ absl::StatusOr DynamicallyQuantizeInput( return qd_id; } +absl::StatusOr DefineZeroTensor( + XnnpackBuildContext& ctx, const graph::TensorInformation& input_info, + size_t num_dims, const size_t* dims, absl::string_view op_name) { + xnn_datatype datatype = GetXnnpackType(input_info); + if (datatype == xnn_datatype_invalid) { + return absl::UnimplementedError( + absl::StrFormat("%s: unsupported input type %d", op_name, + static_cast(input_info.type))); + } + + size_t num_elements = 1; + for (size_t i = 0; i < num_dims; ++i) { + num_elements *= dims[i]; + } + + uint32_t zero_id = XNN_INVALID_VALUE_ID; + const void* data_ptr = nullptr; + + if (input_info.quantization) { + if (auto maybe_pcq = + input_info.quantization->As(); + maybe_pcq.ok()) { + const auto& pcq = maybe_pcq.value(); + if (pcq.scales.size() == 1) { + int32_t zero_point = pcq.zero_points.empty() ? 0 : pcq.zero_points[0]; + float scale = pcq.scales[0]; + + if (datatype == xnn_datatype_qint8) { + std::vector zeros(num_elements, + static_cast(zero_point)); + data_ptr = ctx.KeepAlive(std::move(zeros)); + } else { + return absl::UnimplementedError( + absl::StrFormat("%s: unsupported quantized type %d", op_name, + static_cast(datatype))); + } + + LRT_TENSOR_RETURN_IF_ERROR(xnn_define_quantized_tensor_value( + ctx.subgraph(), datatype, zero_point, scale, num_dims, dims, + data_ptr, + /*external_id=*/XNN_INVALID_VALUE_ID, /*flags=*/0, &zero_id)) + << "Could not define quantized zero tensor."; + } else { + return absl::UnimplementedError(absl::StrFormat( + "%s: per-channel quantized zero tensor not supported", op_name)); + } + } else { + return absl::UnimplementedError( + absl::StrFormat("%s: unsupported quantization type", op_name)); + } + } else { + if (datatype == xnn_datatype_fp32) { + std::vector zeros(num_elements, 0.0f); + data_ptr = ctx.KeepAlive(std::move(zeros)); + } else if (datatype == xnn_datatype_int32) { + std::vector zeros(num_elements, 0); + data_ptr = ctx.KeepAlive(std::move(zeros)); + } else { + return absl::UnimplementedError( + absl::StrFormat("%s: unsupported type %d for zero tensor", op_name, + static_cast(datatype))); + } + + LRT_TENSOR_RETURN_IF_ERROR(xnn_define_tensor_value( + ctx.subgraph(), datatype, num_dims, dims, data_ptr, + /*external_id=*/XNN_INVALID_VALUE_ID, /*flags=*/0, &zero_id)) + << "Could not define zero tensor."; + } + + return zero_id; +} + template absl::Status ValidateTensorType(const graph::Tensor& tensor, absl::string_view op_name) { @@ -1316,7 +1388,7 @@ absl::Status OpMixin::ToXnnpack( input_info.shape.size(), num_dims)); } - std::vector new_shape(num_dims); + std::vector zero_shape(num_dims); for (size_t i = 0; i < num_dims; ++i) { int mult = multiples_data[i]; int dim = input_info.shape[i]; @@ -1326,13 +1398,21 @@ absl::Status OpMixin::ToXnnpack( "Dimension %d has size %d but multiples[%d]=%d", op_name, i, dim, i, mult)); } - new_shape[i] = static_cast(dim * mult); + zero_shape[i] = static_cast(mult); } - LRT_TENSOR_RETURN_IF_ERROR(xnn_define_static_broadcast( - ctx.subgraph(), num_dims, new_shape.data(), input_id, output_id, + LRT_TENSOR_ASSIGN_OR_RETURN( + uint32_t zero_id, DefineZeroTensor(ctx, input_info, zero_shape.size(), + zero_shape.data(), op_name)); + + LRT_TENSOR_ASSIGN_OR_RETURN(auto params, + BuildBinaryParams(kActNone, op_name)); + + LRT_TENSOR_RETURN_IF_ERROR(xnn_define_binary( + ctx.subgraph(), xnn_binary_add, ¶ms, input_id, zero_id, output_id, /*flags=*/0)) << op_name; + return absl::OkStatus(); } @@ -1511,17 +1591,24 @@ OpMixin::ToXnnpack( static_cast(scale_h), static_cast(input_w), static_cast(scale_w), static_cast(channels)}; + xnn_datatype datatype = GetXnnpackType(input_info); + if (datatype == xnn_datatype_invalid) { + return absl::UnimplementedError( + absl::StrFormat("%s: unsupported input type %d", op_name, + static_cast(input_info.type))); + } + uint32_t reshape_id = XNN_INVALID_VALUE_ID; LRT_TENSOR_RETURN_IF_ERROR(xnn_define_tensor_value( - ctx.subgraph(), xnn_datatype_fp32, reshape_dims.size(), - reshape_dims.data(), /*data=*/nullptr, XNN_INVALID_VALUE_ID, + ctx.subgraph(), datatype, reshape_dims.size(), reshape_dims.data(), + /*data=*/nullptr, XNN_INVALID_VALUE_ID, /*flags=*/0, &reshape_id)) << op_name; uint32_t broadcast_id = XNN_INVALID_VALUE_ID; LRT_TENSOR_RETURN_IF_ERROR(xnn_define_tensor_value( - ctx.subgraph(), xnn_datatype_fp32, broadcast_dims.size(), - broadcast_dims.data(), /*data=*/nullptr, XNN_INVALID_VALUE_ID, + ctx.subgraph(), datatype, broadcast_dims.size(), broadcast_dims.data(), + /*data=*/nullptr, XNN_INVALID_VALUE_ID, /*flags=*/0, &broadcast_id)) << op_name; @@ -1530,9 +1617,20 @@ OpMixin::ToXnnpack( reshape_id, /*flags=*/0)) << op_name; - LRT_TENSOR_RETURN_IF_ERROR(xnn_define_static_broadcast( - ctx.subgraph(), broadcast_dims.size(), broadcast_dims.data(), reshape_id, - broadcast_id, /*flags=*/0)) + const std::array zero_dims = { + 1, 1, static_cast(scale_h), 1, static_cast(scale_w), 1}; + + LRT_TENSOR_ASSIGN_OR_RETURN( + uint32_t zero_id, DefineZeroTensor(ctx, input_info, zero_dims.size(), + zero_dims.data(), op_name)); + + LRT_TENSOR_ASSIGN_OR_RETURN(auto params, + BuildBinaryParams(kActNone, op_name)); + + LRT_TENSOR_RETURN_IF_ERROR(xnn_define_binary(ctx.subgraph(), xnn_binary_add, + ¶ms, reshape_id, zero_id, + broadcast_id, + /*flags=*/0)) << op_name; const std::array output_dims = { diff --git a/litert/tensor/backends/xnnpack/arithmetic.h b/litert/tensor/backends/xnnpack/arithmetic.h index dfb735bdb67..f79d316617f 100644 --- a/litert/tensor/backends/xnnpack/arithmetic.h +++ b/litert/tensor/backends/xnnpack/arithmetic.h @@ -19,6 +19,7 @@ limitations under the License. #include #include #include +#include #include #include "include/xnnpack.h" @@ -37,6 +38,21 @@ struct xnn_subgraph; namespace litert::tensor { +class BufferHolder { + public: + virtual ~BufferHolder() = default; +}; + +template +class TypedBufferHolder : public BufferHolder { + public: + explicit TypedBufferHolder(std::vector&& vec) : vec_(std::move(vec)) {} + const T* data() const { return vec_.data(); } + + private: + std::vector vec_; +}; + // Tag to identify the XNNPACK mixin. struct XnnpackMixinTag {}; @@ -65,6 +81,14 @@ class XnnpackBuildContext { // Returns the XNNPACK subgraph. ::xnn_subgraph* subgraph(); + template + const T* KeepAlive(std::vector&& buffer) { + auto holder = std::make_unique>(std::move(buffer)); + const T* ptr = holder->data(); + custom_buffers_.push_back(std::move(holder)); + return ptr; + } + private: xnn_subgraph* subgraph_ = nullptr; std::vector outputs_; @@ -74,6 +98,7 @@ class XnnpackBuildContext { absl::flat_hash_map external_ids_; std::vector> dequantized_buffers_; std::vector> fp16_buffers_; + std::vector> custom_buffers_; friend absl::StatusOr> BuildXnnpackGraph( std::vector outputs); diff --git a/litert/tensor/backends/xnnpack/conversion.cc b/litert/tensor/backends/xnnpack/conversion.cc index 0e1f9adb558..d770353e21a 100644 --- a/litert/tensor/backends/xnnpack/conversion.cc +++ b/litert/tensor/backends/xnnpack/conversion.cc @@ -56,50 +56,6 @@ absl::Status EnsureXnnInitialized() { return *g_xnn_init_status; } -xnn_datatype GetXnnpackType(const XnnpackValue& value) { - switch (value.info.type) { - case Type::kUnknown: - case Type::kBOOL: - case Type::kI2: - case Type::kI4: - if (value.info.quantization) { - if (value.info.quantization->As().ok()) { - return xnn_datatype_qcint4; - } else if (value.info.quantization->As().ok()) { - return xnn_datatype_qbint4; - } - } - break; - case Type::kI8: - if (value.info.quantization) { - if (auto it = - value.info.quantization->As(); - it.ok()) { - return it->scales.size() > 1 ? xnn_datatype_qcint8 - : xnn_datatype_qint8; - } - } - break; - case Type::kI16: - case Type::kI64: - case Type::kU4: - case Type::kU8: - case Type::kU16: - case Type::kU32: - case Type::kU64: - case Type::kFP16: - return xnn_datatype_fp16; - case Type::kI32: - return xnn_datatype_int32; - case Type::kFP32: - return xnn_datatype_fp32; - case Type::kFP64: - break; - case Type::kBF16: - return xnn_datatype_bf16; - } - return xnn_datatype_invalid; -} // TODO: b/493560478 - Decide whether to delete this from here. [[maybe_unused]] @@ -225,13 +181,15 @@ XnnpackGraph::XnnpackGraph( absl::flat_hash_map tensor_index, absl::flat_hash_set external_outputs, std::vector> dequantized_buffers, - std::vector> fp16_buffers) + std::vector> fp16_buffers, + std::vector> custom_buffers) : subgraph_(subgraph), values_(std::move(values)), tensor_index_(std::move(tensor_index)), external_outputs_(std::move(external_outputs)), dequantized_buffers_(std::move(dequantized_buffers)), - fp16_buffers_(std::move(fp16_buffers)) {} + fp16_buffers_(std::move(fp16_buffers)), + custom_buffers_(std::move(custom_buffers)) {} XnnpackGraph::~XnnpackGraph() { if (subgraph_ != nullptr) { @@ -285,7 +243,7 @@ absl::StatusOr> XnnpackBuildContext::Finalize() { return std::make_unique( subgraph, std::move(values_), std::move(tensor_index_), std::move(external_outputs_), std::move(dequantized_buffers_), - std::move(fp16_buffers_)); + std::move(fp16_buffers_), std::move(custom_buffers_)); } absl::StatusOr> BuildXnnpackGraph( @@ -370,10 +328,10 @@ absl::StatusOr XnnpackBuildContext::DefineValue( } if (!info.quantization) { - LRT_TENSOR_RETURN_IF_ERROR( - xnn_define_tensor_value(subgraph_, GetXnnpackType(value), dims.size(), - dims.empty() ? nullptr : dims.data(), data_ptr, - external_id, value.flags, &value.id)) + LRT_TENSOR_RETURN_IF_ERROR(xnn_define_tensor_value( + subgraph_, GetXnnpackType(value.info), dims.size(), + dims.empty() ? nullptr : dims.data(), data_ptr, external_id, + value.flags, &value.id)) << "Could not define a new tensor value."; } else if (auto maybe_pcq = info.quantization->As(); @@ -381,7 +339,7 @@ absl::StatusOr XnnpackBuildContext::DefineValue( const auto& pcq = maybe_pcq.value(); if (pcq.scales.size() == 1) { LRT_TENSOR_RETURN_IF_ERROR(xnn_define_quantized_tensor_value( - subgraph_, GetXnnpackType(value), + subgraph_, GetXnnpackType(value.info), pcq.zero_points.empty() ? 0 : pcq.zero_points[0], pcq.scales[0], dims.size(), dims.empty() ? nullptr : dims.data(), data_ptr, external_id, value.flags, &value.id)) @@ -408,7 +366,7 @@ absl::StatusOr XnnpackBuildContext::DefineValue( } else { LRT_TENSOR_RETURN_IF_ERROR( xnn_define_channelwise_quantized_tensor_value_v3( - subgraph_, GetXnnpackType(value), /*zero_point=*/0, + subgraph_, GetXnnpackType(value.info), /*zero_point=*/0, pcq.scales.data(), dims.size(), pcq.quantized_dimension, dims.empty() ? nullptr : dims.data(), data_ptr, external_id, value.flags, &value.id, /*channelwise_zero_point=*/nullptr)) @@ -422,8 +380,8 @@ absl::StatusOr XnnpackBuildContext::DefineValue( const void* scale_ptr = fp16_buffers_.back().data(); int32_t zero_point = bwq.zero_points.empty() ? 0 : bwq.zero_points[0]; LRT_TENSOR_RETURN_IF_ERROR(xnn_define_blockwise_quantized_tensor_value_v2( - subgraph_, GetXnnpackType(value), zero_point, scale_ptr, dims.size(), - bwq.quantized_dimension, bwq.block_size, + subgraph_, GetXnnpackType(value.info), zero_point, scale_ptr, + dims.size(), bwq.quantized_dimension, bwq.block_size, dims.empty() ? nullptr : dims.data(), data_ptr, external_id, value.flags, xnn_datatype_fp16, &value.id)) << "Could not define a new blockwise quantized tensor value."; diff --git a/litert/tensor/backends/xnnpack/conversion.h b/litert/tensor/backends/xnnpack/conversion.h index 5fc18776d76..e44993060b7 100644 --- a/litert/tensor/backends/xnnpack/conversion.h +++ b/litert/tensor/backends/xnnpack/conversion.h @@ -39,7 +39,8 @@ class XnnpackGraph { absl::flat_hash_map tensor_index, absl::flat_hash_set external_outputs, std::vector> dequantized_buffers = {}, - std::vector> fp16_buffers = {}); + std::vector> fp16_buffers = {}, + std::vector> custom_buffers = {}); ~XnnpackGraph(); // Returns the XNNPACK subgraph. @@ -66,6 +67,7 @@ class XnnpackGraph { absl::flat_hash_set external_outputs_; std::vector> dequantized_buffers_; std::vector> fp16_buffers_; + std::vector> custom_buffers_; }; // Builds an XNNPACK graph from the given outputs. diff --git a/litert/tensor/backends/xnnpack/utils.h b/litert/tensor/backends/xnnpack/utils.h index 15ba1e9f3d9..ac0d5dac272 100644 --- a/litert/tensor/backends/xnnpack/utils.h +++ b/litert/tensor/backends/xnnpack/utils.h @@ -20,6 +20,7 @@ limitations under the License. #include "absl/status/status.h" #include "absl/strings/str_cat.h" #include "absl/strings/string_view.h" +#include "litert/tensor/internal/graph.h" #include "litert/tensor/utils/macros.h" namespace litert::tensor { @@ -44,6 +45,51 @@ inline absl::Status XnnStatusToAbsl(enum xnn_status status, absl::StrCat("xnn_status=", static_cast(status), ";", label)); } +inline xnn_datatype GetXnnpackType(const graph::TensorInformation& info) { + switch (info.type) { + case Type::kUnknown: + case Type::kBOOL: + case Type::kI2: + case Type::kI4: + if (info.quantization) { + if (info.quantization->As().ok()) { + return xnn_datatype_qcint4; + } else if (info.quantization->As().ok()) { + return xnn_datatype_qbint4; + } + } + break; + case Type::kI8: + if (info.quantization) { + if (auto it = + info.quantization->As(); + it.ok()) { + return it->scales.size() > 1 ? xnn_datatype_qcint8 + : xnn_datatype_qint8; + } + } + break; + case Type::kI16: + case Type::kI64: + case Type::kU4: + case Type::kU8: + case Type::kU16: + case Type::kU32: + case Type::kU64: + case Type::kFP16: + return xnn_datatype_fp16; + case Type::kI32: + return xnn_datatype_int32; + case Type::kFP32: + return xnn_datatype_fp32; + case Type::kFP64: + break; + case Type::kBF16: + return xnn_datatype_bf16; + } + return xnn_datatype_invalid; +} + } // namespace litert::tensor #endif // LITERT_TENSOR_BACKENDS_XNNPACK_UTILS_H_