From 0c10cd43d55f5a8d87ebf8c42602e1bbacf27a57 Mon Sep 17 00:00:00 2001 From: Chi Lo Date: Thu, 6 Feb 2025 15:16:26 -0800 Subject: [PATCH] don't create optimizer instances until EP requests it by calling GetOptimizerByName --- .../optimizer/graph_optimizer_registry.cc | 70 ++++++++++- .../core/optimizer/graph_optimizer_registry.h | 67 ++++++++-- .../constant_folding_dq_node.cc | 11 +- .../constant_folding_dq_node.h | 6 +- .../selection_and_optimization_func.cc | 108 ++++++++++++++++ .../selection_and_optimization_func.h | 20 +++ .../shared_library/provider_interfaces.h | 1 + .../tensorrt/tensorrt_execution_provider.cc | 3 +- onnxruntime/core/session/inference_session.cc | 11 +- .../core/session/provider_bridge_ort.cc | 115 +----------------- 10 files changed, 276 insertions(+), 136 deletions(-) create mode 100644 onnxruntime/core/optimizer/selection_and_optimization_func.cc create mode 100644 onnxruntime/core/optimizer/selection_and_optimization_func.h diff --git a/onnxruntime/core/optimizer/graph_optimizer_registry.cc b/onnxruntime/core/optimizer/graph_optimizer_registry.cc index 2811acdeea..4784b11f3b 100644 --- a/onnxruntime/core/optimizer/graph_optimizer_registry.cc +++ b/onnxruntime/core/optimizer/graph_optimizer_registry.cc @@ -3,23 +3,77 @@ #include "core/optimizer/graph_optimizer_registry.h" #include "core/optimizer/graph_transformer_utils.h" +#include "core/optimizer/selection_and_optimization_func.h" +#include "core/optimizer/qdq_transformer/constant_folding_dq_node.h" using namespace onnxruntime; using namespace ::onnxruntime::common; namespace onnxruntime { +GraphOptimizerRegistry::GraphOptimizerRegistry() { + logger_ = &logging::LoggingManager::DefaultLogger(); +} + common::Status GraphOptimizerRegistry::Register(std::unique_ptr transformer) { const auto& name = transformer->Name(); - if (name_to_transformer_map_.find(name) != name_to_transformer_map_.end()) { - return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "This transformer is already registered " + name); + if (name_to_transformer_map_.find(name) != name_to_transformer_map_.end() && + name_to_transformer_map_.at(name)) { + LOGS(*logger_, WARNING) << "This optimizer is already created and registered " << name; + return Status::OK(); } name_to_transformer_map_[name] = transformer.get(); transformer_list_.push_back(std::move(transformer)); + + if (name == kCONSTANT_FOLDING_DQ) { + transformer_name_to_selection_func_[name] = ConstantFoldingDQ_selection; + } + return Status::OK(); } +common::Status GraphOptimizerRegistry::AddPredefinedOptimizerNames(std::vector& optimizer_names) { + for (auto name : optimizer_names) { + if (name_to_transformer_map_.find(name) != name_to_transformer_map_.end()) { + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "This transformer name is already added " + name); + } + name_to_transformer_map_[name] = nullptr; // The transformer will be instantizted only when EP requests it + } + return Status::OK(); +} + +std::unique_ptr GraphOptimizerRegistry::CreateOptimizer(std::string& name, std::unordered_map& key_value_configs) { + std::unique_ptr transformer; + if (name == kCONSTANT_FOLDING_DQ) { + const InlinedHashSet node_index_set = {}; + return std::make_unique(*cpu_ep_, false /*skip_dequantize_linear*/, + session_options_->config_options, node_index_set); + } + LOGS(*logger_, WARNING) << "Can't create optimizer " << name; + return transformer; +} + +std::optional>(const GraphViewer&)>> GraphOptimizerRegistry::GetSelectionFunc(std::string& name, + std::unordered_map& key_value_configs) const { + if (name_to_transformer_map_.find(name) == name_to_transformer_map_.end()) { + LOGS(*logger_, WARNING) << "Can't find optimizer " << name; + return std::nullopt; + } + + // Create and register if the transformer instance is not created. + if (!name_to_transformer_map_.at(name)) { + auto new_transformer = Get()->CreateOptimizer(name, key_value_configs); + Get()->Register(std::move(new_transformer)); + } + + auto lookup = transformer_name_to_selection_func_.find(name); + if (lookup != transformer_name_to_selection_func_.end()) { + return transformer_name_to_selection_func_.at(name); + } + return std::nullopt; +} + GraphTransformer* GraphOptimizerRegistry::GetTransformerByName(std::string& name) const { if (name_to_transformer_map_.find(name) != name_to_transformer_map_.end()) { return name_to_transformer_map_.at(name); @@ -27,7 +81,7 @@ GraphTransformer* GraphOptimizerRegistry::GetTransformerByName(std::string& name return nullptr; } -// Registers all the predefined transformers for EP +// Create and register all the predefined transformers for EP common::Status GraphOptimizerRegistry::AddPredefinedOptimizers( const onnxruntime::SessionOptions& sess_options, const onnxruntime::IExecutionProvider& cpu_ep, @@ -56,6 +110,16 @@ common::Status GraphOptimizerRegistry::ApplyTransformer(Graph& graph, std::strin return Status::OK(); } +common::Status GraphOptimizerRegistry::AddCpuEpReference(onnxruntime::IExecutionProvider* cpu_ep) { + cpu_ep_ = cpu_ep; + return Status::OK(); +} + +common::Status GraphOptimizerRegistry::AddSessionOptionsReference(onnxruntime::SessionOptions* session_options) { + session_options_ = session_options; + return Status::OK(); +} + // Initialize static members std::shared_ptr onnxruntime::GraphOptimizerRegistry::graph_optimizer_registry = nullptr; std::mutex GraphOptimizerRegistry::registry_mutex; diff --git a/onnxruntime/core/optimizer/graph_optimizer_registry.h b/onnxruntime/core/optimizer/graph_optimizer_registry.h index fc425b50de..99ddc6542d 100644 --- a/onnxruntime/core/optimizer/graph_optimizer_registry.h +++ b/onnxruntime/core/optimizer/graph_optimizer_registry.h @@ -5,16 +5,22 @@ #include "core/common/inlined_containers.h" #include "core/common/logging/logging.h" -//#include "core/common/common.h" #include "core/optimizer/graph_transformer.h" #include "core/framework/execution_providers.h" +#include "core/framework/compute_capability.h" namespace onnxruntime { +/** + * A registration/lookup class for re-usable optimizers for EPs. + */ class GraphOptimizerRegistry { public: - explicit GraphOptimizerRegistry() {} + explicit GraphOptimizerRegistry(); GraphOptimizerRegistry(const GraphOptimizerRegistry&) = delete; + /** + * Get GraphOptimizerRegistry instance as a singleton. + */ static std::shared_ptr Get() { if (!graph_optimizer_registry) { // First Check (without locking) std::lock_guard lock(registry_mutex); @@ -25,21 +31,66 @@ class GraphOptimizerRegistry { return graph_optimizer_registry; } - common::Status AddPredefinedOptimizers(const onnxruntime::SessionOptions& sess_options, - const onnxruntime::IExecutionProvider& cpu_ep, - const logging::Logger& logger); + /** + * Register all the predefined optimizer names, only name not the optimizer instance. + * + * The optimizer will later be instantizted only when EP requests it by calling GetOptimizerByName in provider bridge. + */ + common::Status GraphOptimizerRegistry::AddPredefinedOptimizerNames(std::vector& optimizer_names); + /** + * Create and register all predefined optimizers. + */ + common::Status AddPredefinedOptimizers(const onnxruntime::SessionOptions& sess_options, + const onnxruntime::IExecutionProvider& cpu_ep, + const logging::Logger& logger); + + /** + * Create optimizer instance. + */ + std::unique_ptr CreateOptimizer(std::string& name, std::unordered_map& key_value_configs); + + /** + * Get optimizer by name. + */ + GraphTransformer* GraphOptimizerRegistry::GetTransformerByName(std::string& name) const; + + /** + * Run the optimizer. + */ common::Status ApplyTransformer(Graph& graph, std::string& name, const logging::Logger& logger) const; + /** + * Register optimizer and its optimization selection function. + */ common::Status Register(std::unique_ptr transformer); - - // Get transformer by name. Return nullptr if not found. - GraphTransformer* GetTransformerByName(std::string& name) const; + + /** + * Get optimizer selection function requested by EP. If the optimizer name can't be found, return nullopt. + * + * Please note that this function also creates and registers the optimizer if its instance is not existed. + */ + std::optional>(const GraphViewer&)>> GraphOptimizerRegistry::GetSelectionFunc(std::string& name, + std::unordered_map& key_value_configs) const; + + /** + * Add CPU EP reference from InferenceSession as it's needed for some optimizers, ex: ConstantFoldingDQ. + */ + common::Status AddCpuEpReference(onnxruntime::IExecutionProvider* cpu_ep); + + /** + * Add Session Options reference from InferenceSession as it's needed for some optimizers, ex: ConstantFoldingDQ. + */ + common::Status AddSessionOptionsReference(onnxruntime::SessionOptions* session_options); private: InlinedVector> transformer_list_; InlinedHashMap name_to_transformer_map_; + InlinedHashMap>(const GraphViewer&)>> transformer_name_to_selection_func_; + const logging::Logger* logger_; + onnxruntime::IExecutionProvider* cpu_ep_; + onnxruntime::SessionOptions* session_options_; static std::shared_ptr graph_optimizer_registry; static std::mutex registry_mutex; diff --git a/onnxruntime/core/optimizer/qdq_transformer/constant_folding_dq_node.cc b/onnxruntime/core/optimizer/qdq_transformer/constant_folding_dq_node.cc index 9bb09b542a..afbf68f8bb 100644 --- a/onnxruntime/core/optimizer/qdq_transformer/constant_folding_dq_node.cc +++ b/onnxruntime/core/optimizer/qdq_transformer/constant_folding_dq_node.cc @@ -1,18 +1,9 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. -#include - #include "core/optimizer/qdq_transformer/constant_folding_dq_node.h" -#include "core/optimizer/initializer.h" -#include "core/optimizer/utils.h" +#include "core/optimizer/graph_optimizer_registry.h" #include "core/graph/graph_utils.h" -#include "core/optimizer/optimizer_execution_frame.h" -#include "core/optimizer/utils.h" -#include "core/framework/op_kernel.h" -#include "core/framework/tensorprotoutils.h" - -using namespace onnxruntime::common; namespace onnxruntime { diff --git a/onnxruntime/core/optimizer/qdq_transformer/constant_folding_dq_node.h b/onnxruntime/core/optimizer/qdq_transformer/constant_folding_dq_node.h index 6b2f8aabda..a708f33040 100644 --- a/onnxruntime/core/optimizer/qdq_transformer/constant_folding_dq_node.h +++ b/onnxruntime/core/optimizer/qdq_transformer/constant_folding_dq_node.h @@ -12,15 +12,13 @@ namespace onnxruntime { /** -@class ConstantFolding +@class ConstantFoldingDQ -Transformer that traverses the graph top-down and performs constant folding, i.e., -it statically computes parts of the graph that rely only on constant initializers. +It's the derived class from ConstantFolding. */ class ConstantFoldingDQ : public ConstantFolding { public: /*! Constant folding will not be applied to nodes that have one of initializers from excluded_initializers as input. - For pre-training, the trainable weights are those initializers to be excluded. \param execution_provider Execution provider instance to execute constant folding. */ ConstantFoldingDQ(const IExecutionProvider& execution_provider, diff --git a/onnxruntime/core/optimizer/selection_and_optimization_func.cc b/onnxruntime/core/optimizer/selection_and_optimization_func.cc new file mode 100644 index 0000000000..a3fe45d89f --- /dev/null +++ b/onnxruntime/core/optimizer/selection_and_optimization_func.cc @@ -0,0 +1,108 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "selection_and_optimization_func.h" +#include "core/optimizer/graph_optimizer_registry.h" +#include "core/graph/graph_utils.h" +#include "core/framework/compute_capability.h" +#include "core/optimizer/qdq_transformer/constant_folding_dq_node.h" + +namespace onnxruntime { + +std::vector> ConstantFoldingDQ_selection(const GraphViewer& graph_viewer) { + std::vector> result; + std::unique_ptr sub_graph = std::make_unique(); + const std::vector& node_index = graph_viewer.GetNodesInTopologicalOrder(ExecutionOrder::PRIORITY_BASED /*priority-based topological sort*/); + InitializedTensorSet constant_inputs; + const InlinedHashSet excluded_initializers; + + // Select DequantizeLinear node where all inputs are constant + for (const auto& index : node_index) { + const auto& node = graph_viewer.GetNode(index); + if (node->OpType() != "DequantizeLinear") { + continue; + } + if (!graph_utils::AllNodeInputsAreConstant(graph_viewer.GetGraph(), *node, constant_inputs, excluded_initializers)) { + continue; + } + sub_graph->nodes.push_back(index); + } + + result.push_back(std::make_unique(std::move(sub_graph))); + result.back()->optimization_func = ConstantFoldingDQ_optimization; + return result; +} + +Status ConstantFoldingDQ_optimization(Graph& graph, const ComputeCapability& optimization_cc, ComputeCapability& cc_to_update) { + std::string optimizer_name = kCONSTANT_FOLDING_DQ; + auto logger = const_cast(&logging::LoggingManager::DefaultLogger()); + std::unordered_set original_initializers_to_remove; + std::unordered_set new_initializers_to_add; + InlinedHashSet dq_node_index_set; + + // iterate node_to_optimize to: + // 1. get original initializers to remove + // 2. add new initializers + // 3. create dq node index set + for (const auto& index : optimization_cc.sub_graph->nodes) { + auto node = graph.GetNode(index); + if (node->OpType() != "DequantizeLinear") { + continue; + } + auto input_0 = node->InputDefs()[0]; + auto output_0 = node->OutputDefs()[0]; + original_initializers_to_remove.insert(input_0->Name()); + new_initializers_to_add.insert(output_0->Name()); + dq_node_index_set.insert(index); + } + + auto optimizer_registry = onnxruntime::GraphOptimizerRegistry::Get(); + + ConstantFoldingDQ* transformer = static_cast(optimizer_registry->GetTransformerByName(optimizer_name)); + transformer->UpdateNodeIndexSet(dq_node_index_set); + + // apply constant folding on DQ nodes + optimizer_registry->ApplyTransformer(graph, optimizer_name, *logger); + + // update the overall ComputeCapability + std::vector updated_nodes; + for (auto index : cc_to_update.sub_graph->nodes) { + if (dq_node_index_set.find(index) != dq_node_index_set.end()) { + continue; + } + updated_nodes.push_back(index); + } + cc_to_update.sub_graph->nodes = updated_nodes; + + auto original_meta_def = cc_to_update.sub_graph->GetMetaDef(); + std::unique_ptr updated_meta_def = std::make_unique(); + updated_meta_def->name = original_meta_def->name; + updated_meta_def->domain = original_meta_def->domain; + updated_meta_def->since_version = original_meta_def->since_version; + updated_meta_def->status = original_meta_def->status; + updated_meta_def->inputs = original_meta_def->inputs; + updated_meta_def->outputs = original_meta_def->outputs; + updated_meta_def->attributes = original_meta_def->attributes; + updated_meta_def->doc_string = original_meta_def->doc_string; +#if !defined(ORT_MINIMAL_BUILD) + updated_meta_def->type_and_shape_inference_function = original_meta_def->type_and_shape_inference_function; +#endif + for (auto constant_initializer : original_meta_def->constant_initializers) { + if (original_initializers_to_remove.find(constant_initializer) != original_initializers_to_remove.end()) { + continue; + } + updated_meta_def->constant_initializers.push_back(constant_initializer); + } + + for (auto constant_initializer : new_initializers_to_add) { + updated_meta_def->constant_initializers.push_back(constant_initializer); + } + + cc_to_update.sub_graph->SetMetaDef(std::move(updated_meta_def)); + + return Status::OK(); +} + + + +} // namespace onnxruntime diff --git a/onnxruntime/core/optimizer/selection_and_optimization_func.h b/onnxruntime/core/optimizer/selection_and_optimization_func.h new file mode 100644 index 0000000000..08ce605c3d --- /dev/null +++ b/onnxruntime/core/optimizer/selection_and_optimization_func.h @@ -0,0 +1,20 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#pragma once + +#include "core/framework/compute_capability.h" +#include "core/graph/graph_viewer.h" + +namespace onnxruntime { +static const std::string kCONSTANT_FOLDING_DQ = "ConstantFoldingDQ"; + +// ConstantFoldingDQ selection function +std::vector> ConstantFoldingDQ_selection(const GraphViewer& graph_viewer); + +// ConstantFoldingDQ optimization function +Status ConstantFoldingDQ_optimization(Graph& graph, const ComputeCapability& optimization_cc, ComputeCapability& cc_to_update); + + + +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/shared_library/provider_interfaces.h b/onnxruntime/core/providers/shared_library/provider_interfaces.h index 16fc15ea76..ec6cd85b42 100644 --- a/onnxruntime/core/providers/shared_library/provider_interfaces.h +++ b/onnxruntime/core/providers/shared_library/provider_interfaces.h @@ -143,6 +143,7 @@ struct ProviderHost { virtual const OrtApiBase* OrtGetApiBase() = 0; virtual Status GetOptimizerByName(const std::string& name, + std::unordered_map& key_value_configs, std::function>(const GraphViewer&)>& selection_func) = 0; virtual void* HeapAllocate(size_t size) = 0; diff --git a/onnxruntime/core/providers/tensorrt/tensorrt_execution_provider.cc b/onnxruntime/core/providers/tensorrt/tensorrt_execution_provider.cc index 8a899884e5..4100552257 100644 --- a/onnxruntime/core/providers/tensorrt/tensorrt_execution_provider.cc +++ b/onnxruntime/core/providers/tensorrt/tensorrt_execution_provider.cc @@ -2670,7 +2670,8 @@ TensorrtExecutionProvider::GetCapability(const GraphViewer& graph, */ std::function>(const GraphViewer&)> selection_func; - auto status = g_host->GetOptimizerByName("ConstantFoldingDQ", selection_func); + std::unordered_map key_value_configs = {}; + auto status = g_host->GetOptimizerByName("ConstantFoldingDQ", key_value_configs, selection_func); std::vector> selection_cc; if (selection_func) { selection_cc = selection_func(graph); diff --git a/onnxruntime/core/session/inference_session.cc b/onnxruntime/core/session/inference_session.cc index a852c195c1..e17562eb84 100644 --- a/onnxruntime/core/session/inference_session.cc +++ b/onnxruntime/core/session/inference_session.cc @@ -1845,11 +1845,20 @@ common::Status InferenceSession::Initialize() { record_runtime_optimization_produced_op_schema, *session_logger_)); - // register predefined transformers for EP + // Register predefined optimizer names for EPs. + // It won't create optimizer instance until EP later requests it by calling GetOptimizerByName in provider bridge + std::vector predefined_optimizer_names_for_ep; + predefined_optimizer_names_for_ep.push_back("ConstantFoldingDQ"); auto graph_optimizer_registry = onnxruntime::GraphOptimizerRegistry::Get(); + graph_optimizer_registry->AddPredefinedOptimizerNames(predefined_optimizer_names_for_ep); + graph_optimizer_registry->AddCpuEpReference(execution_providers_.Get(onnxruntime::kCpuExecutionProvider)); + + // Don't create optimizer instances upfront. + /* graph_optimizer_registry->AddPredefinedOptimizers(session_options_, *execution_providers_. Get(onnxruntime::kCpuExecutionProvider), *session_logger_); + */ #ifdef USE_DML const IExecutionProvider* dmlExecutionProvider = execution_providers_.Get(kDmlExecutionProvider); diff --git a/onnxruntime/core/session/provider_bridge_ort.cc b/onnxruntime/core/session/provider_bridge_ort.cc index c954141ae0..10b5a4138b 100644 --- a/onnxruntime/core/session/provider_bridge_ort.cc +++ b/onnxruntime/core/session/provider_bridge_ort.cc @@ -4,6 +4,7 @@ // This is the Onnxruntime side of the bridge to allow providers to be built as a DLL // It implements onnxruntime::ProviderHost +#include #include "core/common/inlined_containers.h" #include "core/common/path_string.h" #include "core/framework/allocator_utils.h" @@ -37,8 +38,6 @@ #include "core/framework/model_metadef_id_generator.h" #include "core/optimizer/qdq_transformer/selectors_actions/qdq_selectors.h" #include "core/optimizer/qdq_transformer/selectors_actions/shared/utils.h" -#include "core/optimizer/qdq_transformer/constant_folding_dq_node.h" -#include "core/session/onnxruntime_session_options_config_keys.h" #include "core/session/onnxruntime_c_api.h" #include "core/common/string_helper.h" @@ -215,119 +214,17 @@ struct ProviderHostImpl : ProviderHost { const OrtApiBase* OrtGetApiBase() override { return ::OrtGetApiBase(); } Status GetOptimizerByName(const std::string& name, + std::unordered_map& key_value_configs, std::function>(const GraphViewer&)>& selection_func) override { - //static const GraphTransformerManager& graph_transformer_mgr = transformer_mgr; std::string optimizer_name(name); - // Pre-defined graph transformers/optimizers - static const std::string kEP_GRAPH_TRANSFORMER_CONSTANT_FOLDING_DQ = "ConstantFoldingDQ"; + auto optimizer_registry = onnxruntime::GraphOptimizerRegistry::Get(); + auto func = optimizer_registry->GetSelectionFunc(optimizer_name, key_value_configs); - // ConstantFoldingDQ optimization function - auto constant_folding_dq_optimization = [&](Graph& graph, const ComputeCapability& optimization_cc, ComputeCapability& cc_to_update) -> Status { - std::string optimizer_name = kEP_GRAPH_TRANSFORMER_CONSTANT_FOLDING_DQ; - auto logger = const_cast(&logging::LoggingManager::DefaultLogger()); - std::unordered_set original_initializers_to_remove; - std::unordered_set new_initializers_to_add; - InlinedHashSet dq_node_index_set; - - // iterate node_to_optimize to: - // 1. get original initializers to remove - // 2. add new initializers - // 3. create dq node index set - for (const auto& index : optimization_cc.sub_graph->nodes) { - auto node = graph.GetNode(index); - if (node->OpType() != "DequantizeLinear") { - continue; - } - auto input_0 = node->InputDefs()[0]; - auto output_0 = node->OutputDefs()[0]; - original_initializers_to_remove.insert(input_0->Name()); - new_initializers_to_add.insert(output_0->Name()); - dq_node_index_set.insert(index); - } - - auto optimizer_registry = onnxruntime::GraphOptimizerRegistry::Get(); - - ConstantFoldingDQ* transformer = static_cast(optimizer_registry->GetTransformerByName(optimizer_name)); - transformer->UpdateNodeIndexSet(dq_node_index_set); - - // apply constant folding on DQ nodes - optimizer_registry->ApplyTransformer(graph, optimizer_name, *logger); - - // update the overall ComputeCapability - std::vector updated_nodes; - for (auto index : cc_to_update.sub_graph->nodes) { - if (dq_node_index_set.find(index) != dq_node_index_set.end()) { - continue; - } - updated_nodes.push_back(index); - } - cc_to_update.sub_graph->nodes = updated_nodes; - - auto original_meta_def = cc_to_update.sub_graph->GetMetaDef(); - std::unique_ptr updated_meta_def = std::make_unique(); - updated_meta_def->name = original_meta_def->name; - updated_meta_def->domain = original_meta_def->domain; - updated_meta_def->since_version = original_meta_def->since_version; - updated_meta_def->status = original_meta_def->status; - updated_meta_def->inputs = original_meta_def->inputs; - updated_meta_def->outputs = original_meta_def->outputs; - updated_meta_def->attributes = original_meta_def->attributes; - updated_meta_def->doc_string = original_meta_def->doc_string; -#if !defined(ORT_MINIMAL_BUILD) - updated_meta_def->type_and_shape_inference_function = original_meta_def->type_and_shape_inference_function; -#endif - for (auto constant_initializer : original_meta_def->constant_initializers) { - if (original_initializers_to_remove.find(constant_initializer) != original_initializers_to_remove.end()) { - continue; - } - updated_meta_def->constant_initializers.push_back(constant_initializer); - } - - for (auto constant_initializer : new_initializers_to_add) { - updated_meta_def->constant_initializers.push_back(constant_initializer); - } - - cc_to_update.sub_graph->SetMetaDef(std::move(updated_meta_def)); - - return Status::OK(); - }; - - // ConstantFoldingDQ selection function - auto constant_folding_dq_selection = [&](const GraphViewer& graph_viewer) -> std::vector> { - std::vector> result; - std::unique_ptr sub_graph = std::make_unique(); - const std::vector& node_index = graph_viewer.GetNodesInTopologicalOrder(ExecutionOrder::PRIORITY_BASED /*priority-based topological sort*/); - InitializedTensorSet constant_inputs; - const InlinedHashSet excluded_initializers; - - // Select DequantizeLinear node where all inputs are constant - for (const auto& index : node_index) { - const auto& node = graph_viewer.GetNode(index); - if (node->OpType() != "DequantizeLinear") { - continue; - } - if (!graph_utils::AllNodeInputsAreConstant(graph_viewer.GetGraph(), *node, constant_inputs, excluded_initializers)) { - continue; - } - sub_graph->nodes.push_back(index); - } - - result.push_back(std::make_unique(std::move(sub_graph))); - result.back()->optimization_func = constant_folding_dq_optimization; - return result; - }; - - // Optimizer lookup table - static std::unordered_map>(const GraphViewer&)>> optimizer_to_selection_function; - if (optimizer_to_selection_function.find(kEP_GRAPH_TRANSFORMER_CONSTANT_FOLDING_DQ) == optimizer_to_selection_function.end()) { - optimizer_to_selection_function[kEP_GRAPH_TRANSFORMER_CONSTANT_FOLDING_DQ] = constant_folding_dq_selection; + if (func.has_value()) { + selection_func = func.value(); } - auto look_up = optimizer_to_selection_function.find(optimizer_name); - if (look_up != optimizer_to_selection_function.end()) { - selection_func = optimizer_to_selection_function[optimizer_name]; - } return Status::OK(); };