mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-20 19:12:24 +00:00
[NNAPI EP] Move GetCapability independent of ModelBuilder (#5767)
* Move GetCapability independent of ModelBuilder * minor code style fix * Move ort_enforce for same number of op_builders and op_support_checkers * minor code fix
This commit is contained in:
parent
d8ace07ad7
commit
4207e99be3
10 changed files with 254 additions and 232 deletions
|
|
@ -9,9 +9,11 @@
|
|||
#include <core/common/safeint.h>
|
||||
#include <core/common/logging/logging.h>
|
||||
#include <core/graph/graph.h>
|
||||
#include <core/graph/graph_viewer.h>
|
||||
#include <core/providers/common.h>
|
||||
|
||||
#include "helper.h"
|
||||
#include "op_support_checker.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace nnapi {
|
||||
|
|
@ -93,7 +95,7 @@ bool HasValidBinaryOpQuantizedInputs(const Node& node) {
|
|||
return true;
|
||||
}
|
||||
|
||||
bool HasValidQuantizationScales(const InitializerMap& initializers, const Node& node,
|
||||
bool HasValidQuantizationScales(const InitializedTensorSet& initializers, const Node& node,
|
||||
const std::vector<size_t>& indices) {
|
||||
const auto& op = node.OpType();
|
||||
const auto input_defs(node.InputDefs());
|
||||
|
|
@ -105,7 +107,7 @@ bool HasValidQuantizationScales(const InitializerMap& initializers, const Node&
|
|||
}
|
||||
const auto scale_name = input_defs[idx]->Name();
|
||||
if (Contains(initializers, scale_name)) {
|
||||
const auto& tensor = initializers.at(scale_name);
|
||||
const auto& tensor = *initializers.at(scale_name);
|
||||
if (!tensor.dims().empty() && tensor.dims()[0] != 1) {
|
||||
LOGS_DEFAULT(VERBOSE) << op << " does not support per-channel quantization";
|
||||
return false;
|
||||
|
|
@ -119,7 +121,7 @@ bool HasValidQuantizationScales(const InitializerMap& initializers, const Node&
|
|||
return true;
|
||||
}
|
||||
|
||||
bool HasValidQuantizationZeroPoints(const InitializerMap& initializers, const Node& node,
|
||||
bool HasValidQuantizationZeroPoints(const InitializedTensorSet& initializers, const Node& node,
|
||||
const std::vector<size_t>& indices) {
|
||||
const auto& op = node.OpType();
|
||||
const auto input_defs(node.InputDefs());
|
||||
|
|
@ -131,7 +133,7 @@ bool HasValidQuantizationZeroPoints(const InitializerMap& initializers, const No
|
|||
}
|
||||
const auto zero_point_name = node.InputDefs()[idx]->Name();
|
||||
if (Contains(initializers, zero_point_name)) {
|
||||
const auto& tensor = initializers.at(zero_point_name);
|
||||
const auto& tensor = *initializers.at(zero_point_name);
|
||||
if (!tensor.dims().empty() && tensor.dims()[0] != 1) {
|
||||
LOGS_DEFAULT(VERBOSE) << op << " does not support per-channel quantization";
|
||||
return false;
|
||||
|
|
@ -191,7 +193,7 @@ bool GetType(const NodeArg& node_arg, int32_t& type) {
|
|||
return true;
|
||||
}
|
||||
|
||||
bool GetClipMinMax(const InitializerMap& initializers, const Node& node, float& min, float& max) {
|
||||
bool GetClipMinMax(const InitializedTensorSet& initializers, const Node& node, float& min, float& max) {
|
||||
min = std::numeric_limits<float>::lowest();
|
||||
max = std::numeric_limits<float>::max();
|
||||
if (node.SinceVersion() < 11) { // Clip opset 1, 6 is using attributes for min/max
|
||||
|
|
@ -205,7 +207,7 @@ bool GetClipMinMax(const InitializerMap& initializers, const Node& node, float&
|
|||
LOGS_DEFAULT(VERBOSE) << "Input min of Clip must be known";
|
||||
return false;
|
||||
}
|
||||
min = GetTensorFloatData(initializers.at(min_name))[0];
|
||||
min = GetTensorFloatData(*initializers.at(min_name))[0];
|
||||
}
|
||||
|
||||
if (node.InputDefs().size() > 2) { // we have input max
|
||||
|
|
@ -214,7 +216,7 @@ bool GetClipMinMax(const InitializerMap& initializers, const Node& node, float&
|
|||
LOGS_DEFAULT(VERBOSE) << "Input max of Clip must be known";
|
||||
return false;
|
||||
}
|
||||
max = GetTensorFloatData(initializers.at(max_name))[0];
|
||||
max = GetTensorFloatData(*initializers.at(max_name))[0];
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -234,6 +236,75 @@ void GetFlattenOutputShape(const Node& node, const Shape& input_shape, int32_t&
|
|||
dim_2 = std::accumulate(input_shape.cbegin() + axis, input_shape.cend(), 1, std::multiplies<int32_t>());
|
||||
}
|
||||
|
||||
bool IsValidSupportedNodesVec(const std::vector<int>& supported_node_vec, const GraphViewer& graph_viewer) {
|
||||
if (supported_node_vec.empty())
|
||||
return false;
|
||||
|
||||
if (supported_node_vec.size() == 1) {
|
||||
const auto& node_indices = graph_viewer.GetNodesInTopologicalOrder();
|
||||
const auto* node(graph_viewer.GetNode(node_indices[supported_node_vec[0]]));
|
||||
const auto& op = node->OpType();
|
||||
// It is not worth it to perform a single Reshape/Flatten/Identity operator
|
||||
// which is only copying the data in NNAPI
|
||||
// If this is the case, let it fall back
|
||||
if (op == "Reshape" ||
|
||||
op == "Flatten" ||
|
||||
op == "Identity") {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
bool IsNodeSupported(const Node& node, const GraphViewer& graph_viewer, const OpSupportCheckParams& params) {
|
||||
const auto& op_support_checkers = GetOpSupportCheckers();
|
||||
if (Contains(op_support_checkers, node.OpType())) {
|
||||
const auto op_support_checker = op_support_checkers.at(node.OpType());
|
||||
return op_support_checker->IsOpSupported(graph_viewer.GetAllInitializedTensors(), node, params);
|
||||
} else {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
std::vector<std::vector<int>> GetSupportedNodes(const GraphViewer& graph_viewer, const OpSupportCheckParams& params) {
|
||||
std::vector<std::vector<int>> supported_node_vecs;
|
||||
if (params.android_sdk_ver < ORT_NNAPI_MIN_API_LEVEL) {
|
||||
LOGS_DEFAULT(WARNING) << "All ops will fallback to CPU EP, because Android API level [" << params.android_sdk_ver
|
||||
<< "] is lower than minimal supported API level [" << ORT_NNAPI_MIN_API_LEVEL
|
||||
<< "] of this build for NNAPI";
|
||||
return supported_node_vecs;
|
||||
}
|
||||
|
||||
std::vector<int> supported_node_vec;
|
||||
const auto& node_indices = graph_viewer.GetNodesInTopologicalOrder();
|
||||
for (size_t i = 0; i < node_indices.size(); i++) {
|
||||
const auto* node(graph_viewer.GetNode(node_indices[i]));
|
||||
bool supported = IsNodeSupported(*node, graph_viewer, params);
|
||||
LOGS_DEFAULT(VERBOSE) << "Operator type: [" << node->OpType()
|
||||
<< "] index: [" << i
|
||||
<< "] name: [" << node->Name()
|
||||
<< "] supported: [" << supported
|
||||
<< "]";
|
||||
if (supported) {
|
||||
supported_node_vec.push_back(i);
|
||||
} else {
|
||||
if (IsValidSupportedNodesVec(supported_node_vec, graph_viewer)) {
|
||||
supported_node_vecs.push_back(supported_node_vec);
|
||||
supported_node_vec.clear();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (IsValidSupportedNodesVec(supported_node_vec, graph_viewer))
|
||||
supported_node_vecs.push_back(supported_node_vec);
|
||||
|
||||
LOGS_DEFAULT(VERBOSE) << "Support vectors size is " << supported_node_vecs.size();
|
||||
for (const auto& group : supported_node_vecs)
|
||||
LOGS_DEFAULT(VERBOSE) << "Support vector size is " << group.size();
|
||||
|
||||
return supported_node_vecs;
|
||||
}
|
||||
|
||||
std::string Shape2String(const std::vector<uint32_t>& shape) {
|
||||
std::ostringstream os;
|
||||
os << "[ ";
|
||||
|
|
|
|||
|
|
@ -7,6 +7,11 @@
|
|||
#include "core/graph/basic_types.h"
|
||||
#include "core/providers/nnapi/nnapi_builtin/nnapi_lib/NeuralNetworksTypes.h"
|
||||
|
||||
// This is the minimal Android API Level required by ORT NNAPI EP to run
|
||||
#ifndef ORT_NNAPI_MIN_API_LEVEL
|
||||
#define ORT_NNAPI_MIN_API_LEVEL 27
|
||||
#endif
|
||||
|
||||
namespace onnxruntime {
|
||||
|
||||
using Shape = std::vector<uint32_t>;
|
||||
|
|
@ -14,9 +19,13 @@ using InitializerMap = std::unordered_map<std::string, const ONNX_NAMESPACE::Ten
|
|||
|
||||
class Node;
|
||||
class NodeArg;
|
||||
class GraphViewer;
|
||||
|
||||
namespace nnapi {
|
||||
|
||||
class IOpSupportChecker;
|
||||
struct OpSupportCheckParams;
|
||||
|
||||
#define THROW_ON_ERROR(val) \
|
||||
{ \
|
||||
const auto ret = (val); \
|
||||
|
|
@ -78,10 +87,10 @@ bool IsQLinearBinaryOp(QLinearOpType qlinear_op_type);
|
|||
// Check if a qlinear binary op has valid inputs
|
||||
bool HasValidBinaryOpQuantizedInputs(const Node& node);
|
||||
// Check if a qlinear op has valid scales for given indices
|
||||
bool HasValidQuantizationScales(const InitializerMap& initializers, const Node& node,
|
||||
bool HasValidQuantizationScales(const InitializedTensorSet& initializers, const Node& node,
|
||||
const std::vector<size_t>& indices);
|
||||
// Check if a qlinear op has valid zero points for given indices
|
||||
bool HasValidQuantizationZeroPoints(const InitializerMap& initializers, const Node& node,
|
||||
bool HasValidQuantizationZeroPoints(const InitializedTensorSet& initializers, const Node& node,
|
||||
const std::vector<size_t>& indices);
|
||||
|
||||
// Get initialize tensort float/int32/int64 data without unpacking
|
||||
|
|
@ -96,11 +105,17 @@ bool GetType(const NodeArg& node_arg, int32_t& type);
|
|||
|
||||
// Get the min/max value from Clip op
|
||||
// If the min/max are inputs be not initializers (value not preset), will return false
|
||||
bool GetClipMinMax(const InitializerMap& initializers, const Node& node, float& min, float& max);
|
||||
bool GetClipMinMax(const InitializedTensorSet& initializers, const Node& node, float& min, float& max);
|
||||
|
||||
// Get the output shape of Flatten Op
|
||||
void GetFlattenOutputShape(const Node& node, const Shape& input_shape, int32_t& dim_1, int32_t& dim_2);
|
||||
|
||||
// If a node is supported by NNAPI
|
||||
bool IsNodeSupported(const Node& node, const GraphViewer& graph_viewer, const OpSupportCheckParams& params);
|
||||
|
||||
// Get a list of groups of supported nodes, each group represents a subgraph supported by NNAPI EP
|
||||
std::vector<std::vector<int>> GetSupportedNodes(const GraphViewer& graph_viewer, const OpSupportCheckParams& params);
|
||||
|
||||
// Get string representation of a Shape
|
||||
std::string Shape2String(const std::vector<uint32_t>& shape);
|
||||
|
||||
|
|
|
|||
|
|
@ -19,90 +19,12 @@ using std::vector;
|
|||
ModelBuilder::ModelBuilder(const GraphViewer& graph_viewer)
|
||||
: nnapi_(NnApiImplementation()), graph_viewer_(graph_viewer) {
|
||||
GetAllInitializers();
|
||||
op_builders_ = CreateOpBuilders();
|
||||
op_support_checkers_ = CreateOpSupportCheckers();
|
||||
ORT_ENFORCE(op_builders_.size() == op_support_checkers_.size(),
|
||||
"We should have same number of OpBuilder and OpSupportChecker");
|
||||
}
|
||||
|
||||
int32_t ModelBuilder::GetAndroidSdkVer() const {
|
||||
return nnapi_ ? nnapi_->android_sdk_version : 0;
|
||||
}
|
||||
|
||||
bool ModelBuilder::IsNodeSupported(const Node& node) {
|
||||
if (auto* op_support_checker = GetOPSupportChecker(node)) {
|
||||
OPSupportCheckParams param{
|
||||
GetAndroidSdkVer(), // android_sdk_ver
|
||||
UseNCHW(), // use_nchw
|
||||
};
|
||||
return op_support_checker->IsOpSupported(GetInitializerTensors(), node, param);
|
||||
} else {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
bool IsValidSupportedNodesVec(const std::vector<int>& supported_node_vec, const GraphViewer& graph_viewer) {
|
||||
if (supported_node_vec.empty())
|
||||
return false;
|
||||
|
||||
if (supported_node_vec.size() == 1) {
|
||||
const auto& node_indices = graph_viewer.GetNodesInTopologicalOrder();
|
||||
const auto* node(graph_viewer.GetNode(node_indices[supported_node_vec[0]]));
|
||||
const auto& op = node->OpType();
|
||||
// It is not worth it to perform a single Reshape/Flatten/Identity operator
|
||||
// which is only copying the data in NNAPI
|
||||
// If this is the case, let it fall back
|
||||
if (op == "Reshape" ||
|
||||
op == "Flatten" ||
|
||||
op == "Identity") {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
std::vector<std::vector<int>> ModelBuilder::GetSupportedNodes() {
|
||||
std::vector<std::vector<int>> supported_node_vecs;
|
||||
int32_t android_sdk_ver = GetAndroidSdkVer();
|
||||
#ifdef __ANDROID__
|
||||
if (android_sdk_ver < ORT_NNAPI_MIN_API_LEVEL) {
|
||||
LOGS_DEFAULT(WARNING) << "All ops will fallback to CPU EP, because Android API level [" << android_sdk_ver
|
||||
<< "] is lower than minimal supported API level [" << ORT_NNAPI_MIN_API_LEVEL
|
||||
<< "] of this build for NNAPI";
|
||||
return supported_node_vecs;
|
||||
}
|
||||
#endif
|
||||
|
||||
std::vector<int> supported_node_vec;
|
||||
const auto& node_indices = graph_viewer_.GetNodesInTopologicalOrder();
|
||||
for (size_t i = 0; i < node_indices.size(); i++) {
|
||||
const auto* node(graph_viewer_.GetNode(node_indices[i]));
|
||||
bool supported = IsNodeSupported(*node);
|
||||
LOGS_DEFAULT(VERBOSE) << "Operator type: [" << node->OpType()
|
||||
<< "] index: [" << i
|
||||
<< "] name: [" << node->Name()
|
||||
<< "] supported: [" << supported
|
||||
<< "]";
|
||||
if (supported) {
|
||||
supported_node_vec.push_back(i);
|
||||
} else {
|
||||
if (IsValidSupportedNodesVec(supported_node_vec, graph_viewer_)) {
|
||||
supported_node_vecs.push_back(supported_node_vec);
|
||||
supported_node_vec.clear();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (IsValidSupportedNodesVec(supported_node_vec, graph_viewer_))
|
||||
supported_node_vecs.push_back(supported_node_vec);
|
||||
|
||||
LOGS_DEFAULT(VERBOSE) << "Support vectors size is " << supported_node_vecs.size();
|
||||
for (const auto& group : supported_node_vecs)
|
||||
LOGS_DEFAULT(VERBOSE) << "Support vector size is " << group.size();
|
||||
|
||||
return supported_node_vecs;
|
||||
}
|
||||
|
||||
// Scalar operand is copied into the model, no need to persist
|
||||
#define DEFINE_ADD_OPERAND_FROM_SCALAR(scalar_type, op_type) \
|
||||
Status ModelBuilder::AddOperandFromScalar(scalar_type value, uint32_t& index) { \
|
||||
|
|
@ -614,17 +536,19 @@ int32_t ModelBuilder::FindActivation(const Node& node, const NodeArg& output) {
|
|||
}
|
||||
|
||||
IOpBuilder* ModelBuilder::GetOpBuilder(const Node& node) {
|
||||
if (!Contains(op_builders_, node.OpType()))
|
||||
const auto& op_builders = GetOpBuilders();
|
||||
if (!Contains(op_builders, node.OpType()))
|
||||
return nullptr;
|
||||
|
||||
return op_builders_[node.OpType()].get();
|
||||
return op_builders.at(node.OpType()).get();
|
||||
}
|
||||
|
||||
IOpSupportChecker* ModelBuilder::GetOPSupportChecker(const Node& node) {
|
||||
if (!Contains(op_support_checkers_, node.OpType()))
|
||||
const auto& op_support_checkers = GetOpSupportCheckers();
|
||||
if (!Contains(op_support_checkers, node.OpType()))
|
||||
return nullptr;
|
||||
|
||||
return op_support_checkers_[node.OpType()].get();
|
||||
return op_support_checkers.at(node.OpType()).get();
|
||||
}
|
||||
|
||||
std::string ModelBuilder::GetUniqueName(const std::string& base_name) {
|
||||
|
|
|
|||
|
|
@ -10,11 +10,6 @@
|
|||
#include "core/providers/nnapi/nnapi_builtin/nnapi_lib/NeuralNetworksWrapper.h"
|
||||
#include "shaper.h"
|
||||
|
||||
// This is the minimal Android API Level required by ORT NNAPI EP to run
|
||||
#ifndef ORT_NNAPI_MIN_API_LEVEL
|
||||
#define ORT_NNAPI_MIN_API_LEVEL 27
|
||||
#endif
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace nnapi {
|
||||
|
||||
|
|
@ -37,8 +32,6 @@ class ModelBuilder {
|
|||
ModelBuilder(const GraphViewer& graph_viewer);
|
||||
~ModelBuilder() = default;
|
||||
|
||||
std::vector<std::vector<int>> GetSupportedNodes();
|
||||
|
||||
Status Compile(std::unique_ptr<Model>& model) ORT_MUST_USE_RESULT;
|
||||
|
||||
int32_t GetAndroidSdkVer() const;
|
||||
|
|
@ -98,10 +91,9 @@ class ModelBuilder {
|
|||
const std::unordered_set<std::string>&
|
||||
GetFusedActivations() const { return fused_activations_; }
|
||||
|
||||
const std::unordered_map<std::string, const ONNX_NAMESPACE::TensorProto&>&
|
||||
GetInitializerTensors() const { return initializers_; }
|
||||
const InitializedTensorSet& GetInitializerTensors() const { return graph_viewer_.GetAllInitializedTensors(); }
|
||||
|
||||
const Graph& GetOnnxGraph() const { return graph_viewer_.GetGraph(); }
|
||||
const GraphViewer& GetGraphViewer() const { return graph_viewer_; }
|
||||
|
||||
void RegisterNHWCOperand(const std::string& name);
|
||||
bool IsOperandNHWC(const std::string& name);
|
||||
|
|
@ -115,9 +107,6 @@ class ModelBuilder {
|
|||
Status SetNCHWToNHWCOperandMap(const std::string& nchw_name,
|
||||
const std::string& nhwc_name) ORT_MUST_USE_RESULT;
|
||||
|
||||
// Is the given node supported by NNAPI
|
||||
bool IsNodeSupported(const Node& node);
|
||||
|
||||
private:
|
||||
const NnApi* nnapi_{nullptr};
|
||||
const GraphViewer& graph_viewer_;
|
||||
|
|
@ -144,7 +133,6 @@ class ModelBuilder {
|
|||
// All activation nodes (Relu, Relu1, Relu6) as a map <NodeIndex, activation_code>
|
||||
std::unordered_map<NodeIndex, int32_t> activation_nodes_;
|
||||
|
||||
std::unordered_map<std::string, std::shared_ptr<IOpBuilder>> op_builders_;
|
||||
std::unordered_map<std::string, std::shared_ptr<IOpSupportChecker>> op_support_checkers_;
|
||||
|
||||
// Operands in nhwc
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@
|
|||
#include "helper.h"
|
||||
#include "model_builder.h"
|
||||
#include "op_builder.h"
|
||||
#include "op_support_checker.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace nnapi {
|
||||
|
|
@ -187,7 +188,7 @@ static Status AddInitializerInNewLayout(ModelBuilder& model_builder,
|
|||
const std::string& name,
|
||||
const OperandType& source_operand_type,
|
||||
DataLayout new_layout) {
|
||||
const auto& tensor = model_builder.GetInitializerTensors().at(name);
|
||||
const auto& tensor = *model_builder.GetInitializerTensors().at(name);
|
||||
const Shape& shape = source_operand_type.dimensions;
|
||||
ORT_RETURN_IF_NOT(shape.size() == 4,
|
||||
"The initializer is not 4D: ", name, " actual dim ", shape.size());
|
||||
|
|
@ -267,7 +268,7 @@ static Status AddInitializerTransposed(ModelBuilder& model_builder,
|
|||
static Status AddInitializerTransposed(ModelBuilder& model_builder,
|
||||
const OperandType& source_operand_type,
|
||||
const std::string& name) {
|
||||
const auto& tensor = model_builder.GetInitializerTensors().at(name);
|
||||
const auto& tensor = *model_builder.GetInitializerTensors().at(name);
|
||||
const Shape& shape = source_operand_type.dimensions;
|
||||
|
||||
ORT_RETURN_IF_NOT(shape.size() == 2,
|
||||
|
|
@ -399,7 +400,7 @@ static Status HandleAutoPad(const Shape& input_shape,
|
|||
}
|
||||
|
||||
static float GetQuantizationScale(const ModelBuilder& model_builder, const Node& node, size_t idx) {
|
||||
const auto& scale_tensor = model_builder.GetInitializerTensors().at(node.InputDefs()[idx]->Name());
|
||||
const auto& scale_tensor = *model_builder.GetInitializerTensors().at(node.InputDefs()[idx]->Name());
|
||||
return GetTensorFloatData(scale_tensor)[0];
|
||||
}
|
||||
|
||||
|
|
@ -408,7 +409,7 @@ static Status GetQuantizationZeroPoint(const ModelBuilder& model_builder, const
|
|||
static Status GetQuantizationZeroPoint(const ModelBuilder& model_builder, const Node& node, size_t idx, int32_t& zero_point) {
|
||||
std::unique_ptr<uint8_t[]> unpacked_tensor;
|
||||
size_t tensor_byte_size;
|
||||
const auto& zero_point_tensor = model_builder.GetInitializerTensors().at(node.InputDefs()[idx]->Name());
|
||||
const auto& zero_point_tensor = *model_builder.GetInitializerTensors().at(node.InputDefs()[idx]->Name());
|
||||
ORT_RETURN_IF_ERROR(
|
||||
onnxruntime::utils::UnpackInitializerData(zero_point_tensor, unpacked_tensor, tensor_byte_size));
|
||||
zero_point = static_cast<int32_t>(unpacked_tensor.get()[0]);
|
||||
|
|
@ -535,7 +536,12 @@ class BaseOpBuilder : public IOpBuilder {
|
|||
};
|
||||
|
||||
Status BaseOpBuilder::AddToModelBuilder(ModelBuilder& model_builder, const Node& node) const {
|
||||
ORT_RETURN_IF_NOT(model_builder.IsNodeSupported(node), "Unsupported operator ", node.OpType());
|
||||
OpSupportCheckParams params{
|
||||
model_builder.GetAndroidSdkVer(),
|
||||
model_builder.UseNCHW(),
|
||||
};
|
||||
|
||||
ORT_RETURN_IF_NOT(IsNodeSupported(node, model_builder.GetGraphViewer(), params), "Unsupported operator ", node.OpType());
|
||||
ORT_RETURN_IF_ERROR(AddToModelBuilderImpl(model_builder, node));
|
||||
LOGS_DEFAULT(VERBOSE) << "Operator name: [" << node.Name()
|
||||
<< "] type: [" << node.OpType() << "] was added";
|
||||
|
|
@ -828,7 +834,7 @@ Status ReshapeOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, cons
|
|||
ORT_RETURN_IF_ERROR(GetNCHWInput(model_builder, node, 0, input));
|
||||
}
|
||||
|
||||
const auto& shape_tensor = initializers.at(node.InputDefs()[1]->Name());
|
||||
const auto& shape_tensor = *initializers.at(node.InputDefs()[1]->Name());
|
||||
const int64_t* raw_shape = GetTensorInt64Data(shape_tensor);
|
||||
const auto size = SafeInt<uint32_t>(shape_tensor.dims()[0]);
|
||||
|
||||
|
|
@ -874,10 +880,10 @@ Status BatchNormalizationOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_bu
|
|||
const auto& input = node.InputDefs()[0]->Name();
|
||||
const auto& output = node.OutputDefs()[0]->Name();
|
||||
|
||||
const auto& scale_tensor = initializers.at(node.InputDefs()[1]->Name());
|
||||
const auto& bias_tensor = initializers.at(node.InputDefs()[2]->Name());
|
||||
const auto& mean_tensor = initializers.at(node.InputDefs()[3]->Name());
|
||||
const auto& var_tensor = initializers.at(node.InputDefs()[4]->Name());
|
||||
const auto& scale_tensor = *initializers.at(node.InputDefs()[1]->Name());
|
||||
const auto& bias_tensor = *initializers.at(node.InputDefs()[2]->Name());
|
||||
const auto& mean_tensor = *initializers.at(node.InputDefs()[3]->Name());
|
||||
const auto& var_tensor = *initializers.at(node.InputDefs()[4]->Name());
|
||||
const auto eps = helper.Get("epsilon", 1e-5f);
|
||||
|
||||
const auto size = SafeInt<uint32_t>(scale_tensor.dims()[0]);
|
||||
|
|
@ -1126,7 +1132,7 @@ Status ConvOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const N
|
|||
}
|
||||
|
||||
const auto& weight = input_defs[w_idx]->Name();
|
||||
const auto& weight_tensor = initializers.at(weight);
|
||||
const auto& weight_tensor = *initializers.at(weight);
|
||||
bool conv_2d = false,
|
||||
depthwise_conv_2d = false,
|
||||
grouped_conv_2d = false;
|
||||
|
|
@ -1199,7 +1205,7 @@ Status ConvOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const N
|
|||
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Unknown weight type ", TypeToStr(weight_type));
|
||||
}
|
||||
} else if (is_qlinear_conv) { // QLinearConv's bias type need special handling
|
||||
const auto& bias_tensor = model_builder.GetInitializerTensors().at(bias);
|
||||
const auto& bias_tensor = *model_builder.GetInitializerTensors().at(bias);
|
||||
ORT_RETURN_IF_NOT(bias_tensor.data_type() == ONNX_NAMESPACE::TensorProto_DataType_INT32,
|
||||
"bias of QLinearConv should be int32, actual type: ", bias_tensor.data_type());
|
||||
Shape bias_dimen;
|
||||
|
|
@ -1479,7 +1485,7 @@ Status GemmOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const N
|
|||
else
|
||||
onnx_mat_b_type = Type::TENSOR_QUANT8_ASYMM;
|
||||
|
||||
const auto& mat_b_tensor = initializers.at(input2);
|
||||
const auto& mat_b_tensor = *initializers.at(input2);
|
||||
Shape onnx_mat_b_shape;
|
||||
for (auto dim : mat_b_tensor.dims())
|
||||
onnx_mat_b_shape.push_back(SafeInt<uint32_t>(dim));
|
||||
|
|
@ -1974,7 +1980,7 @@ Status ResizeOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const
|
|||
|
||||
if (input_defs.size() == 3) { // we are using scales
|
||||
const auto& scales_name = input_defs[2]->Name();
|
||||
const auto& scales_tensor = initializers.at(scales_name);
|
||||
const auto& scales_tensor = *initializers.at(scales_name);
|
||||
const float* scales_data = GetTensorFloatData(scales_tensor);
|
||||
float scale_h = scales_data[2];
|
||||
float scale_w = scales_data[3];
|
||||
|
|
@ -1982,7 +1988,7 @@ Status ResizeOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const
|
|||
shaper.ResizeUsingScales(input, scale_h, scale_w, use_nchw, output));
|
||||
} else { // we are using sizes
|
||||
const auto& sizes_name = input_defs[3]->Name();
|
||||
const auto& sizes_tensor = initializers.at(sizes_name);
|
||||
const auto& sizes_tensor = *initializers.at(sizes_name);
|
||||
const int64_t* sizes_data = GetTensorInt64Data(sizes_tensor);
|
||||
ORT_RETURN_IF_ERROR(
|
||||
shaper.ResizeUsingOutputSizes(input, SafeInt<uint32_t>(sizes_data[2]), SafeInt<uint32_t>(sizes_data[3]), use_nchw, output));
|
||||
|
|
@ -2048,10 +2054,9 @@ Status FlattenOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, cons
|
|||
|
||||
#pragma endregion op_reshape
|
||||
|
||||
#pragma region CreateOpBuilders
|
||||
#pragma region CreateGetOpBuilders
|
||||
|
||||
std::unordered_map<std::string, std::shared_ptr<IOpBuilder>>
|
||||
CreateOpBuilders() {
|
||||
static std::unordered_map<std::string, std::shared_ptr<IOpBuilder>> CreateOpBuilders() {
|
||||
std::unordered_map<std::string, std::shared_ptr<IOpBuilder>> op_map;
|
||||
|
||||
{
|
||||
|
|
@ -2115,6 +2120,14 @@ CreateOpBuilders() {
|
|||
op_map.emplace("Resize", std::make_shared<ResizeOpBuilder>());
|
||||
op_map.emplace("Flatten", std::make_shared<FlattenOpBuilder>());
|
||||
|
||||
ORT_ENFORCE(op_map.size() == GetOpSupportCheckers().size(),
|
||||
"We should have same number of OpBuilder and OpSupportChecker");
|
||||
|
||||
return op_map;
|
||||
}
|
||||
|
||||
const std::unordered_map<std::string, std::shared_ptr<IOpBuilder>>& GetOpBuilders() {
|
||||
static const std::unordered_map<std::string, std::shared_ptr<IOpBuilder>> op_map = CreateOpBuilders();
|
||||
return op_map;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -20,10 +20,10 @@ class IOpBuilder {
|
|||
virtual Status AddToModelBuilder(ModelBuilder& model_builder, const Node& node) const ORT_MUST_USE_RESULT = 0;
|
||||
};
|
||||
|
||||
// Generate a lookup table with IOpBuilder delegates for different onnx operators
|
||||
// Get the lookup table with IOpBuilder delegates for different onnx operators
|
||||
// Note, the lookup table should have same number of entries as the result of CreateOpSupportCheckers()
|
||||
// in op_support_checker.h
|
||||
std::unordered_map<std::string, std::shared_ptr<IOpBuilder>> CreateOpBuilders();
|
||||
const std::unordered_map<std::string, std::shared_ptr<IOpBuilder>>& GetOpBuilders();
|
||||
|
||||
// Transpose the NHWC input to NCHW output
|
||||
Status TransposeNHWCToNCHW(ModelBuilder& model_builder, const std::string& input, const std::string& output)
|
||||
|
|
|
|||
|
|
@ -16,13 +16,13 @@ using std::vector;
|
|||
|
||||
#pragma region helpers
|
||||
|
||||
bool HasExternalInitializer(const InitializerMap& initializers, const Node& node) {
|
||||
bool HasExternalInitializer(const InitializedTensorSet& initializers, const Node& node) {
|
||||
for (const auto* node_arg : node.InputDefs()) {
|
||||
const auto& input_name(node_arg->Name());
|
||||
if (!Contains(initializers, input_name))
|
||||
continue;
|
||||
|
||||
const auto& tensor = initializers.at(input_name);
|
||||
const auto& tensor = *initializers.at(input_name);
|
||||
if (tensor.has_data_location() &&
|
||||
tensor.data_location() == ONNX_NAMESPACE::TensorProto_DataLocation_EXTERNAL) {
|
||||
LOGS_DEFAULT(VERBOSE) << "Initializer [" << input_name
|
||||
|
|
@ -41,16 +41,16 @@ bool HasExternalInitializer(const InitializerMap& initializers, const Node& node
|
|||
class BaseOpSupportChecker : public IOpSupportChecker {
|
||||
public:
|
||||
virtual ~BaseOpSupportChecker() = default;
|
||||
bool IsOpSupported(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const override;
|
||||
bool IsOpSupported(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const override;
|
||||
|
||||
protected:
|
||||
virtual bool IsOpSupportedImpl(const InitializerMap& /* initializers */, const Node& /* node */,
|
||||
const OPSupportCheckParams& /* params */) const {
|
||||
virtual bool IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, const Node& /* node */,
|
||||
const OpSupportCheckParams& /* params */) const {
|
||||
return true;
|
||||
}
|
||||
|
||||
virtual int32_t GetMinSupportedSdkVer(const Node& /* node */, const OPSupportCheckParams& /* params */) const {
|
||||
virtual int32_t GetMinSupportedSdkVer(const Node& /* node */, const OpSupportCheckParams& /* params */) const {
|
||||
// Android API level 27 is the baseline version of NNAPI,
|
||||
// There is no NNAPI support for Android API level 26-
|
||||
return 27;
|
||||
|
|
@ -63,8 +63,8 @@ class BaseOpSupportChecker : public IOpSupportChecker {
|
|||
bool HasSupportedOpSet(const Node& node) const;
|
||||
};
|
||||
|
||||
bool BaseOpSupportChecker::IsOpSupported(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const {
|
||||
bool BaseOpSupportChecker::IsOpSupported(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const {
|
||||
int32_t required_sdk_ver = GetMinSupportedSdkVer(node, params);
|
||||
if (required_sdk_ver > params.android_sdk_ver) {
|
||||
LOGS_DEFAULT(VERBOSE) << "Current Android API level [" << params.android_sdk_ver
|
||||
|
|
@ -129,15 +129,15 @@ bool BaseOpSupportChecker::HasSupportedOpSet(const Node& node) const {
|
|||
|
||||
class BinaryOpSupportChecker : public BaseOpSupportChecker {
|
||||
private:
|
||||
int32_t GetMinSupportedSdkVer(const Node& node, const OPSupportCheckParams& params) const override;
|
||||
bool IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const override;
|
||||
int32_t GetMinSupportedSdkVer(const Node& node, const OpSupportCheckParams& params) const override;
|
||||
bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const override;
|
||||
bool HasSupportedInputs(const Node& node) const override;
|
||||
int GetMinSupportedOpSet(const Node& node) const override;
|
||||
};
|
||||
|
||||
int32_t BinaryOpSupportChecker::GetMinSupportedSdkVer(
|
||||
const Node& node, const OPSupportCheckParams& /* params */) const {
|
||||
const Node& node, const OpSupportCheckParams& /* params */) const {
|
||||
const auto& op(node.OpType());
|
||||
if (op == "Sub" || op == "Div") {
|
||||
return 28;
|
||||
|
|
@ -166,8 +166,8 @@ bool BinaryOpSupportChecker::HasSupportedInputs(const Node& node) const {
|
|||
return true;
|
||||
}
|
||||
|
||||
bool BinaryOpSupportChecker::IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& /* params */) const {
|
||||
bool BinaryOpSupportChecker::IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& /* params */) const {
|
||||
const auto& op_type(node.OpType());
|
||||
const auto input_defs(node.InputDefs());
|
||||
bool op_is_qlinear = op_type == "QLinearAdd";
|
||||
|
|
@ -221,16 +221,16 @@ bool BinaryOpSupportChecker::IsOpSupportedImpl(const InitializerMap& initializer
|
|||
|
||||
class TransposeOpSupportChecker : public BaseOpSupportChecker {
|
||||
private:
|
||||
bool IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const override;
|
||||
bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const override;
|
||||
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OPSupportCheckParams& /* params */) const override {
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OpSupportCheckParams& /* params */) const override {
|
||||
return 28;
|
||||
}
|
||||
};
|
||||
|
||||
bool TransposeOpSupportChecker::IsOpSupportedImpl(const InitializerMap& /* initializers */, const Node& node,
|
||||
const OPSupportCheckParams& /* params */) const {
|
||||
bool TransposeOpSupportChecker::IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, const Node& node,
|
||||
const OpSupportCheckParams& /* params */) const {
|
||||
Shape input_shape;
|
||||
if (!GetShape(*node.InputDefs()[0], input_shape))
|
||||
return false;
|
||||
|
|
@ -251,15 +251,15 @@ bool TransposeOpSupportChecker::IsOpSupportedImpl(const InitializerMap& /* initi
|
|||
|
||||
class ReshapeOpSupportChecker : public BaseOpSupportChecker {
|
||||
private:
|
||||
bool IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const override;
|
||||
bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const override;
|
||||
|
||||
// Reshape opset 4- uses attributes for new shape which we do not support for now
|
||||
int GetMinSupportedOpSet(const Node& /* node */) const override { return 5; }
|
||||
};
|
||||
|
||||
bool ReshapeOpSupportChecker::IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& /* params */) const {
|
||||
bool ReshapeOpSupportChecker::IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& /* params */) const {
|
||||
const auto& perm_name = node.InputDefs()[1]->Name();
|
||||
if (!Contains(initializers, perm_name)) {
|
||||
LOGS_DEFAULT(VERBOSE) << "New shape of reshape must be known";
|
||||
|
|
@ -276,7 +276,7 @@ bool ReshapeOpSupportChecker::IsOpSupportedImpl(const InitializerMap& initialize
|
|||
return false;
|
||||
}
|
||||
|
||||
const auto& shape_tensor = initializers.at(perm_name);
|
||||
const auto& shape_tensor = *initializers.at(perm_name);
|
||||
const int64_t* raw_shape = GetTensorInt64Data(shape_tensor);
|
||||
const auto size = SafeInt<uint32_t>(shape_tensor.dims()[0]);
|
||||
|
||||
|
|
@ -297,15 +297,15 @@ bool ReshapeOpSupportChecker::IsOpSupportedImpl(const InitializerMap& initialize
|
|||
|
||||
class BatchNormalizationOpSupportChecker : public BaseOpSupportChecker {
|
||||
private:
|
||||
bool IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const override;
|
||||
bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const override;
|
||||
|
||||
// BatchNormalization opset 6- has unsupported attributes
|
||||
int GetMinSupportedOpSet(const Node& /* node */) const override { return 7; }
|
||||
};
|
||||
|
||||
bool BatchNormalizationOpSupportChecker::IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& /* params */) const {
|
||||
bool BatchNormalizationOpSupportChecker::IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& /* params */) const {
|
||||
if (node.OutputDefs().size() != 1) {
|
||||
LOGS_DEFAULT(VERBOSE) << "Your onnx model may be in training mode, please export "
|
||||
"it in test mode.";
|
||||
|
|
@ -361,16 +361,16 @@ bool BatchNormalizationOpSupportChecker::IsOpSupportedImpl(const InitializerMap&
|
|||
|
||||
class PoolOpSupportChecker : public BaseOpSupportChecker {
|
||||
private:
|
||||
bool IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const override;
|
||||
bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const override;
|
||||
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OPSupportCheckParams& params) const override {
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OpSupportCheckParams& params) const override {
|
||||
return params.use_nchw ? 29 : 28;
|
||||
}
|
||||
};
|
||||
|
||||
bool PoolOpSupportChecker::IsOpSupportedImpl(const InitializerMap& /* initializers */, const Node& node,
|
||||
const OPSupportCheckParams& /* params */) const {
|
||||
bool PoolOpSupportChecker::IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, const Node& node,
|
||||
const OpSupportCheckParams& /* params */) const {
|
||||
const auto& op_type = node.OpType();
|
||||
Shape input_shape;
|
||||
if (!GetShape(*node.InputDefs()[0], input_shape))
|
||||
|
|
@ -433,10 +433,10 @@ bool PoolOpSupportChecker::IsOpSupportedImpl(const InitializerMap& /* initialize
|
|||
|
||||
class ConvOpSupportChecker : public BaseOpSupportChecker {
|
||||
private:
|
||||
bool IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const override;
|
||||
bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const override;
|
||||
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OPSupportCheckParams& params) const override {
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OpSupportCheckParams& params) const override {
|
||||
return params.use_nchw ? 29 : 28;
|
||||
}
|
||||
|
||||
|
|
@ -454,8 +454,8 @@ bool ConvOpSupportChecker::HasSupportedInputs(const Node& node) const {
|
|||
return true;
|
||||
}
|
||||
|
||||
bool ConvOpSupportChecker::IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const {
|
||||
bool ConvOpSupportChecker::IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const {
|
||||
const auto& op_type = node.OpType();
|
||||
const auto input_defs = node.InputDefs();
|
||||
NodeAttrHelper helper(node);
|
||||
|
|
@ -465,7 +465,7 @@ bool ConvOpSupportChecker::IsOpSupportedImpl(const InitializerMap& initializers,
|
|||
const auto group = helper.Get("group", 1);
|
||||
const auto weight_name = input_defs[w_idx]->Name();
|
||||
if (Contains(initializers, weight_name)) {
|
||||
const auto& tensor = initializers.at(weight_name);
|
||||
const auto& tensor = *initializers.at(weight_name);
|
||||
if (tensor.dims().size() != 4) {
|
||||
LOGS_DEFAULT(VERBOSE) << "Only conv 2d is supported.";
|
||||
return false;
|
||||
|
|
@ -525,10 +525,10 @@ bool ConvOpSupportChecker::IsOpSupportedImpl(const InitializerMap& initializers,
|
|||
|
||||
class CastOpSupportChecker : public BaseOpSupportChecker {
|
||||
private:
|
||||
bool IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const override;
|
||||
bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const override;
|
||||
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OPSupportCheckParams& /* params */) const override {
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OpSupportCheckParams& /* params */) const override {
|
||||
return 29;
|
||||
}
|
||||
|
||||
|
|
@ -536,8 +536,8 @@ class CastOpSupportChecker : public BaseOpSupportChecker {
|
|||
int GetMinSupportedOpSet(const Node& /* node */) const override { return 6; }
|
||||
};
|
||||
|
||||
bool CastOpSupportChecker::IsOpSupportedImpl(const InitializerMap& /* initializers */, const Node& node,
|
||||
const OPSupportCheckParams& /* params */) const {
|
||||
bool CastOpSupportChecker::IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, const Node& node,
|
||||
const OpSupportCheckParams& /* params */) const {
|
||||
NodeAttrHelper helper(node);
|
||||
const auto to = helper.Get("to", 0);
|
||||
if (to != ONNX_NAMESPACE::TensorProto::FLOAT &&
|
||||
|
|
@ -555,16 +555,16 @@ bool CastOpSupportChecker::IsOpSupportedImpl(const InitializerMap& /* initialize
|
|||
|
||||
class SoftMaxOpSupportChecker : public BaseOpSupportChecker {
|
||||
private:
|
||||
bool IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const override;
|
||||
bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const override;
|
||||
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OPSupportCheckParams& /* params */) const override {
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OpSupportCheckParams& /* params */) const override {
|
||||
return 28;
|
||||
}
|
||||
};
|
||||
|
||||
bool SoftMaxOpSupportChecker::IsOpSupportedImpl(const InitializerMap& /* initializers */, const Node& node,
|
||||
const OPSupportCheckParams& params) const {
|
||||
bool SoftMaxOpSupportChecker::IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, const Node& node,
|
||||
const OpSupportCheckParams& params) const {
|
||||
Shape input_shape;
|
||||
if (!GetShape(*node.InputDefs()[0], input_shape))
|
||||
return false;
|
||||
|
|
@ -596,8 +596,8 @@ bool SoftMaxOpSupportChecker::IsOpSupportedImpl(const InitializerMap& /* initial
|
|||
|
||||
class GemmOpSupportChecker : public BaseOpSupportChecker {
|
||||
private:
|
||||
bool IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const override;
|
||||
bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const override;
|
||||
bool HasSupportedInputs(const Node& node) const override;
|
||||
int GetMinSupportedOpSet(const Node& node) const override;
|
||||
};
|
||||
|
|
@ -623,8 +623,8 @@ int GemmOpSupportChecker::GetMinSupportedOpSet(const Node& node) const {
|
|||
return 1;
|
||||
}
|
||||
|
||||
bool GemmOpSupportChecker::IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& /* params */) const {
|
||||
bool GemmOpSupportChecker::IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& /* params */) const {
|
||||
const auto& op_type = node.OpType();
|
||||
const auto input_defs(node.InputDefs());
|
||||
size_t a_idx = 0, b_idx = 1, c_idx = 2; // A*B+C
|
||||
|
|
@ -733,7 +733,7 @@ bool GemmOpSupportChecker::IsOpSupportedImpl(const InitializerMap& initializers,
|
|||
|
||||
class UnaryOpSupportChecker : public BaseOpSupportChecker {
|
||||
private:
|
||||
int32_t GetMinSupportedSdkVer(const Node& node, const OPSupportCheckParams& params) const override;
|
||||
int32_t GetMinSupportedSdkVer(const Node& node, const OpSupportCheckParams& params) const override;
|
||||
|
||||
// All ops except "Sin" opset 5- uses consumed_inputs attribute which is not supported for now
|
||||
// "Sin" op has support from opset 7, return 6 here for all ops
|
||||
|
|
@ -741,7 +741,7 @@ class UnaryOpSupportChecker : public BaseOpSupportChecker {
|
|||
};
|
||||
|
||||
int32_t UnaryOpSupportChecker::GetMinSupportedSdkVer(
|
||||
const Node& node, const OPSupportCheckParams& /* params */) const {
|
||||
const Node& node, const OpSupportCheckParams& /* params */) const {
|
||||
const auto& op(node.OpType());
|
||||
if (op == "Abs" ||
|
||||
op == "Exp" ||
|
||||
|
|
@ -761,12 +761,12 @@ int32_t UnaryOpSupportChecker::GetMinSupportedSdkVer(
|
|||
|
||||
class ConcatOpSupportChecker : public BaseOpSupportChecker {
|
||||
private:
|
||||
bool IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const override;
|
||||
bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const override;
|
||||
};
|
||||
|
||||
bool ConcatOpSupportChecker::IsOpSupportedImpl(const InitializerMap& /* initializers */, const Node& node,
|
||||
const OPSupportCheckParams& /* params */) const {
|
||||
bool ConcatOpSupportChecker::IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, const Node& node,
|
||||
const OpSupportCheckParams& /* params */) const {
|
||||
Shape input_shape;
|
||||
if (!GetShape(*node.InputDefs()[0], input_shape))
|
||||
return false;
|
||||
|
|
@ -787,10 +787,10 @@ bool ConcatOpSupportChecker::IsOpSupportedImpl(const InitializerMap& /* initiali
|
|||
|
||||
class SqueezeOpSupportChecker : public BaseOpSupportChecker {
|
||||
private:
|
||||
bool IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const override;
|
||||
bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const override;
|
||||
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OPSupportCheckParams& /* params */) const override {
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OpSupportCheckParams& /* params */) const override {
|
||||
return 28;
|
||||
}
|
||||
|
||||
|
|
@ -799,8 +799,8 @@ class SqueezeOpSupportChecker : public BaseOpSupportChecker {
|
|||
int GetMaxSupportedOpSet(const Node& /* node */) const override { return 12; }
|
||||
};
|
||||
|
||||
bool SqueezeOpSupportChecker::IsOpSupportedImpl(const InitializerMap& /* initializers */, const Node& node,
|
||||
const OPSupportCheckParams& /* params */) const {
|
||||
bool SqueezeOpSupportChecker::IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, const Node& node,
|
||||
const OpSupportCheckParams& /* params */) const {
|
||||
Shape input_shape;
|
||||
if (!GetShape(*node.InputDefs()[0], input_shape))
|
||||
return false;
|
||||
|
|
@ -821,16 +821,16 @@ bool SqueezeOpSupportChecker::IsOpSupportedImpl(const InitializerMap& /* initial
|
|||
|
||||
class QuantizeLinearOpSupportChecker : public BaseOpSupportChecker {
|
||||
private:
|
||||
bool IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const override;
|
||||
bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const override;
|
||||
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OPSupportCheckParams& /* params */) const override {
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OpSupportCheckParams& /* params */) const override {
|
||||
return 27;
|
||||
}
|
||||
};
|
||||
|
||||
bool QuantizeLinearOpSupportChecker::IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& /* params */) const {
|
||||
bool QuantizeLinearOpSupportChecker::IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& /* params */) const {
|
||||
const auto input_defs(node.InputDefs());
|
||||
const auto output_defs(node.OutputDefs());
|
||||
|
||||
|
|
@ -862,17 +862,17 @@ bool QuantizeLinearOpSupportChecker::IsOpSupportedImpl(const InitializerMap& ini
|
|||
|
||||
class DequantizeLinearOpSupportChecker : public BaseOpSupportChecker {
|
||||
private:
|
||||
bool IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const override;
|
||||
bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const override;
|
||||
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OPSupportCheckParams& /* params */) const override {
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OpSupportCheckParams& /* params */) const override {
|
||||
return 29;
|
||||
}
|
||||
bool HasSupportedInputs(const Node& node) const override;
|
||||
};
|
||||
|
||||
bool DequantizeLinearOpSupportChecker::IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& /* params */) const {
|
||||
bool DequantizeLinearOpSupportChecker::IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& /* params */) const {
|
||||
const auto input_defs(node.InputDefs());
|
||||
if (!HasValidQuantizationScales(initializers, node, {1}))
|
||||
return false;
|
||||
|
|
@ -906,16 +906,16 @@ bool DequantizeLinearOpSupportChecker::HasSupportedInputs(const Node& node) cons
|
|||
|
||||
class LRNOpSupportChecker : public BaseOpSupportChecker {
|
||||
private:
|
||||
bool IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const override;
|
||||
bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const override;
|
||||
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OPSupportCheckParams& /* params */) const override {
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OpSupportCheckParams& /* params */) const override {
|
||||
return 28;
|
||||
}
|
||||
};
|
||||
|
||||
bool LRNOpSupportChecker::IsOpSupportedImpl(const InitializerMap& /* initializers */, const Node& node,
|
||||
const OPSupportCheckParams& /* params */) const {
|
||||
bool LRNOpSupportChecker::IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, const Node& node,
|
||||
const OpSupportCheckParams& /* params */) const {
|
||||
Shape input_shape;
|
||||
if (!GetShape(*node.InputDefs()[0], input_shape))
|
||||
return false;
|
||||
|
|
@ -936,12 +936,12 @@ bool LRNOpSupportChecker::IsOpSupportedImpl(const InitializerMap& /* initializer
|
|||
|
||||
class ClipOpSupportChecker : public BaseOpSupportChecker {
|
||||
private:
|
||||
bool IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const override;
|
||||
bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const override;
|
||||
};
|
||||
|
||||
bool ClipOpSupportChecker::IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& /* params */) const {
|
||||
bool ClipOpSupportChecker::IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& /* params */) const {
|
||||
float min, max;
|
||||
if (!GetClipMinMax(initializers, node, min, max))
|
||||
return false;
|
||||
|
|
@ -965,10 +965,10 @@ bool ClipOpSupportChecker::IsOpSupportedImpl(const InitializerMap& initializers,
|
|||
|
||||
class ResizeOpSupportChecker : public BaseOpSupportChecker {
|
||||
private:
|
||||
bool IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const override;
|
||||
bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const override;
|
||||
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OPSupportCheckParams& /* params */) const override {
|
||||
int32_t GetMinSupportedSdkVer(const Node& /* node */, const OpSupportCheckParams& /* params */) const override {
|
||||
return 28;
|
||||
}
|
||||
|
||||
|
|
@ -977,8 +977,8 @@ class ResizeOpSupportChecker : public BaseOpSupportChecker {
|
|||
int GetMinSupportedOpSet(const Node& /* node */) const override { return 11; }
|
||||
};
|
||||
|
||||
bool ResizeOpSupportChecker::IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const {
|
||||
bool ResizeOpSupportChecker::IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const {
|
||||
Shape input_shape;
|
||||
if (!GetShape(*node.InputDefs()[0], input_shape))
|
||||
return false;
|
||||
|
|
@ -1035,7 +1035,7 @@ bool ResizeOpSupportChecker::IsOpSupportedImpl(const InitializerMap& initializer
|
|||
|
||||
// We want to check if the scales or sizes are not trying to resize on N/C channels here
|
||||
if (input_defs.size() == 3) { // we are using scales
|
||||
const auto& scales_tensor = initializers.at(input_defs[2]->Name());
|
||||
const auto& scales_tensor = *initializers.at(input_defs[2]->Name());
|
||||
const float* scales_data = GetTensorFloatData(scales_tensor);
|
||||
float scale_n = scales_data[0];
|
||||
float scale_c = scales_data[1];
|
||||
|
|
@ -1048,7 +1048,7 @@ bool ResizeOpSupportChecker::IsOpSupportedImpl(const InitializerMap& initializer
|
|||
} else {
|
||||
// we are using sizes
|
||||
const auto& sizes_name = input_defs[3]->Name();
|
||||
const auto& sizes_tensor = initializers.at(sizes_name);
|
||||
const auto& sizes_tensor = *initializers.at(sizes_name);
|
||||
const int64_t* sizes_data = GetTensorInt64Data(sizes_tensor);
|
||||
uint32_t size_n = SafeInt<uint32_t>(sizes_data[0]);
|
||||
uint32_t size_c = SafeInt<uint32_t>(sizes_data[1]);
|
||||
|
|
@ -1070,12 +1070,12 @@ bool ResizeOpSupportChecker::IsOpSupportedImpl(const InitializerMap& initializer
|
|||
|
||||
class FlattenOpSupportChecker : public BaseOpSupportChecker {
|
||||
private:
|
||||
bool IsOpSupportedImpl(const InitializerMap& initializers, const Node& node,
|
||||
const OPSupportCheckParams& params) const override;
|
||||
bool IsOpSupportedImpl(const InitializedTensorSet& initializers, const Node& node,
|
||||
const OpSupportCheckParams& params) const override;
|
||||
};
|
||||
|
||||
bool FlattenOpSupportChecker::IsOpSupportedImpl(const InitializerMap& /* initializers */, const Node& node,
|
||||
const OPSupportCheckParams& /* params */) const {
|
||||
bool FlattenOpSupportChecker::IsOpSupportedImpl(const InitializedTensorSet& /* initializers */, const Node& node,
|
||||
const OpSupportCheckParams& /* params */) const {
|
||||
Shape input_shape;
|
||||
if (!GetShape(*node.InputDefs()[0], input_shape))
|
||||
return false;
|
||||
|
|
@ -1101,9 +1101,9 @@ bool FlattenOpSupportChecker::IsOpSupportedImpl(const InitializerMap& /* initial
|
|||
|
||||
#pragma endregion
|
||||
|
||||
#pragma region CreateOpSupportCheckers
|
||||
#pragma region CreateGetOpSupportCheckers
|
||||
|
||||
std::unordered_map<std::string, std::shared_ptr<IOpSupportChecker>> CreateOpSupportCheckers() {
|
||||
static std::unordered_map<std::string, std::shared_ptr<IOpSupportChecker>> CreateOpSupportCheckers() {
|
||||
std::unordered_map<std::string, std::shared_ptr<IOpSupportChecker>> op_map;
|
||||
|
||||
// If an OP is always supported, we use BaseOpSupportChecker as default
|
||||
|
|
@ -1174,6 +1174,11 @@ std::unordered_map<std::string, std::shared_ptr<IOpSupportChecker>> CreateOpSupp
|
|||
return op_map;
|
||||
}
|
||||
|
||||
const std::unordered_map<std::string, std::shared_ptr<IOpSupportChecker>>& GetOpSupportCheckers() {
|
||||
static std::unordered_map<std::string, std::shared_ptr<IOpSupportChecker>> op_map = CreateOpSupportCheckers();
|
||||
return op_map;
|
||||
}
|
||||
|
||||
#pragma endregion
|
||||
|
||||
} // namespace nnapi
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@
|
|||
namespace onnxruntime {
|
||||
namespace nnapi {
|
||||
|
||||
struct OPSupportCheckParams {
|
||||
struct OpSupportCheckParams {
|
||||
int32_t android_sdk_ver = 0;
|
||||
bool use_nchw = false;
|
||||
};
|
||||
|
|
@ -18,14 +18,13 @@ class IOpSupportChecker {
|
|||
virtual ~IOpSupportChecker() = default;
|
||||
|
||||
// Check if an operator is supported
|
||||
virtual bool IsOpSupported(const std::unordered_map<std::string, const ONNX_NAMESPACE::TensorProto&>& initializers,
|
||||
const Node& node, const OPSupportCheckParams& params) const = 0;
|
||||
virtual bool IsOpSupported(const InitializedTensorSet& initializers, const Node& node, const OpSupportCheckParams& params) const = 0;
|
||||
};
|
||||
|
||||
// Generate a lookup table with IOpSupportChecker delegates for different onnx operators
|
||||
// Get the lookup table with IOpSupportChecker delegates for different onnx operators
|
||||
// Note, the lookup table should have same number of entries as the result of CreateOpBuilders()
|
||||
// in op_builder.h
|
||||
std::unordered_map<std::string, std::shared_ptr<IOpSupportChecker>> CreateOpSupportCheckers();
|
||||
const std::unordered_map<std::string, std::shared_ptr<IOpSupportChecker>>& GetOpSupportCheckers();
|
||||
|
||||
} // namespace nnapi
|
||||
} // namespace onnxruntime
|
||||
|
|
@ -3,11 +3,12 @@
|
|||
|
||||
#include "nnapi_execution_provider.h"
|
||||
|
||||
#include "builders/model_builder.h"
|
||||
#include "model.h"
|
||||
#include "builders/helper.h"
|
||||
#include "builders/model_builder.h"
|
||||
#include "builders/op_support_checker.h"
|
||||
#include "core/framework/allocatormgr.h"
|
||||
#include "core/framework/compute_capability.h"
|
||||
#include "core/graph/model.h"
|
||||
#include "core/session/onnxruntime_cxx_api.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
|
|
@ -53,7 +54,11 @@ NnapiExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph_view
|
|||
}
|
||||
|
||||
nnapi::ModelBuilder builder(graph_view);
|
||||
const auto supported_nodes_vector = builder.GetSupportedNodes();
|
||||
nnapi::OpSupportCheckParams params{
|
||||
builder.GetAndroidSdkVer(),
|
||||
!!(nnapi_flags_ & NNAPI_FLAG_USE_NCHW),
|
||||
};
|
||||
const auto supported_nodes_vector = GetSupportedNodes(graph_view, params);
|
||||
|
||||
// Find inputs, initializers and outputs for each supported subgraph
|
||||
const std::vector<NodeIndex>& node_index = graph_view.GetNodesInTopologicalOrder();
|
||||
|
|
|
|||
|
|
@ -4,10 +4,12 @@
|
|||
#pragma once
|
||||
|
||||
#include "core/framework/execution_provider.h"
|
||||
#include "core/providers/nnapi/nnapi_builtin/model.h"
|
||||
#include "core/providers/nnapi/nnapi_provider_factory.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace nnapi {
|
||||
class Model;
|
||||
}
|
||||
|
||||
class NnapiExecutionProvider : public IExecutionProvider {
|
||||
public:
|
||||
|
|
|
|||
Loading…
Reference in a new issue