From bbb5862b9e2144185a71543eb21b1bcfda60d5c2 Mon Sep 17 00:00:00 2001 From: Chi Lo Date: Tue, 14 Jan 2025 15:11:43 -0800 Subject: [PATCH] constant fold initializer for QD --- .../onnxruntime_session_options_config_keys.h | 6 +++ .../core/optimizer/constant_folding.cc | 23 +++++++++- onnxruntime/core/optimizer/constant_folding.h | 2 + .../core/optimizer/graph_transformer_utils.cc | 6 ++- onnxruntime/test/optimizer/cse_test.cc | 5 +- .../test/optimizer/graph_transform_test.cc | 46 +++++++++++++++---- 6 files changed, 76 insertions(+), 12 deletions(-) diff --git a/include/onnxruntime/core/session/onnxruntime_session_options_config_keys.h b/include/onnxruntime/core/session/onnxruntime_session_options_config_keys.h index 64a4dd19c1..d6bf493f0c 100644 --- a/include/onnxruntime/core/session/onnxruntime_session_options_config_keys.h +++ b/include/onnxruntime/core/session/onnxruntime_session_options_config_keys.h @@ -300,3 +300,9 @@ static const char* const kOrtSessionOptionsQDQMatMulNBitsAccuracyLevel = "sessio // “Default”: OS determines the scheduling priority and processor performance to service this workload. [Default] // “Efficient”: OS treats this workload is efficiency oriented with low scheduling priority and efficient processor performance. static const char* const kOrtEpDynamicOptionsWorkloadType = "ep.dynamic.workload_type"; + +// Dequantize initializer using ORT ConstantFolding optimizer for dq node if initializer has specific(? TBD) data type. +// This feature is required by some NPU's. +// "0": disable. ORT doesn't constant fold the DQ node. [DEFAULT] +// "1": enable. ORT constant folds the DQ node. +static const char* const kOrtSessionOptionsDequantizeInitializerForDQNode = "session.dequantize_initializer_for_dq_node"; diff --git a/onnxruntime/core/optimizer/constant_folding.cc b/onnxruntime/core/optimizer/constant_folding.cc index e755b4bfa6..0b12199642 100644 --- a/onnxruntime/core/optimizer/constant_folding.cc +++ b/onnxruntime/core/optimizer/constant_folding.cc @@ -18,11 +18,13 @@ namespace onnxruntime { ConstantFolding::ConstantFolding(const IExecutionProvider& execution_provider, bool skip_dequantize_linear, + bool dequantize_initializer_for_dequantize_linear, const ConfigOptions& config_options, const InlinedHashSet& compatible_execution_providers, const InlinedHashSet& excluded_initializers) noexcept : GraphTransformer("ConstantFolding", compatible_execution_providers), skip_dequantize_linear_(skip_dequantize_linear), + dequantize_initializer_for_dequantize_linear_(dequantize_initializer_for_dequantize_linear), config_options_(config_options), excluded_initializers_(excluded_initializers), execution_provider_(execution_provider) { @@ -219,7 +221,26 @@ Status ConstantFolding::ApplyImpl(Graph& graph, bool& modified, int graph_level, } } - if (!can_constant_fold_qdq_node_unit) { + bool can_constant_fold_dq_node = false; + + // Another scenario where dequantizing initializer (ex: initializer -> DQ -> bias of X) can be constant folded is if: + // - the DQ node does not produce a graph output + // - The data type of initializer is INT32, UINT16 or INT16 (More? How can user specify the data type?) + // Note: Some NPU's only support FP32, FP16, INT8, UNIT8 or INT4 weights if consumed by Q/DQ nodes. + // - Does X need to be Gemm, Conv or LayerNormalization ? + if (!can_constant_fold_qdq_node_unit && dequantize_initializer_for_dequantize_linear_) { + if (!graph.NodeProducesGraphOutput(*node)) { // DQ does not produce graph output + const auto* input_def = node->InputDefs()[0]; // Get NodeArg of the initializer of the DequantizeLinear node; + auto data_type = input_def->TypeAsProto()->tensor_type().elem_type(); + if (data_type == ONNX_NAMESPACE::TensorProto_DataType_INT32 || + data_type == ONNX_NAMESPACE::TensorProto_DataType_INT16 || + data_type == ONNX_NAMESPACE::TensorProto_DataType_UINT16) { + can_constant_fold_dq_node = true; + } + } + } + + if (!can_constant_fold_qdq_node_unit && !can_constant_fold_dq_node) { continue; } } diff --git a/onnxruntime/core/optimizer/constant_folding.h b/onnxruntime/core/optimizer/constant_folding.h index 14eb2a9c5f..c40644392c 100644 --- a/onnxruntime/core/optimizer/constant_folding.h +++ b/onnxruntime/core/optimizer/constant_folding.h @@ -24,6 +24,7 @@ class ConstantFolding : public GraphTransformer { */ ConstantFolding(const IExecutionProvider& execution_provider, bool skip_dequantize_linear, + bool dequantize_initializer_for_dequantize_linear, const ConfigOptions& config_options, const InlinedHashSet& compatible_execution_providers = {}, const InlinedHashSet& excluded_initializers = {}) noexcept; @@ -32,6 +33,7 @@ class ConstantFolding : public GraphTransformer { Status ApplyImpl(Graph& graph, bool& modified, int graph_level, const logging::Logger& logger) const override; bool skip_dequantize_linear_; + bool dequantize_initializer_for_dequantize_linear_; const ConfigOptions& config_options_; const InlinedHashSet excluded_initializers_; const IExecutionProvider& execution_provider_; diff --git a/onnxruntime/core/optimizer/graph_transformer_utils.cc b/onnxruntime/core/optimizer/graph_transformer_utils.cc index 9684394da0..8b8d658f06 100644 --- a/onnxruntime/core/optimizer/graph_transformer_utils.cc +++ b/onnxruntime/core/optimizer/graph_transformer_utils.cc @@ -197,6 +197,8 @@ InlinedVector> GenerateTransformers( InlinedVector> transformers; const bool disable_quant_qdq = session_options.config_options.GetConfigOrDefault(kOrtSessionOptionsDisableQuantQDQ, "0") == "1"; + const bool quantize_initializer_for_dq_node = + session_options.config_options.GetConfigOrDefault(kOrtSessionOptionsDequantizeInitializerForDQNode, "1") == "1"; #ifndef DISABLE_CONTRIB_OPS const InlinedHashSet cpu_ep = {onnxruntime::kCpuExecutionProvider}; const InlinedHashSet cpu_acl_eps = {onnxruntime::kCpuExecutionProvider, @@ -233,7 +235,9 @@ InlinedVector> GenerateTransformers( const InlinedHashSet no_limit_empty_ep_list = {}; transformers.emplace_back(std::make_unique(no_limit_empty_ep_list, excluded_initializers)); transformers.emplace_back(std::make_unique()); - transformers.emplace_back(std::make_unique(cpu_execution_provider, !disable_quant_qdq, + transformers.emplace_back(std::make_unique(cpu_execution_provider, + !disable_quant_qdq, + quantize_initializer_for_dq_node, session_options.config_options)); transformers.emplace_back(std::make_unique()); transformers.emplace_back(std::make_unique()); diff --git a/onnxruntime/test/optimizer/cse_test.cc b/onnxruntime/test/optimizer/cse_test.cc index bad96406df..474e1c235e 100644 --- a/onnxruntime/test/optimizer/cse_test.cc +++ b/onnxruntime/test/optimizer/cse_test.cc @@ -284,7 +284,10 @@ TEST(CseTests, MergeConstants) { TransformerLevel::Level1)); const ConfigOptions empty_config_options; ASSERT_STATUS_OK(graph_transformation_mgr.Register( - std::make_unique(*e.get(), false /*skip_dequantize_linear*/, empty_config_options), + std::make_unique(*e.get(), + false /*skip_dequantize_linear*/, + false /*dequantize_initializer_for_dequantize_linear*/, + empty_config_options), TransformerLevel::Level1)); ASSERT_STATUS_OK(graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level1, DefaultLoggingManager().DefaultLogger())); diff --git a/onnxruntime/test/optimizer/graph_transform_test.cc b/onnxruntime/test/optimizer/graph_transform_test.cc index e069f6ef24..c20dfd10cb 100755 --- a/onnxruntime/test/optimizer/graph_transform_test.cc +++ b/onnxruntime/test/optimizer/graph_transform_test.cc @@ -588,7 +588,10 @@ TEST_F(GraphTransformationTests, ConstantFolding) { onnxruntime::GraphTransformerManager graph_transformation_mgr{5}; const ConfigOptions empty_config_options; ASSERT_STATUS_OK(graph_transformation_mgr.Register( - std::make_unique(*e.get(), false /*skip_dequantize_linear*/, empty_config_options), + std::make_unique(*e.get(), + false /*skip_dequantize_linear*/, + false /*dequantize_initializer_for_dequantize_linear*/, + empty_config_options), TransformerLevel::Level1)); ASSERT_STATUS_OK(graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level1, *logger_)); @@ -609,7 +612,10 @@ TEST_F(GraphTransformationTests, ConstantFoldingNodesOnDifferentEP) { const ConfigOptions empty_config_options; ASSERT_STATUS_OK(graph_transformation_mgr.Register( - std::make_unique(*e.get(), false /*skip_dequantize_linear*/, empty_config_options), + std::make_unique(*e.get(), + false /*skip_dequantize_linear*/, + false /*dequantize_initializer_for_dequantize_linear*/, + empty_config_options), TransformerLevel::Level1)); // assign all nodes to CUDA. the constant folding should override this to perform the constant folding on cpu @@ -639,7 +645,10 @@ TEST_F(GraphTransformationTests, ConstantFoldingUnsupportedFloat16) { onnxruntime::GraphTransformerManager graph_transformation_mgr{5}; const ConfigOptions empty_config_options; ASSERT_STATUS_OK(graph_transformation_mgr.Register( - std::make_unique(*e.get(), false /*skip_dequantize_linear*/, empty_config_options), + std::make_unique(*e.get(), + false /*skip_dequantize_linear*/, + false /*dequantize_initializer_for_dequantize_linear*/, + empty_config_options), TransformerLevel::Level1)); // assign all nodes to CUDA. the constant folding should try folding the node on the CPU and fail, thus leaving the @@ -723,7 +732,10 @@ TEST_F(GraphTransformationTests, ConstantFoldingSubgraph) { onnxruntime::GraphTransformerManager graph_transformation_mgr{5}; const ConfigOptions empty_config_options; ASSERT_STATUS_OK(graph_transformation_mgr.Register( - std::make_unique(*e.get(), false /*skip_dequantize_linear*/, empty_config_options), + std::make_unique(*e.get(), + false /*skip_dequantize_linear*/, + false /*dequantize_initializer_for_dequantize_linear*/, + empty_config_options), TransformerLevel::Level1)); ASSERT_STATUS_OK(graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level1, *logger_)); @@ -752,6 +764,7 @@ TEST_F(GraphTransformationTests, ConstantFoldingWithShapeToInitializer) { ASSERT_STATUS_OK(graph_transformation_mgr.Register( std::make_unique(*e.get(), false /*skip_dequantize_linear*/, + false /*dequantize_initializer_for_dequantize_linear*/, empty_config_options, compatible_eps, excluded_initializers), @@ -781,7 +794,10 @@ TEST_F(GraphTransformationTests, ConstantFoldingWithScalarShapeToInitializer) { std::unique_ptr e = std::make_unique(CPUExecutionProviderInfo()); ASSERT_STATUS_OK(graph_transformation_mgr.Register( - std::make_unique(*e.get(), false /*skip_dequantize_linear*/, empty_config_options, + std::make_unique(*e.get(), + false /*skip_dequantize_linear*/, + false /*dequantize_initializer_for_dequantize_linear*/, + empty_config_options, compatible_eps), TransformerLevel::Level1)); @@ -810,7 +826,10 @@ TEST_F(GraphTransformationTests, ConstantFoldingForOpsWithMissingOptionalInputs) std::unique_ptr e = std::make_unique(CPUExecutionProviderInfo()); ASSERT_STATUS_OK(graph_transformation_mgr.Register( - std::make_unique(*e.get(), false /*skip_dequantize_linear*/, empty_config_options, + std::make_unique(*e.get(), + false /*skip_dequantize_linear*/, + false /*dequantize_initializer_for_dequantize_linear*/, + empty_config_options, compatible_eps), TransformerLevel::Level1)); @@ -988,7 +1007,10 @@ TEST_F(GraphTransformationTests, ConstantFolding_RemoveDanglingInputNodesToConst onnxruntime::GraphTransformerManager graph_transformation_mgr{5}; const ConfigOptions empty_config_options; ASSERT_STATUS_OK(graph_transformation_mgr.Register( - std::make_unique(*e.get(), false /*skip_dequantize_linear*/, empty_config_options), + std::make_unique(*e.get(), + false /*skip_dequantize_linear*/, + false /*dequantize_initializer_for_dequantize_linear*/, + empty_config_options), TransformerLevel::Level1)); ASSERT_STATUS_OK(graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level1, *logger_)); @@ -1012,7 +1034,10 @@ TEST_F(GraphTransformationTests, ConstantFoldingAShapeNodeDeepInTheGraph) { std::unique_ptr e = std::make_unique(CPUExecutionProviderInfo()); ASSERT_STATUS_OK(graph_transformation_mgr.Register( - std::make_unique(*e.get(), false /*skip_dequantize_linear*/, empty_config_options), + std::make_unique(*e.get(), + false /*skip_dequantize_linear*/, + false /*dequantize_initializer_for_dequantize_linear*/, + empty_config_options), TransformerLevel::Level1)); ASSERT_STATUS_OK(graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level1, *logger_)); @@ -1041,7 +1066,10 @@ TEST_F(GraphTransformationTests, ConstantFoldingStringInitializer) { std::unique_ptr e = std::make_unique(CPUExecutionProviderInfo()); ASSERT_STATUS_OK(graph_transformation_mgr.Register( - std::make_unique(*e.get(), false /*skip_dequantize_linear*/, empty_config_options), + std::make_unique(*e.get(), + false /*skip_dequantize_linear*/, + false /*dequantize_initializer_for_dequantize_linear*/, + empty_config_options), TransformerLevel::Level1)); ASSERT_STATUS_OK(graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level1, *logger_));