Transpose MatMul fusion fixes (#4728)

Fix Transpose MatMul fusion handling of existing TransposeScaleMatMul node's attributes and enable support for missing Transpose perm attribute.
Update expected test data to account for floating point calculation differences resulting from the fusion.
This commit is contained in:
edgchen1 2020-08-10 13:00:22 -07:00 committed by GitHub
parent 316d1a9e69
commit 487665c21f
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
8 changed files with 181 additions and 22 deletions

View file

@ -12,6 +12,10 @@ namespace onnxruntime {
* For example, given matrices A and B and constant scalars t, u, and v:
* Mul(v, MatMul(Mul(t, A), Mul(u, B))
* -> TransposeScaleMatMul(A, B, alpha=t*u*v)
*
* Note: Since both leading and following scales may be fused into a single
* scale, the order and number of mathematical operations may change. This may
* yield different results with floating point calculations.
*/
class MatMulScaleFusion : public GraphTransformer {
public:

View file

@ -10,6 +10,28 @@ using namespace ONNX_NAMESPACE;
using namespace ::onnxruntime::common;
namespace onnxruntime {
static bool GetTransposePerms(const Node& transpose_node, std::vector<int64_t>& perms) {
ORT_ENFORCE(transpose_node.InputDefs().size() == 1);
// use perms if present
const auto perm_attr = transpose_node.GetAttributes().find("perm");
if (perm_attr != transpose_node.GetAttributes().end()) {
perms = RetrieveValues<int64_t>(perm_attr->second);
return true;
}
// otherwise, reverse dimensions
const NodeArg& input = *transpose_node.InputDefs()[0];
const TensorShapeProto* shape = input.Shape();
if (!shape) {
return false;
}
perms.resize(shape->dim_size());
std::iota(perms.rbegin(), perms.rend(), 0);
return true;
}
static Node* GetTransposeNodeFromOutput(Graph& graph, NodeArg& node_arg) {
Node* trans_node = graph.GetMutableProducerNode(node_arg.Name());
if (trans_node == nullptr || trans_node->OpType() != "Transpose") {
@ -21,7 +43,11 @@ static Node* GetTransposeNodeFromOutput(Graph& graph, NodeArg& node_arg) {
return nullptr;
}
auto perms = RetrieveValues<int64_t>(trans_node->GetAttributes().at("perm"));
std::vector<int64_t> perms;
if (!GetTransposePerms(*trans_node, perms)) {
return nullptr;
}
int64_t rank = perms.size();
if (rank < 2) {
return nullptr;
@ -109,15 +135,16 @@ Status MatmulTransposeFusion::ApplyImpl(Graph& graph, bool& modified, int graph_
input_defs,
output_defs, {}, kMSDomain);
bool transpose_left = (left != nullptr);
bool transpose_right = (right != nullptr);
float alpha = 1.0f;
if (node.OpType() == "TransposeScaleMatMul") {
transpose_left ^= static_cast<bool>(node.GetAttributes().at("transA").i());
}
bool transpose_right = (right != nullptr);
if (node.OpType() == "TransposeScaleMatMul") {
transpose_right ^= static_cast<bool>(node.GetAttributes().at("transB").i());
alpha = node.GetAttributes().at("alpha").f();
}
matmul_node.AddAttribute("transA", static_cast<int64_t>(transpose_left));
matmul_node.AddAttribute("transB", static_cast<int64_t>(transpose_right));
matmul_node.AddAttribute("alpha", alpha);
// Assign provider to this new node. Provider should be same as the provider for old node.
matmul_node.SetExecutionProviderType(node.GetExecutionProviderType());

View file

@ -751,20 +751,58 @@ TEST_F(GraphTransformationTests, TransposeMatmulFusionOnThreeTranspose) {
}
TEST_F(GraphTransformationTests, TransposeMatmulNoFusionOnInvalidPerm) {
auto model_uri = MODEL_FOLDER "fusion/transpose_matmul_4d_fusion_invalid_perm.onnx";
const std::vector<PathString> model_uris = {
MODEL_FOLDER "fusion/transpose_matmul_4d_fusion_invalid_perm.onnx",
MODEL_FOLDER "fusion/transpose_matmul_4d_fusion_invalid_default_perm.onnx",
};
for (const auto& model_uri : model_uris) {
std::shared_ptr<Model> p_model;
ASSERT_STATUS_OK(Model::Load(model_uri, p_model, nullptr, *logger_));
Graph& graph = p_model->MainGraph();
onnxruntime::GraphTransformerManager graph_transformation_mgr{5};
ASSERT_STATUS_OK(graph_transformation_mgr.Register(
onnxruntime::make_unique<MatmulTransposeFusion>(), TransformerLevel::Level1));
ASSERT_STATUS_OK(graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level1, *logger_));
std::map<std::string, int> op_to_count = CountOpsInGraph(graph);
ASSERT_EQ(op_to_count["Transpose"], 1);
ASSERT_EQ(op_to_count["MatMul"], 1);
ASSERT_EQ(op_to_count["TransposeScaleMatMul"], 0);
}
}
TEST_F(GraphTransformationTests, TransposeMatmulFusionFromTransposeScaleMatMul) {
auto model_uri = MODEL_FOLDER "fusion/transpose_matmul_2d_fusion_from_transpose_scale_matmul.onnx";
std::shared_ptr<Model> p_model;
ASSERT_TRUE(Model::Load(model_uri, p_model, nullptr, *logger_).IsOK());
ASSERT_STATUS_OK(Model::Load(model_uri, p_model, nullptr, *logger_));
Graph& graph = p_model->MainGraph();
float expected_alpha;
{
auto transpose_scale_matmul_node =
std::find_if(
graph.Nodes().cbegin(), graph.Nodes().cend(),
[](const Node& node) { return node.Name() == "TransposeScaleMatMul"; });
ASSERT_NE(transpose_scale_matmul_node, graph.Nodes().cend());
expected_alpha = transpose_scale_matmul_node->GetAttributes().at("alpha").f();
}
onnxruntime::GraphTransformerManager graph_transformation_mgr{5};
graph_transformation_mgr.Register(onnxruntime::make_unique<MatmulTransposeFusion>(), TransformerLevel::Level1);
auto ret = graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level1, *logger_);
ASSERT_TRUE(ret.IsOK());
ASSERT_STATUS_OK(graph_transformation_mgr.Register(
onnxruntime::make_unique<MatmulTransposeFusion>(), TransformerLevel::Level1));
ASSERT_STATUS_OK(graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level1, *logger_));
std::map<std::string, int> op_to_count = CountOpsInGraph(graph);
ASSERT_TRUE(op_to_count["Transpose"] == 1);
ASSERT_TRUE(op_to_count["MatMul"] == 1);
ASSERT_TRUE(op_to_count["TransposeScaleMatMul"] == 0);
ASSERT_EQ(op_to_count["Transpose"], 0);
ASSERT_EQ(op_to_count["MatMul"], 0);
ASSERT_EQ(op_to_count["TransposeScaleMatMul"], 1);
auto& transpose_scale_matmul_node = *graph.Nodes().begin();
ASSERT_EQ(transpose_scale_matmul_node.OpType(), "TransposeScaleMatMul");
ASSERT_FALSE(static_cast<bool>(transpose_scale_matmul_node.GetAttributes().at("transA").i()));
ASSERT_FALSE(static_cast<bool>(transpose_scale_matmul_node.GetAttributes().at("transB").i()));
ASSERT_EQ(transpose_scale_matmul_node.GetAttributes().at("alpha").f(), expected_alpha);
}
TEST_F(GraphTransformationTests, Gemm_LeakyRelu_Fusion) {

View file

@ -0,0 +1,90 @@
import onnx
from onnx import helper
from onnx import TensorProto
from onnx import OperatorSetIdProto
onnxdomain = OperatorSetIdProto()
onnxdomain.version = 12
# The empty string ("") or absence of this field implies the operator set that is defined as part of the ONNX specification.
onnxdomain.domain = ""
msdomain = OperatorSetIdProto()
msdomain.version = 1
msdomain.domain = "com.microsoft"
opsets = [onnxdomain, msdomain]
def save(model_path, nodes, inputs, outputs, initializers):
graph = helper.make_graph(
nodes,
"TransposeMatMulTest",
inputs, outputs, initializers)
model = helper.make_model(
graph, opset_imports=opsets, producer_name="onnxruntime-test")
onnx.save(model, model_path)
def gen_from_transpose_scale_matmul(model_path):
nodes = [
helper.make_node(
"Transpose",
["input_0"],
["transposed_input_0"]),
helper.make_node(
"TransposeScaleMatMul",
["transposed_input_0", "input_1"],
["output"],
"TransposeScaleMatMul",
"",
msdomain.domain,
alpha=3.0, transA=1)
]
inputs = [
helper.make_tensor_value_info(
"input_0", TensorProto.FLOAT, ['M', 'K']),
helper.make_tensor_value_info(
"input_1", TensorProto.FLOAT, ['K', 'N'])
]
outputs = [
helper.make_tensor_value_info(
"output", TensorProto.FLOAT, ['M', 'N'])
]
save(model_path, nodes, inputs, outputs, [])
gen_from_transpose_scale_matmul(
"transpose_matmul_2d_fusion_from_transpose_scale_matmul.onnx")
def gen_invalid_default_perm(model_path):
nodes = [
helper.make_node(
"Transpose",
["input_0"],
["transposed_input_0"]),
helper.make_node(
"MatMul",
["transposed_input_0", "input_1"],
["output"])
]
inputs = [
helper.make_tensor_value_info(
"input_0", TensorProto.FLOAT, ['K', 'M', 3, 2]),
helper.make_tensor_value_info(
"input_1", TensorProto.FLOAT, [2, 3, 'K', 'N'])
]
outputs = [
helper.make_tensor_value_info(
"output", TensorProto.FLOAT, [2, 3, 'M', 'N'])
]
save(model_path, nodes, inputs, outputs, [])
gen_invalid_default_perm(
"transpose_matmul_4d_fusion_invalid_default_perm.onnx")

View file

@ -81,10 +81,10 @@ class ORTGlueTest(unittest.TestCase):
assert_allclose(results['loss'], expected_loss, rtol=self.rtol)
def test_roberta_fp16_with_mrpc(self):
expected_acc = 0.8995098039215687
expected_f1 = 0.9279437609841829
expected_acc_and_f1 = 0.9137267824528758
expected_loss = 0.32052762967114357
expected_acc = 0.8946078431372549
expected_f1 = 0.9244288224956063
expected_acc_and_f1 = 0.9095183328164307
expected_loss = 0.2860557144763423
results = self.run_glue(model_name="roberta-base", task_name="MRPC", fp16=True)
assert_allclose(results['acc'], expected_acc, rtol=self.rtol)
@ -105,10 +105,10 @@ class ORTGlueTest(unittest.TestCase):
assert_allclose(results['loss'], expected_loss, rtol=self.rtol)
def test_bert_fp16_with_mrpc(self):
expected_acc = 0.8651960784313726
expected_f1 = 0.9063032367972743
expected_acc_and_f1 = 0.8857496576143234
expected_loss = 0.38716790532948925
expected_acc = 0.8431372549019608
expected_f1 = 0.888888888888889
expected_acc_and_f1 = 0.8660130718954249
expected_loss = 0.39904916637084065
results = self.run_glue(model_name="bert-base-cased", task_name="MRPC", fp16=True)
assert_allclose(results['acc'], expected_acc, rtol=self.rtol)

View file

@ -99,8 +99,8 @@ class ORTMultipleChoiceTest(unittest.TestCase):
assert_allclose(results['loss'], expected_loss)
def test_bert_fp16_with_swag(self):
expected_acc = 0.7882135359392183
expected_loss = 0.6469693916158167
expected_acc = 0.7876137158852344
expected_loss = 0.6469138865876788
results = self.run_multiple_choice(model_name="bert-base-cased", task_name="swag", fp16=True)
assert_allclose(results['acc'], expected_acc)