mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
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:
parent
316d1a9e69
commit
487665c21f
8 changed files with 181 additions and 22 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
|
||||
|
|
|
|||
|
|
@ -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) {
|
||||
|
|
|
|||
Binary file not shown.
BIN
onnxruntime/test/testdata/transform/fusion/transpose_matmul_4d_fusion_invalid_default_perm.onnx
vendored
Normal file
BIN
onnxruntime/test/testdata/transform/fusion/transpose_matmul_4d_fusion_invalid_default_perm.onnx
vendored
Normal file
Binary file not shown.
90
onnxruntime/test/testdata/transform/fusion/transpose_matmul_gen.py
vendored
Normal file
90
onnxruntime/test/testdata/transform/fusion/transpose_matmul_gen.py
vendored
Normal 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")
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Reference in a new issue