diff --git a/onnxruntime/core/graph/conv_add_fusion.cc b/onnxruntime/core/graph/conv_add_fusion.cc index 89333e275b..095371c35e 100644 --- a/onnxruntime/core/graph/conv_add_fusion.cc +++ b/onnxruntime/core/graph/conv_add_fusion.cc @@ -36,12 +36,25 @@ Status ConvAddFusion::Apply(onnxruntime::Graph& graph, bool& modified) const { // Currently, fusion is only supported for float or double data type. if (!Initializer::IsSupportedDataType(add_B_tensor_proto) || - conv_W_tensor_proto->dims_size() != 4 || - add_B_tensor_proto->dims_size() != 3 || + conv_W_tensor_proto->dims_size() < 4 || + add_B_tensor_proto->dims_size() != conv_W_tensor_proto->dims_size() - 1 || conv_W_tensor_proto->dims(0) != add_B_tensor_proto->dims(0)) { continue; } + // The dimensions of add_B should be equal to 1 except first dimension. + bool flag = false; + for (int i = 1; i < add_B_tensor_proto->dims_size(); i++) { + if (add_B_tensor_proto->dims(i) != 1) { + flag = true; + break; + } + } + + if (flag) { + continue; + } + const ONNX_NAMESPACE::TensorProto* conv_B_tensor_proto = nullptr; if (conv_inputs.size() == 3) { graph.GetInitializedTensor(conv_inputs[2]->Name(), conv_B_tensor_proto); @@ -49,7 +62,6 @@ Status ConvAddFusion::Apply(onnxruntime::Graph& graph, bool& modified) const { if (!Initializer::IsSupportedDataType(conv_B_tensor_proto) || conv_B_tensor_proto->data_type() != add_B_tensor_proto->data_type() || conv_B_tensor_proto->dims_size() != 1 || - add_B_tensor_proto->dims_size() != 3 || conv_B_tensor_proto->dims(0) != add_B_tensor_proto->dims(0)) { continue; } @@ -57,6 +69,9 @@ Status ConvAddFusion::Apply(onnxruntime::Graph& graph, bool& modified) const { auto conv_B = std::make_unique(conv_B_tensor_proto); auto add_B = std::make_unique(add_B_tensor_proto); + if (conv_B->size() != add_B->size()) { + continue; + } // Calculate new value of initializers of conv node conv_B->add(*add_B); diff --git a/onnxruntime/core/graph/conv_mul_fusion.cc b/onnxruntime/core/graph/conv_mul_fusion.cc index 447a57fb80..8fb6450487 100644 --- a/onnxruntime/core/graph/conv_mul_fusion.cc +++ b/onnxruntime/core/graph/conv_mul_fusion.cc @@ -37,18 +37,30 @@ Status ConvMulFusion::Apply(onnxruntime::Graph& graph, bool& modified) const { if (!Initializer::IsSupportedDataType(conv_W_tensor_proto) || !Initializer::IsSupportedDataType(mul_B_tensor_proto) || conv_W_tensor_proto->data_type() != mul_B_tensor_proto->data_type() || - !(conv_W_tensor_proto->dims_size() > 2 && conv_W_tensor_proto->dims(0) == mul_B_tensor_proto->dims(0))) { + conv_W_tensor_proto->dims_size() < 4 || + !(mul_B_tensor_proto->dims_size() == 0 || + (mul_B_tensor_proto->dims_size() == conv_W_tensor_proto->dims_size() - 1 && + conv_W_tensor_proto->dims(0) == mul_B_tensor_proto->dims(0)))) { continue; } + // The dimensions of mul_B should be equal to 1 except first dimension. + if (mul_B_tensor_proto->dims_size() != 0) { + bool flag = false; + for (int i = 1; i < mul_B_tensor_proto->dims_size(); i++) { + if (mul_B_tensor_proto->dims(i) != 1) { + flag = true; + break; + } + } + + if (flag) { + continue; + } + } auto conv_W = std::make_unique(conv_W_tensor_proto); auto mul_B = std::make_unique(mul_B_tensor_proto); - if (conv_W->data_type() != mul_B->data_type() || - !(conv_W->dims().size() > 2 && conv_W->dims()[0] == mul_B->dims()[0])) { - continue; - } - const ONNX_NAMESPACE::TensorProto* conv_B_tensor_proto = nullptr; std::unique_ptr conv_B = nullptr; if (conv_inputs.size() == 3) { @@ -57,8 +69,8 @@ Status ConvMulFusion::Apply(onnxruntime::Graph& graph, bool& modified) const { if (!Initializer::IsSupportedDataType(conv_B_tensor_proto) || conv_B_tensor_proto->data_type() != mul_B_tensor_proto->data_type() || - conv_B_tensor_proto->dims_size() != 1 || mul_B_tensor_proto->dims_size() != 3 || - conv_B_tensor_proto->dims(0) != mul_B_tensor_proto->dims(0)) { + conv_B_tensor_proto->dims_size() != 1 || (mul_B_tensor_proto->dims_size() != 0 && + conv_B_tensor_proto->dims(0) != mul_B_tensor_proto->dims(0))) { continue; } conv_B = std::make_unique(conv_B_tensor_proto); @@ -66,8 +78,13 @@ Status ConvMulFusion::Apply(onnxruntime::Graph& graph, bool& modified) const { // Calculate new value of initializers of conv node conv_W->scale_by_axis(*mul_B, 1); + if (conv_inputs.size() == 3) { - conv_B->mul(*mul_B); + if (mul_B_tensor_proto->dims_size() != 0) { + conv_B->mul(*mul_B); + } else { + conv_B->scale_by_axis(*mul_B, 0); + } } // Create new initializers of conv diff --git a/onnxruntime/core/graph/initializer.h b/onnxruntime/core/graph/initializer.h index 3fda9532d3..67171f84fd 100644 --- a/onnxruntime/core/graph/initializer.h +++ b/onnxruntime/core/graph/initializer.h @@ -190,19 +190,22 @@ class Initializer final { return dims_; } - size_t size() const { return size_; } + int64_t size() const { return size_; } Initializer& add(float value) { + int64_t n = size(); switch (data_type_) { case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: { - for (int i = 0; i < size_; i++) { - data()[i] += value; + float* dst = data(); + for (int i = 0; i < n; i++) { + dst[i] += value; } break; } case ONNX_NAMESPACE::TensorProto_DataType_DOUBLE: { - for (int i = 0; i < size_; i++) { - data()[i] += value; + double* dst = data(); + for (int i = 0; i < n; i++) { + dst[i] += value; } break; } @@ -213,16 +216,21 @@ class Initializer final { } Initializer& add(const Initializer& other) { + int64_t n = size(); switch (data_type_) { case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: { - for (int i = 0; i < size_; i++) { - data()[i] += other.data()[i]; + float* dst = data(); + const float* src = other.data(); + for (int i = 0; i < n; i++) { + dst[i] += src[i]; } break; } case ONNX_NAMESPACE::TensorProto_DataType_DOUBLE: { - for (int i = 0; i < size_; i++) { - data()[i] += other.data()[i]; + double* dst = data(); + const double* src = other.data(); + for (int i = 0; i < n; i++) { + dst[i] += src[i]; } break; } @@ -232,16 +240,21 @@ class Initializer final { return *this; } Initializer& sub(const Initializer& other) { + int64_t n = size(); switch (data_type_) { case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: { - for (int i = 0; i < size_; i++) { - data()[i] -= other.data()[i]; + float* dst = data(); + const float* src = other.data(); + for (int i = 0; i < n; i++) { + dst[i] -= src[i]; } break; } case ONNX_NAMESPACE::TensorProto_DataType_DOUBLE: { - for (int i = 0; i < size_; i++) { - data()[i] -= other.data()[i]; + double* dst = data(); + const double* src = other.data(); + for (int i = 0; i < n; i++) { + dst[i] -= src[i]; } break; } @@ -252,16 +265,21 @@ class Initializer final { } Initializer& mul(const Initializer& other) { + int64_t n = size(); switch (data_type_) { case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: { - for (int i = 0; i < size_; i++) { - data()[i] *= other.data()[i]; + float* dst = data(); + const float* src = other.data(); + for (int i = 0; i < n; i++) { + dst[i] *= src[i]; } break; } case ONNX_NAMESPACE::TensorProto_DataType_DOUBLE: { - for (int i = 0; i < size_; i++) { - data()[i] *= other.data()[i]; + double* dst = data(); + const double* src = other.data(); + for (int i = 0; i < n; i++) { + dst[i] *= src[i]; } break; } @@ -271,16 +289,21 @@ class Initializer final { return *this; } Initializer& div(const Initializer& other) { + int64_t n = size(); switch (data_type_) { case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: { - for (int i = 0; i < size_; i++) { - data()[i] /= other.data()[i]; + float* dst = data(); + const float* src = other.data(); + for (int i = 0; i < n; i++) { + dst[i] /= src[i]; } break; } case ONNX_NAMESPACE::TensorProto_DataType_DOUBLE: { - for (int i = 0; i < size_; i++) { - data()[i] /= other.data()[i]; + double* dst = data(); + const double* src = other.data(); + for (int i = 0; i < n; i++) { + dst[i] /= src[i]; } break; } @@ -291,16 +314,19 @@ class Initializer final { } Initializer& sqrt() { + int64_t n = size(); switch (data_type_) { case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: { - for (int i = 0; i < size_; i++) { - data()[i] = std::sqrt(data()[i]); + float* dst = data(); + for (int i = 0; i < n; i++) { + dst[i] = std::sqrt(dst[i]); } break; } case ONNX_NAMESPACE::TensorProto_DataType_DOUBLE: { - for (int i = 0; i < size_; i++) { - data()[i] = std::sqrt(data()[i]); + double* dst = data(); + for (int i = 0; i < n; i++) { + dst[i] = std::sqrt(dst[i]); } break; } @@ -316,19 +342,26 @@ class Initializer final { num *= dims_[k]; } + int64_t n = size()/num; switch (data_type_) { case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: { - for (int64_t i = 0; i < dims_[0]; i++) { + float* dst = data(); + const float* src = other.data(); + for (int i = 0; i < n; i++) { + int index = other.size() == 1 ? 0 : i; for (int64_t j = 0; j < num; j++) { - data()[i * num + j] *= other.data()[i]; + dst[i * num + j] *= src[index]; } } break; } case ONNX_NAMESPACE::TensorProto_DataType_DOUBLE: { - for (int64_t i = 0; i < dims_[0]; i++) { + double* dst = data(); + const double* src = other.data(); + for (int i = 0; i < n; i++) { + int index = other.size() == 1 ? 0 : i; for (int64_t j = 0; j < num; j++) { - data()[i * num + j] *= other.data()[i]; + dst[i * num + j] *= src[index]; } } break; diff --git a/onnxruntime/test/ir/graph_transform_test.cc b/onnxruntime/test/ir/graph_transform_test.cc index 3d1c754dce..365e92eb3f 100644 --- a/onnxruntime/test/ir/graph_transform_test.cc +++ b/onnxruntime/test/ir/graph_transform_test.cc @@ -71,5 +71,107 @@ TEST(GraphTransformationTests, FuseConvBNMulAddUnsqueeze) { ASSERT_TRUE(session_object.Initialize().IsOK()); } +TEST(GraphTransformationTests, FuseConvBNNoBias) { + string model_uri = MODEL_FOLDER + "fusion/fuse-conv-bn-no-bias.onnx"; + + SessionOptions so; + so.session_logid = "GraphTransformationTests.LoadModelToTransform"; + InferenceSession session_object{so, &DefaultLoggingManager()}; + ASSERT_TRUE(session_object.Load(model_uri).IsOK()); + + std::shared_ptr p_model; + ASSERT_TRUE(Model::Load(model_uri, p_model).IsOK()); + + std::unique_ptr ConvBNFusion_transformer = std::make_unique(); + + session_object.RegisterGraphTransformer(std::move(ConvBNFusion_transformer)); + + ASSERT_TRUE(session_object.Initialize().IsOK()); +} + +TEST(GraphTransformationTests, FuseConvMulNoBias) { + string model_uri = MODEL_FOLDER + "fusion/fuse-conv-mul-no-bias.onnx"; + + SessionOptions so; + so.session_logid = "GraphTransformationTests.LoadModelToTransform"; + InferenceSession session_object{so, &DefaultLoggingManager()}; + ASSERT_TRUE(session_object.Load(model_uri).IsOK()); + + std::shared_ptr p_model; + ASSERT_TRUE(Model::Load(model_uri, p_model).IsOK()); + + std::unique_ptr Unsqueeze_transformer = std::make_unique(); + std::unique_ptr ConvMulFusion_transformer = std::make_unique(); + + session_object.RegisterGraphTransformer(std::move(Unsqueeze_transformer)); + session_object.RegisterGraphTransformer(std::move(ConvMulFusion_transformer)); + + ASSERT_TRUE(session_object.Initialize().IsOK()); +} + +TEST(GraphTransformationTests, FuseConvAddNoBias) { + string model_uri = MODEL_FOLDER + "fusion/fuse-conv-add-no-bias.onnx"; + + SessionOptions so; + so.session_logid = "GraphTransformationTests.LoadModelToTransform"; + InferenceSession session_object{so, &DefaultLoggingManager()}; + ASSERT_TRUE(session_object.Load(model_uri).IsOK()); + + std::shared_ptr p_model; + ASSERT_TRUE(Model::Load(model_uri, p_model).IsOK()); + + std::unique_ptr Unsqueeze_transformer = std::make_unique(); + std::unique_ptr ConvAddFusion_transformer = std::make_unique(); + + session_object.RegisterGraphTransformer(std::move(Unsqueeze_transformer)); + session_object.RegisterGraphTransformer(std::move(ConvAddFusion_transformer)); + + ASSERT_TRUE(session_object.Initialize().IsOK()); +} + +TEST(GraphTransformationTests, FuseConvBNMulAddUnsqueezeNoBias) { + string model_uri = MODEL_FOLDER + "fusion/fuse-conv-bn-mul-add-unsqueeze-no-bias.onnx"; + + SessionOptions so; + so.session_logid = "GraphTransformationTests.LoadModelToTransform"; + InferenceSession session_object{so, &DefaultLoggingManager()}; + ASSERT_TRUE(session_object.Load(model_uri).IsOK()); + + std::shared_ptr p_model; + ASSERT_TRUE(Model::Load(model_uri, p_model).IsOK()); + + std::unique_ptr Unsqueeze_transformer = std::make_unique(); + std::unique_ptr ConvBNFusion_transformer = std::make_unique(); + std::unique_ptr ConvMulFusion_transformer = std::make_unique(); + std::unique_ptr ConvAddFusion_transformer = std::make_unique(); + + session_object.RegisterGraphTransformer(std::move(Unsqueeze_transformer)); + session_object.RegisterGraphTransformer(std::move(ConvBNFusion_transformer)); + session_object.RegisterGraphTransformer(std::move(ConvMulFusion_transformer)); + session_object.RegisterGraphTransformer(std::move(ConvAddFusion_transformer)); + + ASSERT_TRUE(session_object.Initialize().IsOK()); +} + +TEST(GraphTransformationTests, FuseConvAddMul3D) { + string model_uri = MODEL_FOLDER + "fusion/fuse-conv-add-mul-3d.onnx"; + + SessionOptions so; + so.session_logid = "GraphTransformationTests.LoadModelToTransform"; + InferenceSession session_object{so, &DefaultLoggingManager()}; + ASSERT_TRUE(session_object.Load(model_uri).IsOK()); + + std::shared_ptr p_model; + ASSERT_TRUE(Model::Load(model_uri, p_model).IsOK()); + + std::unique_ptr ConvMulFusion_transformer = std::make_unique(); + std::unique_ptr ConvAddFusion_transformer = std::make_unique(); + + session_object.RegisterGraphTransformer(std::move(ConvMulFusion_transformer)); + session_object.RegisterGraphTransformer(std::move(ConvAddFusion_transformer)); + + ASSERT_TRUE(session_object.Initialize().IsOK()); +} + } // namespace test } // namespace onnxruntime diff --git a/onnxruntime/test/testdata/transform/fusion/fuse-conv-add-mul-3d.onnx b/onnxruntime/test/testdata/transform/fusion/fuse-conv-add-mul-3d.onnx new file mode 100644 index 0000000000..20de2940c2 Binary files /dev/null and b/onnxruntime/test/testdata/transform/fusion/fuse-conv-add-mul-3d.onnx differ diff --git a/onnxruntime/test/testdata/transform/fusion/fuse-conv-add-no-bias.onnx b/onnxruntime/test/testdata/transform/fusion/fuse-conv-add-no-bias.onnx new file mode 100644 index 0000000000..f2df9e1022 Binary files /dev/null and b/onnxruntime/test/testdata/transform/fusion/fuse-conv-add-no-bias.onnx differ diff --git a/onnxruntime/test/testdata/transform/fusion/fuse-conv-bn-mul-add-unsqueeze-no-bias.onnx b/onnxruntime/test/testdata/transform/fusion/fuse-conv-bn-mul-add-unsqueeze-no-bias.onnx new file mode 100644 index 0000000000..584fbdf535 Binary files /dev/null and b/onnxruntime/test/testdata/transform/fusion/fuse-conv-bn-mul-add-unsqueeze-no-bias.onnx differ diff --git a/onnxruntime/test/testdata/transform/fusion/fuse-conv-bn-no-bias.onnx b/onnxruntime/test/testdata/transform/fusion/fuse-conv-bn-no-bias.onnx new file mode 100644 index 0000000000..ee8928b19e Binary files /dev/null and b/onnxruntime/test/testdata/transform/fusion/fuse-conv-bn-no-bias.onnx differ diff --git a/onnxruntime/test/testdata/transform/fusion/fuse-conv-mul-no-bias.onnx b/onnxruntime/test/testdata/transform/fusion/fuse-conv-mul-no-bias.onnx new file mode 100644 index 0000000000..ed4c77f313 Binary files /dev/null and b/onnxruntime/test/testdata/transform/fusion/fuse-conv-mul-no-bias.onnx differ