mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
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:
parent
65740deb10
commit
e9671e93f0
18 changed files with 168 additions and 65 deletions
|
|
@ -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)>,
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
|
|
@ -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)>,
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ namespace cuda {
|
|||
|
||||
#define REGISTER_KERNEL_TYPED(T) \
|
||||
ONNX_OPERATOR_TYPED_KERNEL_EX( \
|
||||
TransposeScaleMatMul, \
|
||||
TransposeMatMul, \
|
||||
kMSDomain, \
|
||||
1, \
|
||||
T, \
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
BIN
onnxruntime/test/testdata/transform/fusion/matmul_scale_int32.onnx
vendored
Normal file
BIN
onnxruntime/test/testdata/transform/fusion/matmul_scale_int32.onnx
vendored
Normal file
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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))}}));
|
||||
|
|
|
|||
Loading…
Reference in a new issue