mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
generate transformers bug fix (#838)
* fix graph transformer generation * add more tests * cosmetic changes * more changes per review
This commit is contained in:
parent
1818835795
commit
14d63b5f45
5 changed files with 55 additions and 42 deletions
|
|
@ -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<std::unique_ptr<RewriteRule>> GenerateRewriteRules(TransformerLevel level,
|
||||
const std::vector<std::string>* rules_to_enable = nullptr);
|
||||
const std::vector<std::string>& 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<std::unique_ptr<GraphTransformer>> GenerateTransformers(TransformerLevel level,
|
||||
std::vector<std::string>* rules_and_transformers_to_enable = nullptr);
|
||||
const std::vector<std::string>& 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);
|
||||
|
|
|
|||
|
|
@ -17,7 +17,7 @@ namespace onnxruntime {
|
|||
namespace transformer_utils {
|
||||
|
||||
std::vector<std::unique_ptr<RewriteRule>> GenerateRewriteRules(TransformerLevel level,
|
||||
const std::vector<std::string>* rules_to_enable) {
|
||||
const std::vector<std::string>& rules_to_enable) {
|
||||
std::vector<std::unique_ptr<RewriteRule>> rules;
|
||||
switch (level) {
|
||||
case TransformerLevel::Level1:
|
||||
|
|
@ -32,9 +32,9 @@ std::vector<std::unique_ptr<RewriteRule>> GenerateRewriteRules(TransformerLevel
|
|||
ORT_ENFORCE(false, "Unsupported level" + std::to_string(static_cast<uint32_t>(level)));
|
||||
}
|
||||
|
||||
if (rules_to_enable != nullptr && !rules_to_enable->empty()) {
|
||||
if (!rules_to_enable.empty()) {
|
||||
std::vector<std::unique_ptr<RewriteRule>> 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<RewriteRule>& item) {
|
||||
if ((item != nullptr) && (item->Name() == rule_name)) {
|
||||
filtered_list.push_back(std::move(item));
|
||||
|
|
@ -48,11 +48,11 @@ std::vector<std::unique_ptr<RewriteRule>> GenerateRewriteRules(TransformerLevel
|
|||
}
|
||||
|
||||
std::unique_ptr<RuleBasedGraphTransformer> GenerateRuleBasedGraphTransformer(TransformerLevel level,
|
||||
const std::vector<std::string>* rules_to_enable,
|
||||
const std::vector<std::string>& rules_to_enable,
|
||||
const std::unordered_set<std::string>& 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<RuleBasedGraphTransformer>{};
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
std::unique_ptr<RuleBasedGraphTransformer> rule_transformer =
|
||||
|
|
@ -68,33 +68,23 @@ std::unique_ptr<RuleBasedGraphTransformer> GenerateRuleBasedGraphTransformer(Tra
|
|||
}
|
||||
|
||||
std::vector<std::unique_ptr<GraphTransformer>> GenerateTransformers(TransformerLevel level,
|
||||
std::vector<std::string>* transformers_and_rules_to_enable) {
|
||||
const std::vector<std::string>& transformers_and_rules_to_enable) {
|
||||
std::vector<std::unique_ptr<GraphTransformer>> transformers;
|
||||
|
||||
// Generate rule-based transformer.
|
||||
bool non_empty_rule_transformer = false;
|
||||
|
||||
std::unique_ptr<RuleBasedGraphTransformer> rule_transformer = nullptr;
|
||||
switch (level) {
|
||||
case TransformerLevel::Level1: {
|
||||
std::unordered_set<std::string> l1_execution_providers = {};
|
||||
std::unique_ptr<RuleBasedGraphTransformer> 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<std::string> l2_execution_providers = {onnxruntime::kCpuExecutionProvider};
|
||||
std::unique_ptr<RuleBasedGraphTransformer> 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<GemmActivationFusion>(l2_execution_providers));
|
||||
transformers.emplace_back(std::make_unique<MatMulAddFusion>(l2_execution_providers));
|
||||
transformers.emplace_back(std::make_unique<ConvActivationFusion>(l2_execution_providers));
|
||||
|
|
@ -108,14 +98,22 @@ std::vector<std::unique_ptr<GraphTransformer>> 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<std::unique_ptr<GraphTransformer>> 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<GraphTransformer>& item) {
|
||||
if ((item != nullptr) && (item->Name() == t_name)) {
|
||||
|
|
@ -125,8 +123,6 @@ std::vector<std::unique_ptr<GraphTransformer>> GenerateTransformers(TransformerL
|
|||
}
|
||||
return filtered_list;
|
||||
}
|
||||
|
||||
return transformers;
|
||||
}
|
||||
|
||||
std::string GenerateRuleBasedTransformerName(TransformerLevel level) {
|
||||
|
|
|
|||
|
|
@ -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<std::string>& custom_list) {
|
||||
const std::vector<std::string>& 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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -345,7 +345,7 @@ class InferenceSession {
|
|||
|
||||
void AddPredefinedTransformers(GraphTransformerManager& transformer_manager,
|
||||
TransformerLevel graph_optimization_level,
|
||||
std::vector<std::string>& custom_list);
|
||||
const std::vector<std::string>& custom_list);
|
||||
|
||||
void InitLogger(logging::LoggingManager* logging_manager);
|
||||
|
||||
|
|
|
|||
|
|
@ -20,16 +20,16 @@ TEST(GraphTransformerUtilsTests, TestGenerateRewriterules) {
|
|||
|
||||
// Rule name match test
|
||||
std::vector<std::string> 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<std::string> 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<std::string> 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<RuleBasedGraphTransformer*>(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
|
||||
|
|
|
|||
Loading…
Reference in a new issue