generate transformers bug fix (#838)

* fix graph transformer generation

* add more tests

* cosmetic changes

* more changes per review
This commit is contained in:
Ashwini Khade 2019-04-16 14:10:33 -07:00 committed by GitHub
parent 1818835795
commit 14d63b5f45
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
5 changed files with 55 additions and 42 deletions

View file

@ -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);

View file

@ -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) {

View file

@ -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);
}

View file

@ -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);

View file

@ -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