diff --git a/include/onnxruntime/core/optimizer/graph_transformer_utils.h b/include/onnxruntime/core/optimizer/graph_transformer_utils.h index c0998fe015..e592b1f098 100644 --- a/include/onnxruntime/core/optimizer/graph_transformer_utils.h +++ b/include/onnxruntime/core/optimizer/graph_transformer_utils.h @@ -15,13 +15,13 @@ namespace transformer_utils { If rules_to_enable is not empty, it returns the intersection of predefined rules and rules_to_enable. TODO: This is visible for testing at the moment, but we should rather make it private. */ std::vector> GenerateRewriteRules(TransformerLevel level, - const std::vector* rules_to_enable = nullptr); + const std::vector& rules_to_enable = {}); /** Generates all predefined (both rule-based and non-rule-based) transformers for this level. If transformers_and_rules_to_enable is not empty, it returns the intersection between the predefined transformers/rules and the transformers_and_rules_to_enable. */ std::vector> GenerateTransformers(TransformerLevel level, - std::vector* rules_and_transformers_to_enable = nullptr); + const std::vector& rules_and_transformers_to_enable = {}); /** Given a TransformerLevel, this method generates a name for the rule-based graph transformer of that level. */ std::string GenerateRuleBasedTransformerName(TransformerLevel level); diff --git a/onnxruntime/core/optimizer/graph_transformer_utils.cc b/onnxruntime/core/optimizer/graph_transformer_utils.cc index 4fa39a46d4..4033a61323 100644 --- a/onnxruntime/core/optimizer/graph_transformer_utils.cc +++ b/onnxruntime/core/optimizer/graph_transformer_utils.cc @@ -17,7 +17,7 @@ namespace onnxruntime { namespace transformer_utils { std::vector> GenerateRewriteRules(TransformerLevel level, - const std::vector* rules_to_enable) { + const std::vector& rules_to_enable) { std::vector> rules; switch (level) { case TransformerLevel::Level1: @@ -32,9 +32,9 @@ std::vector> GenerateRewriteRules(TransformerLevel ORT_ENFORCE(false, "Unsupported level" + std::to_string(static_cast(level))); } - if (rules_to_enable != nullptr && !rules_to_enable->empty()) { + if (!rules_to_enable.empty()) { std::vector> filtered_list; - for (const auto& rule_name : *rules_to_enable) { + for (const auto& rule_name : rules_to_enable) { std::for_each(rules.begin(), rules.end(), [&](std::unique_ptr& item) { if ((item != nullptr) && (item->Name() == rule_name)) { filtered_list.push_back(std::move(item)); @@ -48,11 +48,11 @@ std::vector> GenerateRewriteRules(TransformerLevel } std::unique_ptr GenerateRuleBasedGraphTransformer(TransformerLevel level, - const std::vector* rules_to_enable, + const std::vector& rules_to_enable, const std::unordered_set& compatible_execution_providers) { auto rewrite_rules_to_register = transformer_utils::GenerateRewriteRules(level, rules_to_enable); if (rewrite_rules_to_register.empty()) { - return std::unique_ptr{}; + return nullptr; } std::unique_ptr rule_transformer = @@ -68,33 +68,23 @@ std::unique_ptr GenerateRuleBasedGraphTransformer(Tra } std::vector> GenerateTransformers(TransformerLevel level, - std::vector* transformers_and_rules_to_enable) { + const std::vector& transformers_and_rules_to_enable) { std::vector> transformers; - - // Generate rule-based transformer. - bool non_empty_rule_transformer = false; - + std::unique_ptr rule_transformer = nullptr; switch (level) { case TransformerLevel::Level1: { std::unordered_set l1_execution_providers = {}; - std::unique_ptr rule_transformer = - transformer_utils::GenerateRuleBasedGraphTransformer(level, transformers_and_rules_to_enable, l1_execution_providers); - if (rule_transformer) { - transformers.emplace_back(std::move(rule_transformer)); - non_empty_rule_transformer = true; - } - + rule_transformer = GenerateRuleBasedGraphTransformer(level, transformers_and_rules_to_enable, l1_execution_providers); // At the moment, we have only a rule-based transformer for Level1. } break; case TransformerLevel::Level2: { std::unordered_set l2_execution_providers = {onnxruntime::kCpuExecutionProvider}; - std::unique_ptr rule_transformer = - transformer_utils::GenerateRuleBasedGraphTransformer(level, transformers_and_rules_to_enable, l2_execution_providers); - if (rule_transformer) { - transformers.emplace_back(std::move(rule_transformer)); - non_empty_rule_transformer = true; - } + + // create rule based transformer consisting of all the level2 rewrite rules + rule_transformer = GenerateRuleBasedGraphTransformer(level, transformers_and_rules_to_enable, l2_execution_providers); + + // create standalone transformers transformers.emplace_back(std::make_unique(l2_execution_providers)); transformers.emplace_back(std::make_unique(l2_execution_providers)); transformers.emplace_back(std::make_unique(l2_execution_providers)); @@ -108,14 +98,22 @@ std::vector> GenerateTransformers(TransformerL break; } - // If the rule-based transformer is not empty, it should be included in the custom transformer list below. - if (non_empty_rule_transformer) { - transformers_and_rules_to_enable->push_back(transformer_utils::GenerateRuleBasedTransformerName(level)); - } - if (transformers_and_rules_to_enable != nullptr && !transformers_and_rules_to_enable->empty()) { - // pick custom transformers enabled for this session + // if the custom list to enable transformers\rules is empty then return the default generated transformers and rules + // otherwise generate a filtered list based on the provided custom list. + if (transformers_and_rules_to_enable.empty()) { + if (rule_transformer != nullptr) { + transformers.emplace_back(std::move(rule_transformer)); + } + return transformers; + + } else { std::vector> filtered_list; - for (const auto& t_name : *transformers_and_rules_to_enable) { + // If the rule-based transformer is not empty, it should be included in the custom transformer list below. + if (rule_transformer != nullptr) { + filtered_list.emplace_back(std::move(rule_transformer)); + } + // pick custom transformers enabled for this session + for (const auto& t_name : transformers_and_rules_to_enable) { std::for_each(transformers.begin(), transformers.end(), [&](std::unique_ptr& item) { if ((item != nullptr) && (item->Name() == t_name)) { @@ -125,8 +123,6 @@ std::vector> GenerateTransformers(TransformerL } return filtered_list; } - - return transformers; } std::string GenerateRuleBasedTransformerName(TransformerLevel level) { diff --git a/onnxruntime/core/session/inference_session.cc b/onnxruntime/core/session/inference_session.cc index e58fa5739a..d2ea3b14ff 100644 --- a/onnxruntime/core/session/inference_session.cc +++ b/onnxruntime/core/session/inference_session.cc @@ -910,10 +910,10 @@ void InferenceSession::InitLogger(logging::LoggingManager* logging_manager) { // Registers all the predefined transformers with transformer manager void InferenceSession::AddPredefinedTransformers(GraphTransformerManager& transformer_manager, TransformerLevel graph_optimization_level, - std::vector& custom_list) { + const std::vector& custom_list) { auto add_transformers = [&](TransformerLevel level) { // Generate and register transformers for level - auto transformers_to_register = transformer_utils::GenerateTransformers(level, &custom_list); + auto transformers_to_register = transformer_utils::GenerateTransformers(level, custom_list); for (auto& entry : transformers_to_register) { transformer_manager.Register(std::move(entry), level); } diff --git a/onnxruntime/core/session/inference_session.h b/onnxruntime/core/session/inference_session.h index ba4a3c8f85..e347ee4c76 100644 --- a/onnxruntime/core/session/inference_session.h +++ b/onnxruntime/core/session/inference_session.h @@ -345,7 +345,7 @@ class InferenceSession { void AddPredefinedTransformers(GraphTransformerManager& transformer_manager, TransformerLevel graph_optimization_level, - std::vector& custom_list); + const std::vector& custom_list); void InitLogger(logging::LoggingManager* logging_manager); diff --git a/onnxruntime/test/optimizer/graph_transform_utils_test.cc b/onnxruntime/test/optimizer/graph_transform_utils_test.cc index ba963a3f82..e9cd822b36 100644 --- a/onnxruntime/test/optimizer/graph_transform_utils_test.cc +++ b/onnxruntime/test/optimizer/graph_transform_utils_test.cc @@ -20,16 +20,16 @@ TEST(GraphTransformerUtilsTests, TestGenerateRewriterules) { // Rule name match test std::vector custom_list = {"EliminateIdentity", "ConvAddFusion", "ConvMulFusion", "abc", "def"}; - rewrite_rules = transformer_utils::GenerateRewriteRules(TransformerLevel::Level1, &custom_list); + rewrite_rules = transformer_utils::GenerateRewriteRules(TransformerLevel::Level1, custom_list); // validate each rule returned is present in the custom list for (const auto& rule : rewrite_rules) { ASSERT_TRUE(std::find(custom_list.begin(), custom_list.end(), rule->Name()) != custom_list.end()); } - // Rule name no match test. Test to validate empty rules list is returned when + // Rule name no match test. Test to validate empty rules list is returned when // there is no match in custom list custom_list = {"abc"}; - rewrite_rules = transformer_utils::GenerateRewriteRules(TransformerLevel::Level1, &custom_list); + rewrite_rules = transformer_utils::GenerateRewriteRules(TransformerLevel::Level1, custom_list); ASSERT_TRUE(rewrite_rules.size() == 0); } @@ -39,7 +39,7 @@ TEST(GraphTransformerUtilsTests, TestGenerateGraphTransformers) { // Transformer name match test std::vector custom_list = {"EliminateIdentity", "ConvAddFusion", "ConvMulFusion", "abc", "def"}; - transformers = transformer_utils::GenerateTransformers(TransformerLevel::Level2, &custom_list); + transformers = transformer_utils::GenerateTransformers(TransformerLevel::Level2, custom_list); ASSERT_TRUE(transformers.size() == 2); // validate each rule returned is present in the custom list for (const auto& transformer : transformers) { @@ -48,9 +48,26 @@ TEST(GraphTransformerUtilsTests, TestGenerateGraphTransformers) { // Transformer name no match test. When there is no match empty list is expected. custom_list = {"EliminateIdentity"}; - transformers = transformer_utils::GenerateTransformers(TransformerLevel::Level2, &custom_list); + transformers = transformer_utils::GenerateTransformers(TransformerLevel::Level2, custom_list); ASSERT_TRUE(transformers.size() == 0); } +TEST(GraphTransformerUtilsTests, TestGenerateGraphTransformers_CustomList) { + // custom list of rules and transformers + std::string l1_rule = "EliminateIdentity"; + std::string l2_transformer = "ConvAddFusion"; + std::vector custom_list = {l1_rule, l2_transformer}; + + auto transformers = transformer_utils::GenerateTransformers(TransformerLevel::Level1, custom_list); + ASSERT_TRUE(transformers.size() == 1); + auto rule_transformer = dynamic_cast(transformers[0].get()); + ASSERT_TRUE(rule_transformer->GetRewriteRules().size() == 1); + ASSERT_TRUE(rule_transformer->GetRewriteRules()[0]->Name() == l1_rule); + + transformers = transformer_utils::GenerateTransformers(TransformerLevel::Level2, custom_list); + ASSERT_TRUE(transformers.size() == 1); + ASSERT_TRUE(transformers[0]->Name() == l2_transformer); +} + } // namespace test } // namespace onnxruntime