Fix TransposeScaleMatMul and MatMulScaleFusion issues (#5230)

- Rename TransposeScaleMatMul back to TransposeMatMul for backwards compatibility
- Fix MatMulScaleFusion issues:
  - Add check for supported execution providers
  - Add check for supported MatMul input types
This commit is contained in:
edgchen1 2020-09-21 12:34:01 -07:00 committed by GitHub
parent 65740deb10
commit e9671e93f0
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
18 changed files with 168 additions and 65 deletions

View file

@ -20,7 +20,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1,
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, Range);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, WordConvEmbedding);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, GatherND);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, TransposeScaleMatMul);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, TransposeMatMul);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, MurmurHash3);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, float, MaxpoolWithMask);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, Pad);
@ -157,7 +157,7 @@ Status RegisterCpuContribKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, WordConvEmbedding)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, GatherND)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, MurmurHash3)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, TransposeScaleMatMul)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, TransposeMatMul)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, float, MaxpoolWithMask)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, Pad)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, Unique)>,

View file

@ -9,21 +9,21 @@ namespace onnxruntime {
namespace contrib {
ONNX_OPERATOR_KERNEL_EX(
TransposeScaleMatMul,
TransposeMatMul,
kMSDomain,
1,
kCpuExecutionProvider,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
TransposeScaleMatMul);
TransposeMatMul);
TransposeScaleMatMul::TransposeScaleMatMul(const OpKernelInfo& info)
TransposeMatMul::TransposeMatMul(const OpKernelInfo& info)
: OpKernel{info} {
ORT_THROW_IF_ERROR(info.GetAttr("alpha", &alpha_attr_));
ORT_THROW_IF_ERROR(info.GetAttr("transA", &trans_a_attr_));
ORT_THROW_IF_ERROR(info.GetAttr("transB", &trans_b_attr_));
}
Status TransposeScaleMatMul::Compute(OpKernelContext* context) const {
Status TransposeMatMul::Compute(OpKernelContext* context) const {
concurrency::ThreadPool* thread_pool = context->GetOperatorThreadPool();
const Tensor* A = context->Input<Tensor>(0);

View file

@ -8,9 +8,9 @@
namespace onnxruntime {
namespace contrib {
class TransposeScaleMatMul final : public OpKernel {
class TransposeMatMul final : public OpKernel {
public:
TransposeScaleMatMul(const OpKernelInfo& info);
TransposeMatMul(const OpKernelInfo& info);
Status Compute(OpKernelContext* context) const override;

View file

@ -17,9 +17,9 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, BiasGelu);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, BiasGelu);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, BiasGelu);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, TransposeScaleMatMul);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, TransposeScaleMatMul);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, TransposeScaleMatMul);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, TransposeMatMul);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, TransposeMatMul);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, TransposeMatMul);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, Rfft);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, Rfft);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, Rfft);
@ -89,9 +89,9 @@ Status RegisterCudaContribKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, BiasGelu)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, BiasGelu)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, BiasGelu)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, TransposeScaleMatMul)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, TransposeScaleMatMul)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, TransposeScaleMatMul)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, TransposeMatMul)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, TransposeMatMul)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, TransposeMatMul)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, float, Rfft)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, double, Rfft)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, MLFloat16, Rfft)>,

View file

@ -9,7 +9,7 @@ namespace cuda {
#define REGISTER_KERNEL_TYPED(T) \
ONNX_OPERATOR_TYPED_KERNEL_EX( \
TransposeScaleMatMul, \
TransposeMatMul, \
kMSDomain, \
1, \
T, \

View file

@ -1774,15 +1774,14 @@ Matrix product that behaves like numpy.matmul: https://docs.scipy.org/doc/numpy-
ONNX_NAMESPACE::matmulShapeInference(ctx, 0, 1);
});
static const char* TransposeScaleMatMul_doc = R"DOC(
static const char* TransposeMatMul_doc = R"DOC(
Matrix product that behaves like numpy.matmul: https://docs.scipy.org/doc/numpy-1.13.0/reference/generated/numpy.matmul.html
)DOC";
ONNX_CONTRIB_OPERATOR_SCHEMA(TransposeScaleMatMul)
ONNX_CONTRIB_OPERATOR_SCHEMA(TransposeMatMul)
.SetDomain(kMSDomain)
.SinceVersion(1)
.SetSupportLevel(OpSchema::SupportType::EXPERIMENTAL)
.SetDoc("TransposeScaleMatMul")
.SetDoc("TransposeMatMul")
.Input(0, "A", "N-dimensional matrix A", "T")
.Input(1, "B", "N-dimensional matrix B", "T")
.Attr(
@ -1805,7 +1804,7 @@ Matrix product that behaves like numpy.matmul: https://docs.scipy.org/doc/numpy-
"T",
{"tensor(float16)", "tensor(float)", "tensor(double)", "tensor(bfloat16)"},
"Constrain input and output types to float tensors.")
.SetDoc(TransposeScaleMatMul_doc)
.SetDoc(TransposeMatMul_doc)
.TypeAndShapeInferenceFunction([](ONNX_NAMESPACE::InferenceContext& ctx) {
propagateElemTypeFromInputToOutput(ctx, 0, 0);
auto transAAttr = ctx.getAttribute("transA");

View file

@ -43,7 +43,7 @@ optional<float> GetScalarConstantInitializer(const Graph& graph, const NodeArg&
return {};
}
float scalar=0.f;
float scalar{};
utils::MLTypeCallDispatcherRet<
Status, ExtractScalarAsFloatDispatchTarget,
uint32_t, uint64_t, int32_t, int64_t, MLFloat16, float, double, BFloat16>
@ -167,11 +167,31 @@ std::vector<ScaleMergeInfo> GetOutputNodeMerges(
return output_node_merges;
}
bool IsMatMulInputTypeSupported(const Node& node) {
// if no matching key is present, any data type is allowed
static const std::map<std::string, std::vector<std::string>> k_supported_data_types{
{kCudaExecutionProvider, {"tensor(float16)", "tensor(float)", "tensor(double)", "tensor(bfloat16)"}},
{kCpuExecutionProvider, {"tensor(float)"}},
};
const auto it = k_supported_data_types.find(node.GetExecutionProviderType());
return it == k_supported_data_types.end() || optimizer_utils::IsSupportedDataType(node, it->second);
}
Status ProcessNode(
Graph& graph, Node& node, bool& modified,
const std::unordered_set<std::string>& excluded_initializer_names) {
const std::unordered_set<std::string>& excluded_initializer_names,
const std::unordered_set<std::string>& compatible_execution_providers) {
if (!graph_utils::IsSupportedProvider(node, compatible_execution_providers)) {
return Status::OK();
}
if (!graph_utils::IsSupportedOptypeVersionAndDomain(node, "MatMul", {9, 13}) &&
!graph_utils::IsSupportedOptypeVersionAndDomain(node, "TransposeScaleMatMul", {1}, kMSDomain)) {
!graph_utils::IsSupportedOptypeVersionAndDomain(node, "TransposeMatMul", {1}, kMSDomain)) {
return Status::OK();
}
if (!IsMatMulInputTypeSupported(node)) {
return Status::OK();
}
@ -185,7 +205,7 @@ Status ProcessNode(
}
NodeAttributes fused_node_attrs =
node.OpType() == "TransposeScaleMatMul" ? node.GetAttributes() : NodeAttributes{};
node.OpType() == "TransposeMatMul" ? node.GetAttributes() : NodeAttributes{};
{
ONNX_NAMESPACE::AttributeProto& alpha_attr = fused_node_attrs["alpha"];
@ -220,7 +240,7 @@ Status ProcessNode(
Node& matmul_scale_node = graph.AddNode(
graph.GenerateNodeName(node.Name() + "_FusedMatMulAndScale"),
"TransposeScaleMatMul",
"TransposeMatMul",
"Fused MatMul and Scale",
fused_node_inputs,
fused_node_outputs,
@ -273,7 +293,8 @@ Status MatMulScaleFusion::ApplyImpl(Graph& graph, bool& modified, int graph_leve
ORT_RETURN_IF_ERROR(Recurse(*node, modified, graph_level, logger));
ORT_RETURN_IF_ERROR(ProcessNode(graph, *node, modified, excluded_initializer_names_));
ORT_RETURN_IF_ERROR(ProcessNode(
graph, *node, modified, excluded_initializer_names_, GetCompatibleExecutionProviders()));
}
return Status::OK();

View file

@ -7,11 +7,11 @@ namespace onnxruntime {
/**
* Fuses MatMul with surrounding scales (multiplies or divides) by a constant
* scalar into TransposeScaleMatMul.
* scalar into TransposeMatMul (which supports scaling the product).
*
* 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)
* -> TransposeMatMul(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

View file

@ -97,7 +97,7 @@ Status MatmulTransposeFusion::ApplyImpl(Graph& graph, bool& modified, int graph_
ORT_RETURN_IF_ERROR(Recurse(node, modified, graph_level, logger));
if ((!graph_utils::IsSupportedOptypeVersionAndDomain(node, "MatMul", {9, 13}) &&
!graph_utils::IsSupportedOptypeVersionAndDomain(node, "TransposeScaleMatMul", {1, 13}, kMSDomain)) ||
!graph_utils::IsSupportedOptypeVersionAndDomain(node, "TransposeMatMul", {1, 13}, kMSDomain)) ||
!graph_utils::IsSupportedProvider(node, GetCompatibleExecutionProviders())) {
continue;
}
@ -130,14 +130,14 @@ Status MatmulTransposeFusion::ApplyImpl(Graph& graph, bool& modified, int graph_
const std::vector<NodeArg*> output_defs{node.MutableOutputDefs()[0]};
Node& matmul_node = graph.AddNode(graph.GenerateNodeName("MatMul_With_Transpose"),
"TransposeScaleMatMul",
"TransposeMatMul",
"fused MatMul and Transpose ",
input_defs,
output_defs, {}, kMSDomain);
bool transpose_left = (left != nullptr);
bool transpose_right = (right != nullptr);
float alpha = 1.0f;
if (node.OpType() == "TransposeScaleMatMul") {
if (node.OpType() == "TransposeMatMul") {
transpose_left ^= static_cast<bool>(node.GetAttributes().at("transA").i());
transpose_right ^= static_cast<bool>(node.GetAttributes().at("transB").i());
alpha = node.GetAttributes().at("alpha").f();

View file

@ -131,10 +131,10 @@ void ProcessInputs(const std::vector<int64_t>& input_dims, const std::vector<T>&
}
template <typename T>
void RunTransposeScaleMatMulTest(int32_t opset_version = 7, bool transa = false, bool transb = false, float alpha = 1.0f) {
void RunTransposeMatMulTest(int32_t opset_version = 7, bool transa = false, bool transb = false, float alpha = 1.0f) {
std::vector<T> common_input_vals{0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11};
for (auto t : GenerateSimpleTestCases<T>()) {
OpTester test("TransposeScaleMatMul", opset_version, onnxruntime::kMSDomain);
OpTester test("TransposeMatMul", opset_version, onnxruntime::kMSDomain);
std::vector<int64_t> input0_dims(t.input0_dims);
std::vector<T> input0_vals;
@ -166,31 +166,31 @@ void RunTransposeScaleMatMulTest(int32_t opset_version = 7, bool transa = false,
}
TEST(TransposeMatMulOpTest, FloatTypeNoTranspose) {
RunTransposeScaleMatMulTest<float>(1);
RunTransposeMatMulTest<float>(1);
}
#ifdef USE_CUDA // double support only implemented in CUDA kernel
TEST(TransposeMatMulOpTest, DoubleTypeNoTranspose) {
RunTransposeScaleMatMulTest<double>(1);
RunTransposeMatMulTest<double>(1);
}
#endif
TEST(TransposeMatMulOpTest, FloatTypeTransposeA) {
RunTransposeScaleMatMulTest<float>(1, true, false);
RunTransposeMatMulTest<float>(1, true, false);
}
TEST(TransposeMatMulOpTest, FloatTypeTransposeB) {
RunTransposeScaleMatMulTest<float>(1, false, true);
RunTransposeMatMulTest<float>(1, false, true);
}
TEST(TransposeMatMulOpTest, FloatTypeTransposeAB) {
RunTransposeScaleMatMulTest<float>(1, true, true);
RunTransposeMatMulTest<float>(1, true, true);
}
TEST(TransposeMatMulOpTest, FloatTypeScale) {
RunTransposeScaleMatMulTest<float>(1, false, false, 0.5f);
RunTransposeScaleMatMulTest<float>(1, true, false, 2.0f);
RunTransposeScaleMatMulTest<float>(1, true, true, 4.0f);
RunTransposeMatMulTest<float>(1, false, false, 0.5f);
RunTransposeMatMulTest<float>(1, true, false, 2.0f);
RunTransposeMatMulTest<float>(1, true, true, 4.0f);
}
} // namespace transpose_matmul

View file

@ -711,7 +711,7 @@ TEST_F(GraphTransformationTests, TransposeMatmulFusion) {
std::map<std::string, int> op_to_count = CountOpsInGraph(graph);
ASSERT_TRUE(op_to_count["Transpose"] == 0);
ASSERT_TRUE(op_to_count["MatMul"] == 0);
ASSERT_TRUE(op_to_count["TransposeScaleMatMul"] == 1);
ASSERT_TRUE(op_to_count["TransposeMatMul"] == 1);
}
TEST_F(GraphTransformationTests, TransposeMatmulFusionOnTwoTranspose) {
@ -728,10 +728,10 @@ TEST_F(GraphTransformationTests, TransposeMatmulFusionOnTwoTranspose) {
std::map<std::string, int> op_to_count = CountOpsInGraph(graph);
ASSERT_TRUE(op_to_count["Transpose"] == 0);
ASSERT_TRUE(op_to_count["MatMul"] == 0);
ASSERT_TRUE(op_to_count["TransposeScaleMatMul"] == 1);
ASSERT_TRUE(op_to_count["TransposeMatMul"] == 1);
auto& node = *graph.Nodes().begin();
ASSERT_TRUE(node.OpType() == "TransposeScaleMatMul");
ASSERT_TRUE(node.OpType() == "TransposeMatMul");
ASSERT_TRUE(static_cast<bool>(node.GetAttributes().at("transA").i()));
ASSERT_TRUE(static_cast<bool>(node.GetAttributes().at("transB").i()));
}
@ -750,10 +750,10 @@ TEST_F(GraphTransformationTests, TransposeMatmulFusionOnThreeTranspose) {
std::map<std::string, int> op_to_count = CountOpsInGraph(graph);
ASSERT_TRUE(op_to_count["Transpose"] == 0);
ASSERT_TRUE(op_to_count["MatMul"] == 0);
ASSERT_TRUE(op_to_count["TransposeScaleMatMul"] == 1);
ASSERT_TRUE(op_to_count["TransposeMatMul"] == 1);
auto& node = *graph.Nodes().begin();
ASSERT_TRUE(node.OpType() == "TransposeScaleMatMul");
ASSERT_TRUE(node.OpType() == "TransposeMatMul");
ASSERT_FALSE(static_cast<bool>(node.GetAttributes().at("transA").i()));
ASSERT_TRUE(static_cast<bool>(node.GetAttributes().at("transB").i()));
}
@ -776,11 +776,11 @@ TEST_F(GraphTransformationTests, TransposeMatmulNoFusionOnInvalidPerm) {
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);
ASSERT_EQ(op_to_count["TransposeMatMul"], 0);
}
}
TEST_F(GraphTransformationTests, TransposeMatmulFusionFromTransposeScaleMatMul) {
TEST_F(GraphTransformationTests, TransposeMatmulFusionFromTransposeMatMul) {
auto model_uri = MODEL_FOLDER "fusion/transpose_matmul_2d_fusion_from_transpose_scale_matmul.onnx";
std::shared_ptr<Model> p_model;
ASSERT_STATUS_OK(Model::Load(model_uri, p_model, nullptr, *logger_));
@ -791,7 +791,7 @@ TEST_F(GraphTransformationTests, TransposeMatmulFusionFromTransposeScaleMatMul)
auto transpose_scale_matmul_node =
std::find_if(
graph.Nodes().cbegin(), graph.Nodes().cend(),
[](const Node& node) { return node.Name() == "TransposeScaleMatMul"; });
[](const Node& node) { return node.Name() == "TransposeMatMul"; });
ASSERT_NE(transpose_scale_matmul_node, graph.Nodes().cend());
expected_alpha = transpose_scale_matmul_node->GetAttributes().at("alpha").f();
}
@ -804,10 +804,10 @@ TEST_F(GraphTransformationTests, TransposeMatmulFusionFromTransposeScaleMatMul)
std::map<std::string, int> op_to_count = CountOpsInGraph(graph);
ASSERT_EQ(op_to_count["Transpose"], 0);
ASSERT_EQ(op_to_count["MatMul"], 0);
ASSERT_EQ(op_to_count["TransposeScaleMatMul"], 1);
ASSERT_EQ(op_to_count["TransposeMatMul"], 1);
auto& transpose_scale_matmul_node = *graph.Nodes().begin();
ASSERT_EQ(transpose_scale_matmul_node.OpType(), "TransposeScaleMatMul");
ASSERT_EQ(transpose_scale_matmul_node.OpType(), "TransposeMatMul");
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);
@ -827,7 +827,7 @@ TEST_F(GraphTransformationTests, TransposeMatmulFusionWithPreservedTranspose) {
std::map<std::string, int> op_to_count = CountOpsInGraph(graph);
ASSERT_EQ(op_to_count["Transpose"], 1);
ASSERT_EQ(op_to_count["MatMul"], 0);
ASSERT_EQ(op_to_count["TransposeScaleMatMul"], 1);
ASSERT_EQ(op_to_count["TransposeMatMul"], 1);
ASSERT_FALSE(graph.GraphResolveNeeded());
}
@ -3120,10 +3120,12 @@ TEST_F(GraphTransformationTests, ComputationReductionTransformer_GatherND_E2E) {
#endif
#ifndef DISABLE_CONTRIB_OPS
template <typename GraphTransformationCheckFn>
template <typename GraphTransformationCheckFn, typename GraphPreprocessFn>
static void TestMatMulScaleFusion(
const PathString& model_path, const Logger& logger,
GraphTransformationCheckFn graph_transformation_check,
GraphPreprocessFn graph_preprocess_fn,
GraphTransformationCheckFn graph_transformation_check_fn,
const std::unordered_set<std::string>& compatible_execution_providers = {},
const std::unordered_set<std::string>& excluded_initializer_names = {}) {
SCOPED_TRACE(ORT_TSTR("model path: ") + model_path);
@ -3131,17 +3133,31 @@ static void TestMatMulScaleFusion(
ASSERT_STATUS_OK(Model::Load(model_path, model, nullptr, logger));
Graph& graph = model->MainGraph();
graph_preprocess_fn(graph);
auto original_op_counts = CountOpsInGraph(graph);
onnxruntime::GraphTransformerManager graph_transformer_manager{5};
ASSERT_STATUS_OK(graph_transformer_manager.Register(
make_unique<MatMulScaleFusion>(std::unordered_set<std::string>{}, excluded_initializer_names),
make_unique<MatMulScaleFusion>(compatible_execution_providers, excluded_initializer_names),
TransformerLevel::Level2));
ASSERT_STATUS_OK(graph_transformer_manager.ApplyTransformers(graph, TransformerLevel::Level2, logger));
auto transformed_op_counts = CountOpsInGraph(graph);
graph_transformation_check(graph, original_op_counts, transformed_op_counts);
graph_transformation_check_fn(graph, original_op_counts, transformed_op_counts);
}
template <typename GraphTransformationCheckFn>
static void TestMatMulScaleFusion(
const PathString& model_path, const Logger& logger,
GraphTransformationCheckFn graph_transformation_check,
const std::unordered_set<std::string>& compatible_execution_providers = {},
const std::unordered_set<std::string>& excluded_initializer_names = {}) {
TestMatMulScaleFusion(
model_path, logger,
[](Graph&) {}, graph_transformation_check,
compatible_execution_providers, excluded_initializer_names);
}
TEST_F(GraphTransformationTests, MatMulScaleFusionFusableModels) {
@ -3161,17 +3177,17 @@ TEST_F(GraphTransformationTests, MatMulScaleFusionFusableModels) {
EXPECT_EQ(transformed_op_counts["Mul"], 0);
EXPECT_EQ(transformed_op_counts["Div"], 0);
EXPECT_EQ(transformed_op_counts["MatMul"], 0);
EXPECT_EQ(transformed_op_counts["TransposeScaleMatMul"], 1);
EXPECT_EQ(transformed_op_counts["TransposeMatMul"], 1);
// check combined scale, individual scales should all have the same value
const float scale_value = 3.0f;
const int num_scales =
original_op_counts["Mul"] + original_op_counts["Div"] + original_op_counts["TransposeScaleMatMul"];
original_op_counts["Mul"] + original_op_counts["Div"] + original_op_counts["TransposeMatMul"];
auto fused_node = std::find_if(
graph.Nodes().cbegin(), graph.Nodes().cend(),
[](const Node& node) { return node.OpType() == "TransposeScaleMatMul"; });
[](const Node& node) { return node.OpType() == "TransposeMatMul"; });
ASSERT_NE(fused_node, graph.Nodes().cend());
auto alpha_attr = fused_node->GetAttributes().find("alpha");
@ -3209,7 +3225,7 @@ TEST_F(GraphTransformationTests, MatMulScaleFusionReusedInputScale) {
EXPECT_EQ(transformed_op_counts["Mul"], 0);
EXPECT_EQ(transformed_op_counts["Div"], 0);
EXPECT_EQ(transformed_op_counts["MatMul"], 0);
EXPECT_EQ(transformed_op_counts["TransposeScaleMatMul"], 2);
EXPECT_EQ(transformed_op_counts["TransposeMatMul"], 2);
});
}
@ -3221,8 +3237,41 @@ TEST_F(GraphTransformationTests, MatMulScaleFusionExcludedInitializerName) {
const std::map<std::string, int>& transformed_op_counts) {
EXPECT_EQ(original_op_counts, transformed_op_counts);
},
{},
{"scale"});
}
TEST_F(GraphTransformationTests, MatMulScaleFusionIncompatibleExecutionProvider) {
TestMatMulScaleFusion(
MODEL_FOLDER "fusion/matmul_scale_in0.onnx", *logger_,
[](Graph& graph) {
for (auto& node : graph.Nodes()) {
node.SetExecutionProviderType(kCudaExecutionProvider);
}
},
[](const Graph&,
const std::map<std::string, int>& original_op_counts,
const std::map<std::string, int>& transformed_op_counts) {
EXPECT_EQ(original_op_counts, transformed_op_counts);
},
{kCpuExecutionProvider});
}
TEST_F(GraphTransformationTests, MatMulScaleFusionUnsupportedInputType) {
TestMatMulScaleFusion(
MODEL_FOLDER "fusion/matmul_scale_int32.onnx", *logger_,
[](Graph& graph) {
for (auto& node : graph.Nodes()) {
node.SetExecutionProviderType(kCpuExecutionProvider);
}
},
[](const Graph&,
const std::map<std::string, int>& original_op_counts,
const std::map<std::string, int>& transformed_op_counts) {
EXPECT_EQ(original_op_counts, transformed_op_counts);
},
{kCpuExecutionProvider});
}
#endif
} // namespace test

View file

@ -31,7 +31,7 @@ def save(model_path, nodes, inputs, outputs, initializers):
def gen(model_path,
use_transpose_matmul,
scale_input_0, scale_input_1, scale_output):
matmul_op = "TransposeScaleMatMul" if use_transpose_matmul else "MatMul"
matmul_op = "TransposeMatMul" if use_transpose_matmul else "MatMul"
matmul_domain = "com.microsoft" if use_transpose_matmul else ""
matmul_attrs = {"alpha": scale_value} if use_transpose_matmul else {}
@ -184,3 +184,37 @@ def gen_reused_input_scale(model_path):
gen_reused_input_scale("matmul_scale_reused_input_scale.onnx")
def gen_int32(model_path):
matmul_op = "MatMul"
nodes = [
helper.make_node(
"Mul", ["input_0", "scale"], ["scaled_input_0"],
"scale input_0"),
helper.make_node(
matmul_op, ["scaled_input_0", "input_1"], ["output_0"],
"MatMul input_0 and input_1"),
]
initializers = [
helper.make_tensor("scale", TensorProto.INT32, [], [int(scale_value)])
]
inputs = [
helper.make_tensor_value_info(
"input_0", TensorProto.INT32, [2, 'M', 'K']),
helper.make_tensor_value_info(
"input_1", TensorProto.INT32, [2, 'K', 'N']),
]
outputs = [
helper.make_tensor_value_info(
"output_0", TensorProto.INT32, [2, 'M', 'N']),
]
save(model_path, nodes, inputs, outputs, initializers)
gen_int32("matmul_scale_int32.onnx")

Binary file not shown.

View file

@ -32,10 +32,10 @@ def gen_from_transpose_scale_matmul(model_path):
["input_0"],
["transposed_input_0"]),
helper.make_node(
"TransposeScaleMatMul",
"TransposeMatMul",
["transposed_input_0", "input_1"],
["output"],
"TransposeScaleMatMul",
"TransposeMatMul",
"",
msdomain.domain,
alpha=3.0, transA=1)

View file

@ -342,7 +342,7 @@ IMPLEMENT_GRADIENT_BUILDER(GetMatMulGradient) {
if (IsGradientRequiredForSrcNodeInput(0)) {
ArgDef pre_reduce_grad_0 = IA("PreReduceGrad0");
result.push_back(
NodeDef(OpDef{"TransposeScaleMatMul", kMSDomain, 1},
NodeDef(OpDef{"TransposeMatMul", kMSDomain, 1},
{GO(0), B},
{pre_reduce_grad_0},
{{"transB", MakeAttribute("transB", int64_t(1))}}));
@ -360,7 +360,7 @@ IMPLEMENT_GRADIENT_BUILDER(GetMatMulGradient) {
} else {
ArgDef pre_reduce_grad_1 = IA("PreReduceGrad1");
result.push_back(
NodeDef(OpDef{"TransposeScaleMatMul", kMSDomain, 1},
NodeDef(OpDef{"TransposeMatMul", kMSDomain, 1},
{A, GO(0)},
{pre_reduce_grad_1},
{{"transA", MakeAttribute("transA", int64_t(1))}}));