Cleanup NNAPI code (#5505)

* Cleanup NNAPI code

* Check return of GetNCHWInput
This commit is contained in:
Guoyu Wang 2020-10-15 17:40:10 -07:00 committed by GitHub
parent 6f65e2ad2c
commit b991ee4c69
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
4 changed files with 118 additions and 141 deletions

View file

@ -122,6 +122,7 @@ Status ModelBuilder::Prepare() {
RETURN_STATUS_ON_ERROR(nnapi_->ANeuralNetworksModel_create(&nnapi_model_->model_));
ORT_RETURN_IF_ERROR(GetTargetDevices());
PreprocessInitializers();
PreprocessActivations();
ORT_RETURN_IF_ERROR(RegisterInitializers());
ORT_RETURN_IF_ERROR(RegisterModelInputs());
ORT_RETURN_IF_ERROR(AddOperations());
@ -190,6 +191,28 @@ void ModelBuilder::PreprocessInitializers() {
}
}
void ModelBuilder::PreprocessActivations() {
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]));
const auto& op_type(node->OpType());
if (op_type == "Relu") {
activation_nodes_.emplace(node->Index(), ANEURALNETWORKS_FUSED_RELU);
} else if (op_type == "Clip") { // Relu1 or Relu6
float min, max;
if (!GetClipMinMax(*this, *node, min, max))
continue;
if (min == -1.0f && max == 1.0f) {
activation_nodes_.emplace(node->Index(), ANEURALNETWORKS_FUSED_RELU1);
} else if (min == 0.0f && max == 6.0f) {
activation_nodes_.emplace(node->Index(), ANEURALNETWORKS_FUSED_RELU6);
}
}
}
}
// Help to get all quantized operators' input and the node(s) using the input
std::unordered_map<std::string, vector<const Node*>> GetAllQuantizedOpInputs(const GraphViewer& graph_viewer) {
std::unordered_map<std::string, vector<const Node*>> all_quantized_op_inputs;
@ -554,9 +577,9 @@ int32_t ModelBuilder::FindActivation(const Node& node, const NodeArg& output) {
for (auto it = node.OutputEdgesBegin(), end = node.OutputEdgesEnd(); it != end; ++it) {
const auto& dst_node = it->GetNode();
const auto* dst_input = dst_node.InputDefs()[it->GetDstArgIndex()];
if (dst_node.OpType() == "Relu") {
if (Contains(activation_nodes_, dst_node.Index())) {
if (&output == dst_input) {
fuse_code = ANEURALNETWORKS_FUSED_RELU;
fuse_code = activation_nodes_.at(dst_node.Index());
}
} else {
// if there is any other non-relu node using the output

View file

@ -132,6 +132,9 @@ class ModelBuilder {
std::unordered_map<std::string, const ONNX_NAMESPACE::TensorProto&> initializers_;
std::unordered_set<std::string> skipped_initializers_;
// All activation nodes (Relu, Relu1, Relu6) as a map <NodeIndex, activeation_code>
std::unordered_map<NodeIndex, int32_t> activation_nodes_;
std::unordered_map<std::string, std::shared_ptr<IOpBuilder>> op_builders_;
// Operands in nhwc
@ -157,12 +160,18 @@ class ModelBuilder {
Status Prepare() ORT_MUST_USE_RESULT;
Status GetTargetDevices() ORT_MUST_USE_RESULT;
// Get names of all the initializers
void GetAllInitializers();
// If a NNAPI operation will use initializers directly, we will add the initializers to the skip list
void PreprocessInitializers();
// Preprocess all the activation nodes (Relu/Relu1/Relu6) for easy query later
void PreprocessActivations();
// Copy and process all the initializers to NNAPI model
Status RegisterInitializers() ORT_MUST_USE_RESULT;
Status RegisterModelInputs() ORT_MUST_USE_RESULT;
Status AddOperations() ORT_MUST_USE_RESULT;
Status RegisterModelOutputs() ORT_MUST_USE_RESULT;
// After constructing the NNAPI model, will set the shape inferencing record to the Model
void RegisterModelShaper();
Status SetOperandValue(uint32_t index, Model::NNMemory* memory,

View file

@ -33,50 +33,6 @@ GET_TENSOR_DATA(Int64Data, int64_t, int64_data)
#undef GET_TENSOR_DATA
// TODO, move this to a shared location
#define CASE_UNPACK(TYPE, ELEMENT_TYPE, DATA_SIZE) \
case ONNX_NAMESPACE::TensorProto_DataType::TensorProto_DataType_##TYPE: { \
size_t element_count = initializer.has_raw_data() \
? initializer.raw_data().size() \
: initializer.DATA_SIZE(); \
tensor_byte_size = element_count * sizeof(ELEMENT_TYPE); \
unpacked_tensor.reset(new uint8_t[tensor_byte_size]); \
return onnxruntime::utils::UnpackTensor( \
initializer, \
initializer.has_raw_data() ? initializer.raw_data().data() : nullptr, \
initializer.has_raw_data() ? initializer.raw_data().size() : 0, \
reinterpret_cast<ELEMENT_TYPE*>(unpacked_tensor.get()), element_count); \
break; \
}
static Status UnpackInitializerTensor(const onnx::TensorProto& initializer,
std::unique_ptr<uint8_t[]>& unpacked_tensor,
size_t& tensor_byte_size) ORT_MUST_USE_RESULT;
static Status UnpackInitializerTensor(const onnx::TensorProto& initializer,
std::unique_ptr<uint8_t[]>& unpacked_tensor,
size_t& tensor_byte_size) {
switch (initializer.data_type()) {
CASE_UNPACK(FLOAT, float, float_data_size);
CASE_UNPACK(DOUBLE, double, double_data_size);
CASE_UNPACK(BOOL, bool, int32_data_size);
CASE_UNPACK(INT8, int8_t, int32_data_size);
CASE_UNPACK(INT16, int16_t, int32_data_size);
CASE_UNPACK(INT32, int32_t, int32_data_size);
CASE_UNPACK(INT64, int64_t, int64_data_size);
CASE_UNPACK(UINT8, uint8_t, int32_data_size);
CASE_UNPACK(UINT16, uint16_t, int32_data_size);
CASE_UNPACK(UINT32, uint32_t, uint64_data_size);
CASE_UNPACK(UINT64, uint64_t, uint64_data_size);
CASE_UNPACK(FLOAT16, onnxruntime::MLFloat16, int32_data_size);
CASE_UNPACK(BFLOAT16, onnxruntime::BFloat16, int32_data_size);
default:
break;
}
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"Unsupported type: ", initializer.data_type());
}
#undef CASE_UNPACK
#define ADD_SCALAR_OPERAND(model_builder, input_indices, scalar_value) \
{ \
uint32_t _index = 0; \
@ -127,7 +83,7 @@ Status TransposeBetweenNCHWAndNHWC(ModelBuilder& model_builder,
ORT_RETURN_IF_NOT(!model_builder.UseNCHW(), "model_builder.UseNCHW() is on");
const auto& shaper(model_builder.GetShaper());
ORT_RETURN_IF_NOT(4 == shaper[input].size(),
"TransposeNCHWToNHWC input has to be a 4d tensor, actual dimensions: ", shaper[input].size());
"TransposeBetweenNCHWAndNHWC input has to be a 4d tensor, actual dimensions: ", shaper[input].size());
std::string perm_name;
vector<int32_t> perm;
@ -175,6 +131,28 @@ Status TransposeNCHWToNHWC(ModelBuilder& model_builder,
return TransposeBetweenNCHWAndNHWC(model_builder, input, output, true /* nchw_to_nhwc */);
}
// Convert the input from nchw to nhwc
Status GetNHWCInput(ModelBuilder& model_builder, const Node& node, size_t input_index, std::string& input) {
const auto& nchw_input = node.InputDefs()[input_index]->Name();
ORT_RETURN_IF(model_builder.IsOperandNHWC(input));
if (!model_builder.GetNHWCOperand(nchw_input, input)) {
input = model_builder.GetUniqueName(nchw_input + "_nchw_to_nhwc");
ORT_RETURN_IF_ERROR(TransposeNCHWToNHWC(model_builder, nchw_input, input));
}
return Status::OK();
}
// Convert the input from nhwc to nchw
Status GetNCHWInput(ModelBuilder& model_builder, const Node& node, size_t input_index, std::string& input) {
const auto& nhwc_input = node.InputDefs()[input_index]->Name();
ORT_RETURN_IF_NOT(model_builder.IsOperandNHWC(input));
if (!model_builder.GetNCHWOperand(nhwc_input, input)) {
input = model_builder.GetUniqueName(nhwc_input + "_nhwc_to_nchw");
ORT_RETURN_IF_ERROR(TransposeNHWCToNCHW(model_builder, nhwc_input, input));
}
return Status::OK();
}
static Status AddBinaryOperator(int32_t op_type,
ModelBuilder& model_builder,
const std::string& input1,
@ -268,7 +246,7 @@ static Status AddInitializerInNewLayout(ModelBuilder& model_builder,
case ONNX_NAMESPACE::TensorProto_DataType_UINT8:
case ONNX_NAMESPACE::TensorProto_DataType_INT8: {
ORT_RETURN_IF_ERROR(
UnpackInitializerTensor(tensor, unpacked_tensor, tensor_byte_size));
onnxruntime::utils::UnpackInitializerData(tensor, unpacked_tensor, tensor_byte_size));
src = unpacked_tensor.get();
break;
}
@ -348,7 +326,7 @@ static Status AddInitializerTransposed(ModelBuilder& model_builder,
case ONNX_NAMESPACE::TensorProto_DataType_UINT8:
case ONNX_NAMESPACE::TensorProto_DataType_INT8: {
ORT_RETURN_IF_ERROR(
UnpackInitializerTensor(tensor, unpacked_tensor, tensor_byte_size));
onnxruntime::utils::UnpackInitializerData(tensor, unpacked_tensor, tensor_byte_size));
src = unpacked_tensor.get();
break;
}
@ -519,7 +497,7 @@ static Status GetQuantizationZeroPoint(const ModelBuilder& model_builder, const
size_t tensor_byte_size;
const auto& zero_point_tensor = model_builder.GetInitializerTensors().at(node.InputDefs()[idx]->Name());
ORT_RETURN_IF_ERROR(
UnpackInitializerTensor(zero_point_tensor, unpacked_tensor, tensor_byte_size));
onnxruntime::utils::UnpackInitializerData(zero_point_tensor, unpacked_tensor, tensor_byte_size));
zero_point = static_cast<int32_t>(unpacked_tensor.get()[0]);
return Status::OK();
}
@ -647,6 +625,38 @@ Status GetQuantizedInputScaleAndZeroPoint(const ModelBuilder& model_builder,
return Status::OK();
}
bool GetClipMinMax(const ModelBuilder& model_builder, 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
NodeAttrHelper helper(node);
min = helper.Get("min", std::numeric_limits<float>::lowest());
max = helper.Get("max", std::numeric_limits<float>::max());
} else {
const auto& initializers(model_builder.GetInitializerTensors());
if (node.InputDefs().size() > 1) { // we have input min
const auto& min_name = node.InputDefs()[1]->Name();
if (!Contains(initializers, min_name)) {
LOGS_DEFAULT(VERBOSE) << "Input min of Clip must be known";
return false;
}
min = GetTensorFloatData(initializers.at(min_name))[0];
}
if (node.InputDefs().size() > 2) { // we have input max
const auto& max_name = node.InputDefs()[2]->Name();
if (!Contains(initializers, max_name)) {
LOGS_DEFAULT(VERBOSE) << "Input max of Clip must be known";
return false;
}
max = GetTensorFloatData(initializers.at(max_name))[0];
}
}
return true;
}
#pragma endregion helpers
#pragma region op_base
@ -877,18 +887,10 @@ Status BinaryOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const
output_is_nhwc = input1_is_nhwc;
} else if (input1_is_nhwc) {
// need transpsoe input1 back to nchw
const auto& nhwc_input = input_defs[a_idx]->Name();
if (!model_builder.GetNCHWOperand(nhwc_input, input1)) {
input1 = model_builder.GetUniqueName(nhwc_input + "_nhwc_to_nchw");
ORT_RETURN_IF_ERROR(TransposeNHWCToNCHW(model_builder, nhwc_input, input1));
}
ORT_RETURN_IF_ERROR(GetNCHWInput(model_builder, node, a_idx, input1));
} else { // input2_is_nhwc
// need transpsoe input2 back to nchw
const auto& nhwc_input = input_defs[b_idx]->Name();
if (!model_builder.GetNCHWOperand(nhwc_input, input2)) {
input2 = model_builder.GetUniqueName(nhwc_input + "_nhwc_to_nchw");
ORT_RETURN_IF_ERROR(TransposeNHWCToNCHW(model_builder, nhwc_input, input2));
}
ORT_RETURN_IF_ERROR(GetNCHWInput(model_builder, node, b_idx, input2));
}
float a_scale = 0.0f,
@ -1075,11 +1077,7 @@ Status ReshapeOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, cons
auto input = node.InputDefs()[0]->Name();
if (model_builder.IsOperandNHWC(input)) {
// We want to transpose nhwc operand back to nchw before reshape
const auto& nhwc_input = node.InputDefs()[0]->Name();
if (!model_builder.GetNCHWOperand(nhwc_input, input)) {
input = model_builder.GetUniqueName(nhwc_input + "_nhwc_to_nchw");
ORT_RETURN_IF_ERROR(TransposeNHWCToNCHW(model_builder, nhwc_input, input));
}
ORT_RETURN_IF_ERROR(GetNCHWInput(model_builder, node, 0, input));
}
const auto& output = node.OutputDefs()[0]->Name();
@ -1353,11 +1351,7 @@ Status PoolOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const N
} else {
output_is_nhwc = true;
if (!input_is_nhwc) {
const auto& nchw_input = node.InputDefs()[0]->Name();
if (!model_builder.GetNHWCOperand(nchw_input, input)) {
input = model_builder.GetUniqueName(nchw_input + "_nchw_to_nhwc");
ORT_RETURN_IF_ERROR(TransposeNCHWToNHWC(model_builder, nchw_input, input));
}
ORT_RETURN_IF_ERROR(GetNHWCInput(model_builder, node, 0, input));
}
}
@ -1581,11 +1575,7 @@ Status ConvOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const N
} else {
output_is_nhwc = true;
if (!input_is_nhwc) {
const auto& nchw_input = input_defs[x_idx]->Name();
if (!model_builder.GetNHWCOperand(nchw_input, input)) {
input = model_builder.GetUniqueName(nchw_input + "_nchw_to_nhwc");
ORT_RETURN_IF_ERROR(TransposeNCHWToNHWC(model_builder, nchw_input, input));
}
ORT_RETURN_IF_ERROR(GetNHWCInput(model_builder, node, x_idx, input));
}
}
@ -1881,11 +1871,7 @@ Status SoftMaxOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, cons
if (model_builder.IsOperandNHWC(input)) {
output_is_nhwc = false;
// We want to transpose nhwc operand back to nchw before softmax
const auto& nhwc_input = node.InputDefs()[0]->Name();
if (!model_builder.GetNCHWOperand(nhwc_input, input)) {
input = model_builder.GetUniqueName(nhwc_input + "_nhwc_to_nchw");
ORT_RETURN_IF_ERROR(TransposeNHWCToNCHW(model_builder, nhwc_input, input));
}
ORT_RETURN_IF_ERROR(GetNCHWInput(model_builder, node, 0, input));
}
}
@ -2313,20 +2299,15 @@ Status ConcatOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const
for (size_t i = 0; i < node_input_size; i++) {
auto input = node.InputDefs()[i]->Name();
if (model_builder.IsOperandNHWC(input)) {
std::string nhwc_input = input;
input = model_builder.GetUniqueName(input + "_nhwc_to_nchw");
ORT_RETURN_IF_ERROR(TransposeNHWCToNCHW(model_builder, nhwc_input, input));
ORT_RETURN_IF_ERROR(GetNCHWInput(model_builder, node, i, input));
}
input_indices.push_back(operand_indices.at(input));
inputs.push_back(input);
}
}
int32_t axis = helper.Get("axis", 1);
int rank = shaper[input0].size();
if (axis < 0) { // NNAPI does not support negative axis
axis = rank + axis;
}
int32_t axis = SafeInt<int32_t>(HandleNegativeAxis(helper.Get("axis", 1), rank));
if (output_is_nhwc) {
ORT_RETURN_IF_NOT(rank == 4,
@ -2383,11 +2364,7 @@ Status SqueezeOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, cons
auto input = node.InputDefs()[0]->Name();
if (model_builder.IsOperandNHWC(input)) {
// We want to transpose nhwc operand back to nchw before squeeze
const auto& nhwc_input = node.InputDefs()[0]->Name();
if (!model_builder.GetNCHWOperand(nhwc_input, input)) {
input = model_builder.GetUniqueName(nhwc_input + "_nhwc_to_nchw");
ORT_RETURN_IF_ERROR(TransposeNHWCToNCHW(model_builder, nhwc_input, input));
}
ORT_RETURN_IF_ERROR(GetNCHWInput(model_builder, node, 0, input));
}
NodeAttrHelper helper(node);
@ -2395,8 +2372,7 @@ Status SqueezeOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, cons
const auto& input_shape(shaper[input]);
auto input_dims = input_shape.size();
for (auto& axis : axes) {
if (axis < 0)
axis += input_dims;
axis = SafeInt<int32_t>(HandleNegativeAxis(axis, input_dims));
}
if (axes.empty()) { // Squeeze all
@ -2631,11 +2607,7 @@ Status LRNOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const No
// on android api level 28, we need to transpose the nchw input to nhwc
output_is_nhwc = true;
if (!model_builder.IsOperandNHWC(input)) {
const auto& nchw_input = node.InputDefs()[0]->Name();
if (!model_builder.GetNHWCOperand(nchw_input, input)) {
input = model_builder.GetUniqueName(nchw_input + "_nchw_to_nhwc");
ORT_RETURN_IF_ERROR(TransposeNCHWToNHWC(model_builder, nchw_input, input));
}
ORT_RETURN_IF_ERROR(GetNHWCInput(model_builder, node, 0, input));
}
}
@ -2681,7 +2653,6 @@ class ClipOpBuilder : public BaseOpBuilder {
private:
bool IsOpSupportedImpl(ModelBuilder& model_builder, const Node& node) override;
Status AddToModelBuilderImpl(ModelBuilder& model_builder, const Node& node) override ORT_MUST_USE_RESULT;
static bool GetMinMax(ModelBuilder& model_builder, const Node& node, float& min, float& max);
};
void ClipOpBuilder::AddInitializersToSkip(ModelBuilder& model_builder, const Node& node) {
@ -2692,40 +2663,9 @@ void ClipOpBuilder::AddInitializersToSkip(ModelBuilder& model_builder, const Nod
model_builder.AddInitializerToSkip(node.InputDefs()[2]->Name()); // max
}
/* static */ bool ClipOpBuilder::GetMinMax(ModelBuilder& model_builder, const Node& node, float& min, float& max) {
if (node.SinceVersion() < 11) { // Clip opset 1, 6 is using attributes for min/max
NodeAttrHelper helper(node);
min = helper.Get("min", std::numeric_limits<float>::lowest());
max = helper.Get("max", std::numeric_limits<float>::max());
} else {
const auto& initializers(model_builder.GetInitializerTensors());
if (node.InputDefs().size() > 1) { // we have input min
const auto& min_name = node.InputDefs()[1]->Name();
if (!Contains(initializers, min_name)) {
LOGS_DEFAULT(VERBOSE) << "Input min of Clip must be known";
return false;
}
min = GetTensorFloatData(initializers.at(min_name))[0];
}
if (node.InputDefs().size() > 2) { // we have input max
const auto& max_name = node.InputDefs()[2]->Name();
if (!Contains(initializers, max_name)) {
LOGS_DEFAULT(VERBOSE) << "Input max of Clip must be known";
return false;
}
max = GetTensorFloatData(initializers.at(max_name))[0];
}
}
return true;
}
bool ClipOpBuilder::IsOpSupportedImpl(ModelBuilder& model_builder, const Node& node) {
float min = std::numeric_limits<float>::lowest();
float max = std::numeric_limits<float>::max();
if (!GetMinMax(model_builder, node, min, max))
float min, max;
if (!GetClipMinMax(model_builder, node, min, max))
return false;
// We only supoort relu6 or relu1
@ -2753,9 +2693,14 @@ Status ClipOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const N
ORT_RETURN_IF_ERROR(shaper.Identity(input, output));
const OperandType output_operand_type(operand_types.at(input).type, shaper[output]);
float min = std::numeric_limits<float>::lowest();
float max = std::numeric_limits<float>::max();
GetMinMax(model_builder, node, min, max);
if (Contains(model_builder.GetFusedActivations(), input)) {
LOGS_DEFAULT(VERBOSE) << "Clip Node [" << node.Name() << "] fused";
model_builder.RegisterOperand(output, operand_indices.at(input), output_operand_type, output_is_nhwc);
return Status::OK();
}
float min, max;
GetClipMinMax(model_builder, node, min, max);
int32_t op_code;
if (min == 0.0f && max == 6.0f)
@ -2917,11 +2862,7 @@ Status ResizeOpBuilder::AddToModelBuilderImpl(ModelBuilder& model_builder, const
} else {
output_is_nhwc = true;
if (!input_is_nhwc) {
const auto& nchw_input = input_defs[0]->Name();
if (!model_builder.GetNHWCOperand(nchw_input, input)) {
input = model_builder.GetUniqueName(nchw_input + "_nchw_to_nhwc");
ORT_RETURN_IF_ERROR(TransposeNCHWToNHWC(model_builder, nchw_input, input));
}
ORT_RETURN_IF_ERROR(GetNHWCInput(model_builder, node, 0, input));
}
}

View file

@ -36,5 +36,9 @@ Status GetQuantizedInputScaleAndZeroPoint(const ModelBuilder& model_builder,
const Node& node, const std::string& input_name,
float& scale, int32_t& zero_point) ORT_MUST_USE_RESULT;
// 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 ModelBuilder& model_builder, const Node& node, float& min, float& max);
} // namespace nnapi
} // namespace onnxruntime