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:
Weixing Zhang 2018-12-03 13:11:04 -08:00 committed by GitHub
parent b7cc611563
commit aa549cd194
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
9 changed files with 208 additions and 41 deletions

View file

@ -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);

View file

@ -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

View file

@ -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;

View file

@ -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