diff --git a/onnxruntime/core/optimizer/initializer.h b/onnxruntime/core/optimizer/initializer.h index d1283e8967..477fde90e4 100644 --- a/onnxruntime/core/optimizer/initializer.h +++ b/onnxruntime/core/optimizer/initializer.h @@ -9,6 +9,7 @@ #include "core/common/common.h" #include "core/graph/onnx_protobuf.h" +#include "core/util/math.h" namespace onnxruntime { @@ -17,6 +18,7 @@ class Initializer final { static bool IsSupportedDataType(const ONNX_NAMESPACE::TensorProto* tensor_proto) { return !(tensor_proto == nullptr || (tensor_proto->data_type() != ONNX_NAMESPACE::TensorProto_DataType_FLOAT && + tensor_proto->data_type() != ONNX_NAMESPACE::TensorProto_DataType_FLOAT16 && tensor_proto->data_type() != ONNX_NAMESPACE::TensorProto_DataType_DOUBLE)); } @@ -30,6 +32,10 @@ class Initializer final { size_ = std::accumulate(dims_.begin(), dims_.end(), static_cast(1), std::multiplies{}); switch (data_type_) { + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: { + float16_data_.assign(size_, 0); + break; + } case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: { float_data_.assign(size_, 0.0f); break; @@ -60,6 +66,14 @@ class Initializer final { raw_data_ = tensor_proto->raw_data(); } else { switch (data_type_) { + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: { + int64_t size = tensor_proto->int32_data_size(); + ORT_ENFORCE(size_ == size, "size is different"); + for (int i = 0; i < size_; i++) { + float16_data_.push_back(static_cast(tensor_proto->int32_data(i))); + } + break; + } case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: { int64_t size = tensor_proto->float_data_size(); ORT_ENFORCE(size_ == size, "size is different"); @@ -104,6 +118,13 @@ class Initializer final { tensor_proto->set_raw_data(raw_data_); } else { switch (data_type_) { + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: { + tensor_proto->clear_int32_data(); + for (int i = 0; i < size_; i++) { + tensor_proto->add_int32_data(float16_data_[i]); + } + break; + } case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: { tensor_proto->clear_float_data(); for (int i = 0; i < size_; i++) { @@ -143,6 +164,10 @@ class Initializer final { return (T*)&raw_data_[0]; } switch (data_type_) { + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: { + return (T*)float16_data_.data(); + break; + } case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: { return (T*)float_data_.data(); break; @@ -164,6 +189,10 @@ class Initializer final { return (T*)&raw_data_[0]; } switch (data_type_) { + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: { + return (T*)float16_data_.data(); + break; + } case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: { return (T*)float_data_.data(); break; @@ -192,6 +221,13 @@ class Initializer final { Initializer& add(float value) { int64_t n = size(); switch (data_type_) { + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: { + uint16_t* dst = data(); + for (int i = 0; i < n; i++) { + dst[i] = math::floatToHalf(math::halfToFloat(dst[i]) + value); + } + break; + } case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: { float* dst = data(); for (int i = 0; i < n; i++) { @@ -215,6 +251,14 @@ class Initializer final { Initializer& add(const Initializer& other) { int64_t n = size(); switch (data_type_) { + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: { + uint16_t* dst = data(); + const uint16_t* src = other.data(); + for (int i = 0; i < n; i++) { + dst[i] = math::floatToHalf(math::halfToFloat(dst[i]) + math::halfToFloat(src[i])); + } + break; + } case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: { float* dst = data(); const float* src = other.data(); @@ -239,6 +283,14 @@ class Initializer final { Initializer& sub(const Initializer& other) { int64_t n = size(); switch (data_type_) { + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: { + uint16_t* dst = data(); + const uint16_t* src = other.data(); + for (int i = 0; i < n; i++) { + dst[i] = math::floatToHalf(math::halfToFloat(dst[i]) - math::halfToFloat(src[i])); + } + break; + } case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: { float* dst = data(); const float* src = other.data(); @@ -264,6 +316,14 @@ class Initializer final { Initializer& mul(const Initializer& other) { int64_t n = size(); switch (data_type_) { + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: { + uint16_t* dst = data(); + const uint16_t* src = other.data(); + for (int i = 0; i < n; i++) { + dst[i] = math::floatToHalf(math::halfToFloat(dst[i]) * math::halfToFloat(src[i])); + } + break; + } case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: { float* dst = data(); const float* src = other.data(); @@ -288,6 +348,14 @@ class Initializer final { Initializer& div(const Initializer& other) { int64_t n = size(); switch (data_type_) { + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: { + uint16_t* dst = data(); + const uint16_t* src = other.data(); + for (int i = 0; i < n; i++) { + dst[i] = math::floatToHalf(math::halfToFloat(dst[i]) / math::halfToFloat(src[i])); + } + break; + } case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: { float* dst = data(); const float* src = other.data(); @@ -313,6 +381,13 @@ class Initializer final { Initializer& sqrt() { int64_t n = size(); switch (data_type_) { + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: { + uint16_t* dst = data(); + for (int i = 0; i < n; i++) { + dst[i] = math::floatToHalf(std::sqrt(math::halfToFloat(dst[i]))); + } + break; + } case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: { float* dst = data(); for (int i = 0; i < n; i++) { @@ -341,6 +416,18 @@ class Initializer final { int64_t n = size() / num; switch (data_type_) { + case ONNX_NAMESPACE::TensorProto_DataType_FLOAT16: { + uint16_t* dst = data(); + const uint16_t* 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++) { + auto k = i * num + j; + dst[k] = math::floatToHalf(math::halfToFloat(dst[k]) * math::halfToFloat(src[index])); + } + } + break; + } case ONNX_NAMESPACE::TensorProto_DataType_FLOAT: { float* dst = data(); const float* src = other.data(); @@ -376,6 +463,7 @@ class Initializer final { std::string raw_data_; std::vector float_data_; + std::vector float16_data_; std::vector double_data_; }; diff --git a/onnxruntime/test/ir/graph_transform_test.cc b/onnxruntime/test/ir/graph_transform_test.cc index d5671b621b..7ea5863641 100644 --- a/onnxruntime/test/ir/graph_transform_test.cc +++ b/onnxruntime/test/ir/graph_transform_test.cc @@ -15,7 +15,11 @@ #include "core/optimizer/conv_activation_fusion.h" #include "core/optimizer/matmul_add_fusion.h" #include "core/optimizer/gemm_activation_fusion.h" +#include "core/framework/data_types.h" +#include "core/framework/ml_value.h" +#include "core/util/math.h" #include "core/platform/env.h" +#include "test/framework/test_utils.h" #include "test/capturing_sink.h" #include "test/test_environment.h" #include "gtest/gtest.h" @@ -281,6 +285,60 @@ TEST(GraphTransformationTests, Gemm_Relu_three_input) { ASSERT_TRUE(session_object.Initialize().IsOK()); } +TEST(GraphTransformationTests, FuseConvBnAddMulFloat16) { + string model_uri = MODEL_FOLDER + "fusion/fuse-conv-bn-add-mul-float16.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(); + std::unique_ptr ConvMulFusion_transformer = std::make_unique(); + std::unique_ptr ConvAddFusion_transformer = std::make_unique(); + 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()); + + NameMLValMap feeds; + RunOptions run_options; + run_options.run_tag = "one session/one tag"; + MLValue ml_value_x; + + auto x_f = MLFloat16(math::floatToHalf(1.0)); + std::vector dims_x = {1,1,3,3}; + std::vector values_x; + for (int i = 0; i < 9; ++i) { + values_x.push_back(x_f); + } + CreateMLValue(TestCPUExecutionProvider()->GetAllocator(0, OrtMemTypeDefault), dims_x, values_x, &ml_value_x); + feeds.insert(std::make_pair("X", ml_value_x)); + + std::vector output_names; + output_names.push_back("PROD"); + std::vector fetches; + + ASSERT_TRUE(session_object.Run(run_options, feeds, output_names, &fetches).IsOK()); + + auto prod_f = MLFloat16(math::floatToHalf(6.0)); + std::vector expected_dims_prod = {1,1,2,2}; + std::vector expected_values_prod; + for (int i = 0; i < 4; ++i) { + expected_values_prod.push_back(prod_f); + } + + ASSERT_EQ(1, fetches.size()); + auto& rtensor = fetches.front().Get(); + TensorShape expected_shape(expected_dims_prod); + ASSERT_EQ(expected_shape, rtensor.Shape()); + const std::vector found(rtensor.template Data(), rtensor.template Data() + expected_dims_prod.size()); + ASSERT_EQ(expected_values_prod, found); +} } // namespace test } // namespace onnxruntime diff --git a/onnxruntime/test/testdata/transform/fusion/fuse-conv-bn-add-mul-float16.onnx b/onnxruntime/test/testdata/transform/fusion/fuse-conv-bn-add-mul-float16.onnx new file mode 100644 index 0000000000..f76a8bcb83 Binary files /dev/null and b/onnxruntime/test/testdata/transform/fusion/fuse-conv-bn-add-mul-float16.onnx differ