mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
Add float16 support for fusion (#476)
* Add float16 support for fusion * update test case * update test case
This commit is contained in:
parent
9add0e9a9f
commit
2a9a924c23
3 changed files with 146 additions and 0 deletions
|
|
@ -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<int64_t>(1), std::multiplies<int64_t>{});
|
||||
|
||||
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<uint16_t>(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<uint16_t>();
|
||||
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<float>();
|
||||
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<uint16_t>();
|
||||
const uint16_t* src = other.data<uint16_t>();
|
||||
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<float>();
|
||||
const float* src = other.data<float>();
|
||||
|
|
@ -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<uint16_t>();
|
||||
const uint16_t* src = other.data<uint16_t>();
|
||||
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<float>();
|
||||
const float* src = other.data<float>();
|
||||
|
|
@ -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<uint16_t>();
|
||||
const uint16_t* src = other.data<uint16_t>();
|
||||
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<float>();
|
||||
const float* src = other.data<float>();
|
||||
|
|
@ -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<uint16_t>();
|
||||
const uint16_t* src = other.data<uint16_t>();
|
||||
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<float>();
|
||||
const float* src = other.data<float>();
|
||||
|
|
@ -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<uint16_t>();
|
||||
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<float>();
|
||||
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<uint16_t>();
|
||||
const uint16_t* src = other.data<uint16_t>();
|
||||
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<float>();
|
||||
const float* src = other.data<float>();
|
||||
|
|
@ -376,6 +463,7 @@ class Initializer final {
|
|||
|
||||
std::string raw_data_;
|
||||
std::vector<float> float_data_;
|
||||
std::vector<uint16_t> float16_data_;
|
||||
std::vector<double> double_data_;
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -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<Model> p_model;
|
||||
ASSERT_TRUE(Model::Load(model_uri, p_model).IsOK());
|
||||
|
||||
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(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<int64_t> dims_x = {1,1,3,3};
|
||||
std::vector<MLFloat16> values_x;
|
||||
for (int i = 0; i < 9; ++i) {
|
||||
values_x.push_back(x_f);
|
||||
}
|
||||
CreateMLValue<MLFloat16>(TestCPUExecutionProvider()->GetAllocator(0, OrtMemTypeDefault), dims_x, values_x, &ml_value_x);
|
||||
feeds.insert(std::make_pair("X", ml_value_x));
|
||||
|
||||
std::vector<std::string> output_names;
|
||||
output_names.push_back("PROD");
|
||||
std::vector<MLValue> fetches;
|
||||
|
||||
ASSERT_TRUE(session_object.Run(run_options, feeds, output_names, &fetches).IsOK());
|
||||
|
||||
auto prod_f = MLFloat16(math::floatToHalf(6.0));
|
||||
std::vector<int64_t> expected_dims_prod = {1,1,2,2};
|
||||
std::vector<MLFloat16> 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<Tensor>();
|
||||
TensorShape expected_shape(expected_dims_prod);
|
||||
ASSERT_EQ(expected_shape, rtensor.Shape());
|
||||
const std::vector<MLFloat16> found(rtensor.template Data<MLFloat16>(), rtensor.template Data<MLFloat16>() + expected_dims_prod.size());
|
||||
ASSERT_EQ(expected_values_prod, found);
|
||||
}
|
||||
|
||||
} // namespace test
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
BIN
onnxruntime/test/testdata/transform/fusion/fuse-conv-bn-add-mul-float16.onnx
vendored
Normal file
BIN
onnxruntime/test/testdata/transform/fusion/fuse-conv-bn-add-mul-float16.onnx
vendored
Normal file
Binary file not shown.
Loading…
Reference in a new issue