From 62eab67f796bf6ae693542cacb3f9fa5ac69829e Mon Sep 17 00:00:00 2001 From: Yi-Hong Lyu Date: Wed, 19 Jan 2022 06:47:33 +0800 Subject: [PATCH] Fuse DQ -> ArgMax into ArgMax (#10274) --- .../optimizer/qdq_transformer/qdq_util.cc | 26 +++++++++++++ .../core/optimizer/qdq_transformer/qdq_util.h | 7 ++++ .../qdq_selector_action_transformer.cc | 24 ++++++++++++ .../selectors_actions/qdq_selectors.cc | 24 ++++++++++++ .../selectors_actions/qdq_selectors.h | 15 +++++++ .../test/optimizer/qdq_transformer_test.cc | 39 +++++++++++++++++++ 6 files changed, 135 insertions(+) diff --git a/onnxruntime/core/optimizer/qdq_transformer/qdq_util.cc b/onnxruntime/core/optimizer/qdq_transformer/qdq_util.cc index 14ed791cb0..0c5710b6fd 100644 --- a/onnxruntime/core/optimizer/qdq_transformer/qdq_util.cc +++ b/onnxruntime/core/optimizer/qdq_transformer/qdq_util.cc @@ -58,5 +58,31 @@ bool IsQDQPairSupported( *q_scale.data() == *dq_scale.data(); } +bool IsDQSupported( + const Node& dq_node, + const std::function& get_const_initializer) { + ConstPointerContainer> dq_input_defs = dq_node.InputDefs(); + + // DQ contains optional input is not supported + // non-scalar DQ scale and zero point needs are not supported + if (dq_input_defs.size() != InputIndex::TOTAL_COUNT || + !optimizer_utils::IsScalar(*dq_input_defs[InputIndex::SCALE_ID]) || + !optimizer_utils::IsScalar(*dq_input_defs[InputIndex::ZERO_POINT_ID])) { + return false; + } + + // if DQ scale and zero point are not constant, return false + const ONNX_NAMESPACE::TensorProto* dq_scale_tensor_proto = + get_const_initializer(dq_input_defs[InputIndex::SCALE_ID]->Name()); + const ONNX_NAMESPACE::TensorProto* dq_zp_tensor_proto = + get_const_initializer(dq_input_defs[InputIndex::ZERO_POINT_ID]->Name()); + if (nullptr == dq_zp_tensor_proto || + nullptr == dq_scale_tensor_proto) { + return false; + } + + return true; +} + } // namespace QDQ } // namespace onnxruntime diff --git a/onnxruntime/core/optimizer/qdq_transformer/qdq_util.h b/onnxruntime/core/optimizer/qdq_transformer/qdq_util.h index a6690cd782..22a06e7506 100644 --- a/onnxruntime/core/optimizer/qdq_transformer/qdq_util.h +++ b/onnxruntime/core/optimizer/qdq_transformer/qdq_util.h @@ -36,5 +36,12 @@ bool IsQDQPairSupported( const std::function& get_const_initializer, const Path& model_path); +// Check if DQ is supported in the QDQ transformer. It requires: +// 1. DQ doesn't have optional input. +// 2. scale and zero point is constant scalar +bool IsDQSupported( + const Node& dq_node, + const std::function& get_const_initializer); + } // namespace QDQ } // namespace onnxruntime diff --git a/onnxruntime/core/optimizer/qdq_transformer/selectors_actions/qdq_selector_action_transformer.cc b/onnxruntime/core/optimizer/qdq_transformer/selectors_actions/qdq_selector_action_transformer.cc index e39a11a1e8..50bc405378 100644 --- a/onnxruntime/core/optimizer/qdq_transformer/selectors_actions/qdq_selector_action_transformer.cc +++ b/onnxruntime/core/optimizer/qdq_transformer/selectors_actions/qdq_selector_action_transformer.cc @@ -45,6 +45,29 @@ void DropQDQNodesRules(SelectorActionRegistry& qdq_selector_action_registry) { #endif } +// create rules for ops that don't change the data +void DropDQNodesRules(SelectorActionRegistry& qdq_selector_action_registry) { + // 2 nodes. DQ, target. Merge into target and remove DQ. + const std::string action_name{"dropDQ"}; + NTO::NodeLocation dq{NTO::NodeType::kInput, 0}; + + // Move DQ input 0 to target input 0. + std::vector moves{ + MoveToSlot(dq, ArgType::kInput, 0, ArgType::kInput, 0)}; + + std::unique_ptr action = std::make_unique(std::move(moves)); + +#if !defined(ORT_MINIMAL_BUILD) + std::unique_ptr selector = std::make_unique(); + qdq_selector_action_registry.RegisterSelectorAndAction(action_name, + {{"ArgMax", {}}}, + std::move(selector), + std::move(action)); +#else + qdq_selector_action_registry.RegisterAction(action_name, std::move(action)); +#endif +} + void UnaryOpQDQRules(SelectorActionRegistry& qdq_selector_action_registry) { // 3 nodes. DQ, target, Q // Replace with internal QLinear version of operator. Delete all original nodes. @@ -148,6 +171,7 @@ SelectorActionRegistry CreateSelectorActionRegistry(bool is_int8_allowed) { SelectorActionRegistry qdq_selector_action_registry; DropQDQNodesRules(qdq_selector_action_registry); + DropDQNodesRules(qdq_selector_action_registry); UnaryOpQDQRules(qdq_selector_action_registry); BinaryOpQDQRules(qdq_selector_action_registry); VariadicOpQDQRules(qdq_selector_action_registry); diff --git a/onnxruntime/core/optimizer/qdq_transformer/selectors_actions/qdq_selectors.cc b/onnxruntime/core/optimizer/qdq_transformer/selectors_actions/qdq_selectors.cc index 582d215cfc..56f3015903 100644 --- a/onnxruntime/core/optimizer/qdq_transformer/selectors_actions/qdq_selectors.cc +++ b/onnxruntime/core/optimizer/qdq_transformer/selectors_actions/qdq_selectors.cc @@ -104,6 +104,30 @@ bool DropQDQNodeGroupSelector::Check(const GraphViewer& graph_viewer, return IsQDQPairSupported(q_node, dq_node, get_const_initializer, graph_viewer.ModelPath()); } +bool DropDQNodeGroupSelector::CheckDQNodes(const Node& node, const std::vector& dq_nodes) const { + int num_dq_inputs = NumActualValues(node, true); + + return num_dq_inputs == gsl::narrow_cast(dq_nodes.size()); +} + +bool DropDQNodeGroupSelector::Check(const GraphViewer& graph_viewer, + const Node& node, + const std::vector& dq_nodes, + const std::vector& q_nodes) const { + if (!CheckDQNodes(node, dq_nodes)) { + return false; + } + + (void)q_nodes; + const Node& dq_node = *dq_nodes.front(); + + auto get_const_initializer = [&graph_viewer](const std::string& initializer_name) { + return graph_viewer.GetConstantInitializer(initializer_name, true); + }; + + return IsDQSupported(dq_node, get_const_initializer); +} + bool UnaryNodeGroupSelector::Check(const GraphViewer& graph_viewer, const Node& node, const std::vector& dq_nodes, const std::vector& q_nodes) const { diff --git a/onnxruntime/core/optimizer/qdq_transformer/selectors_actions/qdq_selectors.h b/onnxruntime/core/optimizer/qdq_transformer/selectors_actions/qdq_selectors.h index 75ec972c9a..398cfc1cce 100644 --- a/onnxruntime/core/optimizer/qdq_transformer/selectors_actions/qdq_selectors.h +++ b/onnxruntime/core/optimizer/qdq_transformer/selectors_actions/qdq_selectors.h @@ -56,6 +56,16 @@ class DropQDQNodeGroupSelector : public NodeGroupSelector { const std::vector& q_nodes) const override; }; +// Single DQ -> node. +class DropDQNodeGroupSelector : public NodeGroupSelector { + // base check that we have the expected number of DQ inputs. + bool CheckDQNodes(const Node& node, const std::vector& dq_nodes) const; + + bool Check(const GraphViewer& graph_viewer, const Node& node, + const std::vector& dq_nodes, + const std::vector& q_nodes) const override; +}; + // single input. default is to only support uint8. class UnaryNodeGroupSelector : public NodeGroupSelector { bool Check(const GraphViewer& graph_viewer, const Node& node, @@ -142,6 +152,11 @@ class DropQDQNodesSelector : public BaseSelector { DropQDQNodesSelector() : BaseSelector(std::make_unique()) {} }; +class DropDQNodesSelector : public BaseSelector { + public: + DropDQNodesSelector() : BaseSelector(std::make_unique()) {} +}; + class UnarySelector : public BaseSelector { public: UnarySelector() : BaseSelector(std::make_unique()) {} diff --git a/onnxruntime/test/optimizer/qdq_transformer_test.cc b/onnxruntime/test/optimizer/qdq_transformer_test.cc index 4eea3be246..cbb716bfdd 100644 --- a/onnxruntime/test/optimizer/qdq_transformer_test.cc +++ b/onnxruntime/test/optimizer/qdq_transformer_test.cc @@ -819,6 +819,45 @@ TEST(QDQTransformerTests, ResizeReshape) { test_case({1, 2, 26, 42}, {4}); } +TEST(QDQTransformerTests, ArgMax) { + auto test_case = [&](const std::vector& input_shape, + int axis, + int keepdims, + int select_last_index) { + auto build_test_case = [&](ModelTestBuilder& builder) { + auto* input_arg = builder.MakeInput(input_shape, + std::numeric_limits::min(), + std::numeric_limits::max()); + auto* output_arg = builder.MakeOutput(); + + // add DQ + auto* dq_output = builder.MakeIntermediate(); + builder.AddDequantizeLinearNode(input_arg, .003f, 1, dq_output); + + // add ArgMax + Node& argmax_node = builder.AddNode("ArgMax", {dq_output}, {output_arg}); + argmax_node.AddAttribute("axis", static_cast(axis)); + argmax_node.AddAttribute("keepdims", static_cast(keepdims)); + argmax_node.AddAttribute("select_last_index", static_cast(select_last_index)); + }; + + auto check_argmax_graph = [&](InferenceSessionWrapper& session) { + auto op_to_count = CountOpsInGraph(session.GetGraph()); + EXPECT_EQ(op_to_count["ArgMax"], 1); + EXPECT_EQ(op_to_count["DequantizeLinear"], 0); + }; + + TransformerTester(build_test_case, check_argmax_graph, + TransformerLevel::Level1, + TransformerLevel::Level2, + /* opset_version */ 13); + }; + + test_case({2, 13, 12, 37}, 1, 0, 0); + test_case({2, 13, 12, 37}, 0, 1, 0); + test_case({2, 13, 12, 37}, 0, 0, 1); +} + TEST(QDQTransformerTests, QLinearMatMul) { auto test_case = [&](const std::vector& input1_shape, const std::vector& input2_shape) { auto build_test_case = [&](ModelTestBuilder& builder) {