mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-28 20:11:22 +00:00
Support fusing 3D Conv with Add/Mul. (#23)
* Support fusing 3D Conv with Add/Mul. With this PR, the subgraph 3D Conve->Add->Mul in Resnet3D can be fused into one 3D Conv. * Updated it based on feedback. * Updated it based on review feedback. * Change the implementation of scale_by_axis rather than Mul. * Refactor the code to make the compiler optimize it easily.
This commit is contained in:
parent
b7cc611563
commit
aa549cd194
9 changed files with 208 additions and 41 deletions
|
|
@ -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<Initializer>(conv_B_tensor_proto);
|
||||
auto add_B = std::make_unique<Initializer>(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);
|
||||
|
||||
|
|
|
|||
|
|
@ -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<Initializer>(conv_W_tensor_proto);
|
||||
auto mul_B = std::make_unique<Initializer>(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<Initializer> 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<Initializer>(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
|
||||
|
|
|
|||
|
|
@ -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<float>()[i] += value;
|
||||
float* dst = data<float>();
|
||||
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<double>()[i] += value;
|
||||
double* dst = data<double>();
|
||||
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<float>()[i] += other.data<float>()[i];
|
||||
float* dst = data<float>();
|
||||
const float* src = other.data<float>();
|
||||
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<double>()[i] += other.data<double>()[i];
|
||||
double* dst = data<double>();
|
||||
const double* src = other.data<double>();
|
||||
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<float>()[i] -= other.data<float>()[i];
|
||||
float* dst = data<float>();
|
||||
const float* src = other.data<float>();
|
||||
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<double>()[i] -= other.data<double>()[i];
|
||||
double* dst = data<double>();
|
||||
const double* src = other.data<double>();
|
||||
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<float>()[i] *= other.data<float>()[i];
|
||||
float* dst = data<float>();
|
||||
const float* src = other.data<float>();
|
||||
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<double>()[i] *= other.data<double>()[i];
|
||||
double* dst = data<double>();
|
||||
const double* src = other.data<double>();
|
||||
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<float>()[i] /= other.data<float>()[i];
|
||||
float* dst = data<float>();
|
||||
const float* src = other.data<float>();
|
||||
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<double>()[i] /= other.data<double>()[i];
|
||||
double* dst = data<double>();
|
||||
const double* src = other.data<double>();
|
||||
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<float>()[i] = std::sqrt(data<float>()[i]);
|
||||
float* dst = data<float>();
|
||||
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<double>()[i] = std::sqrt(data<double>()[i]);
|
||||
double* dst = data<double>();
|
||||
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<float>();
|
||||
const float* src = other.data<float>();
|
||||
for (int i = 0; i < n; i++) {
|
||||
int index = other.size() == 1 ? 0 : i;
|
||||
for (int64_t j = 0; j < num; j++) {
|
||||
data<float>()[i * num + j] *= other.data<float>()[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<double>();
|
||||
const double* src = other.data<double>();
|
||||
for (int i = 0; i < n; i++) {
|
||||
int index = other.size() == 1 ? 0 : i;
|
||||
for (int64_t j = 0; j < num; j++) {
|
||||
data<double>()[i * num + j] *= other.data<double>()[i];
|
||||
dst[i * num + j] *= src[index];
|
||||
}
|
||||
}
|
||||
break;
|
||||
|
|
|
|||
|
|
@ -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<Model> p_model;
|
||||
ASSERT_TRUE(Model::Load(model_uri, p_model).IsOK());
|
||||
|
||||
std::unique_ptr<ConvBNFusion> ConvBNFusion_transformer = std::make_unique<ConvBNFusion>();
|
||||
|
||||
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<Model> p_model;
|
||||
ASSERT_TRUE(Model::Load(model_uri, p_model).IsOK());
|
||||
|
||||
std::unique_ptr<UnsqueezeElimination> Unsqueeze_transformer = std::make_unique<UnsqueezeElimination>();
|
||||
std::unique_ptr<ConvMulFusion> ConvMulFusion_transformer = std::make_unique<ConvMulFusion>();
|
||||
|
||||
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<Model> p_model;
|
||||
ASSERT_TRUE(Model::Load(model_uri, p_model).IsOK());
|
||||
|
||||
std::unique_ptr<UnsqueezeElimination> Unsqueeze_transformer = std::make_unique<UnsqueezeElimination>();
|
||||
std::unique_ptr<ConvAddFusion> ConvAddFusion_transformer = std::make_unique<ConvAddFusion>();
|
||||
|
||||
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<Model> p_model;
|
||||
ASSERT_TRUE(Model::Load(model_uri, p_model).IsOK());
|
||||
|
||||
std::unique_ptr<UnsqueezeElimination> Unsqueeze_transformer = std::make_unique<UnsqueezeElimination>();
|
||||
std::unique_ptr<ConvBNFusion> ConvBNFusion_transformer = std::make_unique<ConvBNFusion>();
|
||||
std::unique_ptr<ConvMulFusion> ConvMulFusion_transformer = std::make_unique<ConvMulFusion>();
|
||||
std::unique_ptr<ConvAddFusion> ConvAddFusion_transformer = std::make_unique<ConvAddFusion>();
|
||||
|
||||
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<Model> p_model;
|
||||
ASSERT_TRUE(Model::Load(model_uri, p_model).IsOK());
|
||||
|
||||
std::unique_ptr<ConvMulFusion> ConvMulFusion_transformer = std::make_unique<ConvMulFusion>();
|
||||
std::unique_ptr<ConvAddFusion> ConvAddFusion_transformer = std::make_unique<ConvAddFusion>();
|
||||
|
||||
session_object.RegisterGraphTransformer(std::move(ConvMulFusion_transformer));
|
||||
session_object.RegisterGraphTransformer(std::move(ConvAddFusion_transformer));
|
||||
|
||||
ASSERT_TRUE(session_object.Initialize().IsOK());
|
||||
}
|
||||
|
||||
} // namespace test
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
BIN
onnxruntime/test/testdata/transform/fusion/fuse-conv-add-mul-3d.onnx
vendored
Normal file
BIN
onnxruntime/test/testdata/transform/fusion/fuse-conv-add-mul-3d.onnx
vendored
Normal file
Binary file not shown.
BIN
onnxruntime/test/testdata/transform/fusion/fuse-conv-add-no-bias.onnx
vendored
Normal file
BIN
onnxruntime/test/testdata/transform/fusion/fuse-conv-add-no-bias.onnx
vendored
Normal file
Binary file not shown.
BIN
onnxruntime/test/testdata/transform/fusion/fuse-conv-bn-mul-add-unsqueeze-no-bias.onnx
vendored
Normal file
BIN
onnxruntime/test/testdata/transform/fusion/fuse-conv-bn-mul-add-unsqueeze-no-bias.onnx
vendored
Normal file
Binary file not shown.
BIN
onnxruntime/test/testdata/transform/fusion/fuse-conv-bn-no-bias.onnx
vendored
Normal file
BIN
onnxruntime/test/testdata/transform/fusion/fuse-conv-bn-no-bias.onnx
vendored
Normal file
Binary file not shown.
BIN
onnxruntime/test/testdata/transform/fusion/fuse-conv-mul-no-bias.onnx
vendored
Normal file
BIN
onnxruntime/test/testdata/transform/fusion/fuse-conv-mul-no-bias.onnx
vendored
Normal file
Binary file not shown.
Loading…
Reference in a new issue