mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-23 19:32:23 +00:00
Fix float 16 type support for some CUDA kernels (#436)
* Correct the Consts::Zero & Consts::One for half type * 1. Fix the CreateConstantOnes for float16 type 2. Add cuda kernel code in the BatchNorm for float 16 type, there's issue to run cudnnBatchNormalizationForwardInference with float 16 type 3. Add float 16 test case for Gemm & BatchNorm CUDA kernel only * Fix build * fix Linux build * fix build * Update the fix for BatchNorm, still use cuddn API cudnnBatchNormalizationForwardInference. The root case is, for half type, should use alpha, beta, scale, B, mean, var with float type. * fix build * enable 2 fp16 models for GPU test * enable fp16 test for MaxPool * Need to adjust per_sample_tolerance configuration in the model test
This commit is contained in:
parent
5866e853c4
commit
f14b258a5c
8 changed files with 275 additions and 15 deletions
|
|
@ -9,6 +9,7 @@
|
|||
#include <thrust/fill.h>
|
||||
#include "core/providers/cuda/shared_inc/cuda_utils.h"
|
||||
#include "core/providers/cuda/cu_inc/common.cuh"
|
||||
#include "cudnn_common.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace cuda {
|
||||
|
|
@ -31,7 +32,7 @@ class ConstantBufferImpl : public IConstantBuffer<T> {
|
|||
|
||||
template <typename T>
|
||||
std::unique_ptr<IConstantBuffer<T>> CreateConstantOnes() {
|
||||
return std::make_unique<ConstantBufferImpl<T>>((T)1);
|
||||
return std::make_unique<ConstantBufferImpl<T>>(Consts<T>::One);
|
||||
}
|
||||
|
||||
template std::unique_ptr<IConstantBuffer<float>> CreateConstantOnes<float>();
|
||||
|
|
|
|||
|
|
@ -5,6 +5,8 @@
|
|||
#include "core/providers/common.h"
|
||||
#include "core/providers/cuda/cudnn_common.h"
|
||||
#include "core/providers/cpu/nn/batch_norm_helper.h"
|
||||
#include "core/providers/cuda/math/unary_elementwise_ops_impl.h"
|
||||
|
||||
using namespace std;
|
||||
namespace onnxruntime {
|
||||
namespace cuda {
|
||||
|
|
@ -46,17 +48,54 @@ Status BatchNorm<T>::ComputeInternal(OpKernelContext* p_op_kernel_context) const
|
|||
auto mean_data = reinterpret_cast<const CudaT*>(mean->template Data<T>());
|
||||
auto var_data = reinterpret_cast<const CudaT*>(var->template Data<T>());
|
||||
|
||||
const auto alpha = Consts<CudaT>::One;
|
||||
const auto beta = Consts<CudaT>::Zero;
|
||||
|
||||
CudnnTensor data_desc;
|
||||
vector<int64_t> new_dims;
|
||||
BatchNormHelper::NormalizeDims(x_shape, new_dims);
|
||||
ORT_RETURN_IF_ERROR(data_desc.Set(new_dims, CudnnTensor::GetDataType<CudaT>()));
|
||||
|
||||
// For half data type, the alpha, beta, scale, B, mean, var need to be float type
|
||||
if (X->DataType() == DataTypeImpl::GetType<MLFloat16>()) {
|
||||
CudnnTensor scale_desc;
|
||||
ORT_RETURN_IF_ERROR(scale_desc.Set(new_dims, CudnnTensor::GetDataType<float>()));
|
||||
CudnnTensor bn_tensor_desc;
|
||||
ORT_RETURN_IF_ERROR(bn_tensor_desc.Set(data_desc, cudnn_batch_norm_mode_));
|
||||
|
||||
// Convert the scale, B, mean, var to float
|
||||
const int64_t C = x_shape.GetDims()[1];
|
||||
auto f_scale = GetScratchBuffer<float>(C);
|
||||
auto f_B = GetScratchBuffer<float>(C);
|
||||
auto f_mean = GetScratchBuffer<float>(C);
|
||||
auto f_var = GetScratchBuffer<float>(C);
|
||||
Impl_Cast<CudaT, float>(scale_data, f_scale.get(), C);
|
||||
Impl_Cast<CudaT, float>(b_data, f_B.get(), C);
|
||||
Impl_Cast<CudaT, float>(mean_data, f_mean.get(), C);
|
||||
Impl_Cast<CudaT, float>(var_data, f_var.get(), C);
|
||||
|
||||
CUDNN_RETURN_IF_ERROR(cudnnBatchNormalizationForwardInference(
|
||||
CudnnHandle(),
|
||||
cudnn_batch_norm_mode_,
|
||||
&alpha,
|
||||
&beta,
|
||||
data_desc,
|
||||
x_data,
|
||||
data_desc,
|
||||
y_data,
|
||||
bn_tensor_desc,
|
||||
f_scale.get(),
|
||||
f_B.get(),
|
||||
f_mean.get(),
|
||||
f_var.get(),
|
||||
epsilon_));
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
CudnnTensor bn_tensor_desc;
|
||||
ORT_RETURN_IF_ERROR(bn_tensor_desc.Set(data_desc, cudnn_batch_norm_mode_));
|
||||
|
||||
const auto alpha = Consts<CudaT>::One;
|
||||
const auto beta = Consts<CudaT>::Zero;
|
||||
|
||||
CUDNN_RETURN_IF_ERROR(cudnnBatchNormalizationForwardInference(
|
||||
CudnnHandle(),
|
||||
cudnn_batch_norm_mode_,
|
||||
|
|
|
|||
|
|
@ -341,9 +341,9 @@ int real_main(int argc, char* argv[]) {
|
|||
broken_tests["maxpool_3d_default"] = "cudnn pooling only support input dimension >= 3";
|
||||
broken_tests["maxpool_1d_default"] = "cudnn pooling only support input dimension >= 3";
|
||||
|
||||
broken_tests["fp16_tiny_yolov2"] = "unknown failure on CUDA";
|
||||
broken_tests["fp16_shufflenet"] = "unknown failure on CUDA";
|
||||
broken_tests["fp16_inception_v1"] = "unknown failure on CUDA";
|
||||
broken_tests["fp16_tiny_yolov2"] = "Need to adjust the per_sample_tolerance: 0.2";
|
||||
broken_tests["fp16_shufflenet"] = "still have issue on Linux";
|
||||
broken_tests["fp16_inception_v1"] = "need to adjust the per_sample_tolerance: 0.002";
|
||||
#endif
|
||||
|
||||
int result = 0;
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@
|
|||
namespace onnxruntime {
|
||||
namespace test {
|
||||
|
||||
TEST(MathOpTest, GemmNoTrans) {
|
||||
TEST(GemmOpTest, GemmNoTrans) {
|
||||
OpTester test("Gemm");
|
||||
|
||||
test.AddAttribute("transA", (int64_t)0);
|
||||
|
|
@ -26,7 +26,41 @@ TEST(MathOpTest, GemmNoTrans) {
|
|||
test.Run();
|
||||
}
|
||||
|
||||
TEST(MathOpTest, GemmBroadcast) {
|
||||
// Only CUDA kernel has float 16 support
|
||||
#ifdef USE_CUDA
|
||||
TEST(GemmOpTest, GemmNoTrans_f16) {
|
||||
OpTester test("Gemm");
|
||||
|
||||
test.AddAttribute("transA", (int64_t)0);
|
||||
test.AddAttribute("transB", (int64_t)0);
|
||||
test.AddAttribute("alpha", 1.0f);
|
||||
test.AddAttribute("beta", 1.0f);
|
||||
|
||||
std::vector<float> A{1.0f, 2.0f, 3.0f, 4.0f,
|
||||
-1.0f, -2.0f, -3.0f, -4.0f};
|
||||
std::vector<float> B(12, 1.0f);
|
||||
std::vector<float> C(6, 1.0f);
|
||||
std::vector<float> Y{11.0f, 11.0f, 11.0f,
|
||||
-9.0f, -9.0f, -9.0f};
|
||||
|
||||
std::vector<MLFloat16> f_A(8);
|
||||
std::vector<MLFloat16> f_B(12);
|
||||
std::vector<MLFloat16> f_C(6);
|
||||
std::vector<MLFloat16> f_Y(6);
|
||||
ConvertFloatToMLFloat16(A.data(), f_A.data(), 8);
|
||||
ConvertFloatToMLFloat16(B.data(), f_B.data(), 12);
|
||||
ConvertFloatToMLFloat16(C.data(), f_C.data(), 6);
|
||||
ConvertFloatToMLFloat16(Y.data(), f_Y.data(), 6);
|
||||
|
||||
test.AddInput<MLFloat16>("A", {2, 4}, f_A);
|
||||
test.AddInput<MLFloat16>("B", {4, 3}, f_B);
|
||||
test.AddInput<MLFloat16>("C", {2, 3}, f_C);
|
||||
test.AddOutput<MLFloat16>("Y", {2, 3}, f_Y);
|
||||
test.Run();
|
||||
}
|
||||
#endif
|
||||
|
||||
TEST(GemmOpTest, GemmBroadcast) {
|
||||
OpTester test("Gemm");
|
||||
|
||||
test.AddAttribute("transA", (int64_t)0);
|
||||
|
|
@ -45,7 +79,7 @@ TEST(MathOpTest, GemmBroadcast) {
|
|||
test.Run();
|
||||
}
|
||||
|
||||
TEST(MathOpTest, GemmTrans) {
|
||||
TEST(GemmOpTest, GemmTrans) {
|
||||
OpTester test("Gemm");
|
||||
|
||||
test.AddAttribute("transA", (int64_t)1);
|
||||
|
|
@ -66,7 +100,7 @@ TEST(MathOpTest, GemmTrans) {
|
|||
test.Run();
|
||||
}
|
||||
|
||||
TEST(MathOpTest, GemmAlphaBeta) {
|
||||
TEST(GemmOpTest, GemmAlphaBeta) {
|
||||
OpTester test("Gemm");
|
||||
|
||||
test.AddAttribute("transA", (int64_t)0);
|
||||
|
|
@ -85,7 +119,7 @@ TEST(MathOpTest, GemmAlphaBeta) {
|
|||
test.Run();
|
||||
}
|
||||
|
||||
TEST(MathOpTest, GemmNaN) {
|
||||
TEST(GemmOpTest, GemmNaN) {
|
||||
OpTester test("Gemm");
|
||||
|
||||
test.AddAttribute("transA", (int64_t)0);
|
||||
|
|
@ -104,7 +138,7 @@ TEST(MathOpTest, GemmNaN) {
|
|||
test.Run();
|
||||
}
|
||||
|
||||
TEST(MathOpTest, GemmScalarBroadcast) {
|
||||
TEST(GemmOpTest, GemmScalarBroadcast) {
|
||||
OpTester test("Gemm");
|
||||
|
||||
test.AddAttribute("transA", (int64_t)0);
|
||||
|
|
@ -142,7 +176,7 @@ TEST(MathOpTest, Gemm2DBroadcast) {
|
|||
test.Run();
|
||||
}
|
||||
|
||||
TEST(MathOpTest, GemmFalseBroadcast) {
|
||||
TEST(GemmOpTest, GemmFalseBroadcast) {
|
||||
OpTester test("Gemm");
|
||||
|
||||
test.AddAttribute("transA", (int64_t)0);
|
||||
|
|
@ -161,7 +195,7 @@ TEST(MathOpTest, GemmFalseBroadcast) {
|
|||
test.Run();
|
||||
}
|
||||
|
||||
TEST(MathOpTest, GemmEmptyTensor) {
|
||||
TEST(GemmOpTest, GemmEmptyTensor) {
|
||||
OpTester test("Gemm");
|
||||
|
||||
test.AddAttribute("transA", static_cast<int64_t>(0));
|
||||
|
|
|
|||
|
|
@ -512,5 +512,98 @@ TEST(BatchNormTest, InvalidVarDim) {
|
|||
OpTester::ExpectResult::kExpectFailure,
|
||||
"Invalid input var");
|
||||
}
|
||||
|
||||
// Only CUDA kernel has float 16 support
|
||||
#ifdef USE_CUDA
|
||||
TEST(BatchNormTest, BatchNorm2d_fp16) {
|
||||
vector<float> X{-0.91221f, -0.283559f, 0.937637f, 2.09818f, -0.100199f, -0.608113f, 0.444562f, -1.07505f, 0.940591f,
|
||||
-0.922262f, 0.0931303f, 0.69611f, 1.55187f, 0.159808f, 0.914874f, -1.24856f, -1.98928f, -0.331621f,
|
||||
2.33131f, 0.260409f, 0.0944811f, 0.442397f, 0.76461f, -0.203334f, -0.244228f, -0.387267f, -1.65039f,
|
||||
-0.815409f, 0.931696f, -1.15328f, 0.773952f, -1.28195f, -0.437349f, 0.0644882f, -0.087637f, 1.74999f,
|
||||
0.640154f, -0.505641f, -1.84014f, -0.00135415f, 0.782006f, -1.21172f, -0.621273f, -0.0977471f,
|
||||
-0.941333f, -0.170302f, 0.18923f, 0.436322f, 0.870412f, -0.582312f, 0.679017f, 0.510252f, 0.0786005f,
|
||||
0.160138f, -2.61889f, 0.402828f, 0.551144f, -1.39366f, -1.15191f, 0.160008f, -0.57735f, 0.210758f,
|
||||
1.0541f, -2.12569f, 0.101656f, 1.10223f, 0.725811f, -1.5019f, -0.0892582f, 0.063546f, 0.822734f,
|
||||
1.67707f, 0.478121f, -1.07438f, -0.0487855f, 0.0972885f, -1.54122f, 2.47422f, 0.596108f, 0.0026957f,
|
||||
-0.967677f, -2.08882f, 0.469692f, 0.630784f, 0.196915f, -1.91331f, 1.26255f, 0.0491993f, -0.358415f,
|
||||
0.720588f, 0.976776f, -0.418116f, 1.70979f, 2.49971f, 1.30942f, -1.18304f, -1.64901f, -1.11048f,
|
||||
1.41467f, -0.275486f, -1.20602f, -0.545566f, -0.918059f, 1.48513f, 2.04224f, -0.96909f, -1.92804f,
|
||||
0.634147f, -1.02079f, -0.000786079f, 0.72428f, 0.893569f, 1.14604f, -1.3423f, -1.05061f, -0.617524f,
|
||||
-0.12619f, -0.203127f, -0.941956f, 2.06916f, 2.03025f, 0.37269f, -0.340471f, -1.27962f, 0.159472f,
|
||||
0.643999f, 0.881773f, -0.50873f, 1.04599f, -0.287968f, 1.84344f, -0.728637f, 0.668021f, -2.00452f,
|
||||
-0.585523f, -0.24982f, -0.379091f, 0.213692f, 1.21336f, -0.499157f, -1.50841f, -1.01256f, 0.745338f,
|
||||
0.591107f, 1.33781f, -0.258927f, -1.87304f, 0.884799f, 1.63174f, 0.500887f, 1.80608f, -1.25441f,
|
||||
-0.655316f, 1.22439f, 0.384174f, 0.401395f, 1.43172f, 1.85338f, -0.644909f, -1.46975f, -1.06138f,
|
||||
1.09724f, 0.013438f, 0.589742f, -0.695768f, 0.758401f, 0.924533f, -0.0988563f, -0.197066f, 1.01118f,
|
||||
0.195163f, 0.975466f, 1.7682f, 0.977977f, -0.88963f, -0.251431f, 0.115828f, -0.230065f, -1.08882f,
|
||||
1.62318f, 0.502684f, 0.789724f, 1.13057f, -0.890021f, -0.614755f, 1.11055f, -1.21681f, 0.133085f,
|
||||
0.564458f, 0.723117f, 1.67088f, -0.111012f, 1.39732f, -0.846095f, -0.194408f, -0.381931f, -0.735943f,
|
||||
-0.788814f, -0.910318f, 1.16345f, -1.98542f, 0.742905f, -0.749476f, -0.110805f, 0.307949f, -1.66811f,
|
||||
0.294031f, -0.522837f, -0.774399f, -0.264072f, -0.426894f, 0.965971f, 0.173348f, -0.991018f, 1.9406f,
|
||||
0.0853744f};
|
||||
vector<float> scale{0.736494f, 0.580251f, 0.374834f};
|
||||
vector<float> B{0.0f, 0.0f, 0.0f};
|
||||
vector<float> mean{0.0f, 0.0f, 0.0f};
|
||||
vector<float> var{1.0f, 1.0f, 1.0f};
|
||||
|
||||
vector<float> expected_output{-0.671834f, -0.208838f, 0.69056f, 1.54529f, -0.0737958f, -0.447869f, 0.327415f, -0.791764f,
|
||||
0.692736f, -0.679237f, 0.0685895f, 0.512678f, 1.14294f, 0.117697f, 0.673796f, -0.91955f,
|
||||
-1.46508f, -0.244235f, 1.71699f, 0.191789f, 0.0695843f, 0.325821f, 0.563128f, -0.149753f,
|
||||
-0.179871f, -0.285218f, -1.2155f, -0.60054f, 0.686185f, -0.84938f, 0.570008f, -0.944146f,
|
||||
-0.322103f, 0.0474949f, -0.0645437f, 1.28885f, 0.371448f, -0.293397f, -1.06774f, -0.00078574f,
|
||||
0.453757f, -0.703098f, -0.360492f, -0.0567175f, -0.546206f, -0.0988174f, 0.1098f, 0.253175f,
|
||||
0.505055f, -0.337885f, 0.393998f, 0.296073f, 0.0456077f, 0.0929198f, -1.5196f, 0.23374f, 0.3198f,
|
||||
-0.808671f, -0.668391f, 0.0928441f, -0.335006f, 0.122292f, 0.611642f, -1.23343f, 0.0589855f,
|
||||
0.639564f, 0.42115f, -0.871472f, -0.0517919f, 0.0368724f, 0.477389f, 0.973116f, 0.179215f,
|
||||
-0.402711f, -0.0182864f, 0.0364668f, -0.577699f, 0.927416f, 0.22344f, 0.00101043f, -0.362716f,
|
||||
-0.782957f, 0.176056f, 0.236438f, 0.0738101f, -0.717172f, 0.473243f, 0.0184415f, -0.134346f,
|
||||
0.2701f, 0.366127f, -0.156723f, 0.640883f, 0.936974f, 0.490814f, -0.443443f, -0.618102f,
|
||||
-0.416243f, 0.530263f, -0.103261f, -0.452055f, -0.204496f, -0.344118f, 0.556675f, 0.765497f,
|
||||
-0.363246f, -0.722693f, 0.237699f, -0.751804f, -0.00057894f, 0.533425f, 0.658105f, 0.844047f,
|
||||
-0.988591f, -0.773767f, -0.4548f, -0.0929374f, -0.149601f, -0.693741f, 1.52392f, 1.49526f,
|
||||
0.274482f, -0.250754f, -0.942429f, 0.11745f, 0.474299f, 0.649417f, -0.374675f, 0.770361f,
|
||||
-0.212086f, 1.35768f, -0.536634f, 0.491991f, -1.47631f, -0.431231f, -0.18399f, -0.279197f,
|
||||
0.157382f, 0.89363f, -0.367624f, -1.11093f, -0.745738f, 0.548934f, 0.435344f, 0.776261f,
|
||||
-0.150242f, -1.08683f, 0.513403f, 0.946813f, 0.290638f, 1.04797f, -0.72787f, -0.380246f,
|
||||
0.710451f, 0.222916f, 0.232908f, 0.830753f, 1.07542f, -0.374207f, -0.852818f, -0.615864f,
|
||||
0.63667f, 0.00779735f, 0.342196f, -0.403718f, 0.44006f, 0.536458f, -0.0573611f, -0.114347f,
|
||||
0.586737f, 0.113243f, 0.566012f, 1.02599f, 0.567469f, -0.516206f, -0.145892f, 0.0672092f,
|
||||
-0.133495f, -0.631785f, 0.941846f, 0.188422f, 0.296014f, 0.423775f, -0.333609f, -0.23043f,
|
||||
0.416269f, -0.456101f, 0.0498845f, 0.211577f, 0.271048f, 0.626298f, -0.0416111f, 0.523762f,
|
||||
-0.317144f, -0.0728705f, -0.14316f, -0.275855f, -0.295673f, -0.341217f, 0.436097f, -0.7442f,
|
||||
0.278465f, -0.280928f, -0.0415335f, 0.115429f, -0.625263f, 0.110212f, -0.195976f, -0.29027f,
|
||||
-0.0989828f, -0.160014f, 0.362077f, 0.0649763f, -0.371465f, 0.727401f, 0.0320011f};
|
||||
float epsilon = 1e-05f;
|
||||
|
||||
OpTester test("BatchNormalization");
|
||||
test.AddAttribute("epsilon", epsilon);
|
||||
|
||||
vector<int64_t> input_shape{2, 3, 6, 6};
|
||||
int input_size = 2 * 3 * 6 * 6;
|
||||
|
||||
vector<MLFloat16> f_X(input_size);
|
||||
vector<MLFloat16> f_output(input_size);
|
||||
vector<MLFloat16> f_scale(3);
|
||||
vector<MLFloat16> f_B(3);
|
||||
vector<MLFloat16> f_mean(3);
|
||||
vector<MLFloat16> f_var(3);
|
||||
|
||||
ConvertFloatToMLFloat16(X.data(), f_X.data(), input_size);
|
||||
ConvertFloatToMLFloat16(scale.data(), f_scale.data(), 3);
|
||||
ConvertFloatToMLFloat16(B.data(), f_B.data(), 3);
|
||||
ConvertFloatToMLFloat16(mean.data(), f_mean.data(), 3);
|
||||
ConvertFloatToMLFloat16(var.data(), f_var.data(), 3);
|
||||
ConvertFloatToMLFloat16(expected_output.data(), f_output.data(), input_size);
|
||||
|
||||
test.AddInput<MLFloat16>("X", input_shape, f_X);
|
||||
test.AddInput<MLFloat16>("scale", {3}, f_scale);
|
||||
test.AddInput<MLFloat16>("B", {3}, f_B);
|
||||
test.AddInput<MLFloat16>("mean", {3}, f_mean);
|
||||
test.AddInput<MLFloat16>("var", {3}, f_var);
|
||||
test.AddOutput<MLFloat16>("output", input_shape, f_output);
|
||||
test.Run(OpTester::ExpectResult::kExpectSuccess, "");
|
||||
}
|
||||
#endif
|
||||
|
||||
} // namespace test
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
|
|
@ -52,6 +52,60 @@ TEST(PoolTest, MaxPool) {
|
|||
test.Run();
|
||||
}
|
||||
|
||||
// Only CUDA kernel has float 16 support
|
||||
// Disable for now, still investigating the issue with cudnn lib
|
||||
#ifdef USE_CUDA
|
||||
TEST(PoolTest, MaxPool_F16) {
|
||||
OpTester test("MaxPool");
|
||||
|
||||
test.AddAttribute("auto_pad", "");
|
||||
test.AddAttribute("strides", std::vector<int64_t>{1, 1});
|
||||
test.AddAttribute("pads", vector<int64_t>{0, 0, 0, 0});
|
||||
test.AddAttribute("kernel_shape", vector<int64_t>{8, 8});
|
||||
|
||||
std::vector<float> x_vals = {
|
||||
0.19151945412158966, 0.6221087574958801, 0.43772774934768677, 0.7853586077690125, 0.7799758315086365, 0.27259260416030884, 0.2764642536640167, 0.801872193813324,
|
||||
0.9581393599510193, 0.8759326338768005, 0.35781726241111755, 0.5009950995445251, 0.683462917804718, 0.7127020359039307, 0.37025076150894165, 0.5611962080001831,
|
||||
0.5030831694602966, 0.013768449425697327, 0.772826611995697, 0.8826411962509155, 0.36488598585128784, 0.6153962016105652, 0.07538124173879623, 0.3688240051269531,
|
||||
0.9331400990486145, 0.6513781547546387, 0.39720258116722107, 0.7887301445007324, 0.3168361186981201, 0.5680986642837524, 0.8691273927688599, 0.4361734092235565,
|
||||
0.802147626876831, 0.14376682043075562, 0.7042609453201294, 0.7045813202857971, 0.2187921106815338, 0.9248676300048828, 0.44214075803756714, 0.9093159437179565,
|
||||
0.05980922281742096, 0.18428708612918854, 0.047355279326438904, 0.6748809218406677, 0.5946247577667236, 0.5333101749420166, 0.043324064463377, 0.5614330768585205,
|
||||
0.32966843247413635, 0.5029668211936951, 0.11189431697130203, 0.6071937084197998, 0.5659446716308594, 0.006764062214642763, 0.617441713809967, 0.912122905254364,
|
||||
0.7905241250991821, 0.9920814633369446, 0.9588017463684082, 0.7919641137123108, 0.2852509617805481, 0.6249167323112488, 0.47809380292892456, 0.19567517936229706,
|
||||
|
||||
0.382317453622818, 0.053873684257268906, 0.45164841413497925, 0.9820047616958618, 0.12394270300865173, 0.1193808987736702, 0.7385230660438538, 0.587303638458252,
|
||||
0.47163254022598267, 0.10712681710720062, 0.22921857237815857, 0.8999651670455933, 0.41675353050231934, 0.5358516573905945, 0.0062085166573524475, 0.3006417155265808,
|
||||
0.43689316511154175, 0.6121490001678467, 0.9181980490684509, 0.625736653804779, 0.7059975862503052, 0.14983370900154114, 0.7460634112358093, 0.8310070037841797,
|
||||
0.6337257623672485, 0.4383098781108856, 0.15257278084754944, 0.5684096217155457, 0.5282242894172668, 0.9514287710189819, 0.48035916686058044, 0.5025595426559448,
|
||||
0.5368781685829163, 0.8192020654678345, 0.05711563676595688, 0.6694217324256897, 0.7671166062355042, 0.7081153392791748, 0.7968671917915344, 0.5577608346939087,
|
||||
0.9658365249633789, 0.14715689420700073, 0.02964700013399124, 0.5938934683799744, 0.11406569927930832, 0.9508098363876343, 0.32570740580558777, 0.19361868500709534,
|
||||
0.4578116536140442, 0.9204025864601135, 0.8790691494941711, 0.252615749835968, 0.34800878167152405, 0.18258872628211975, 0.9017960429191589, 0.7065281867980957,
|
||||
0.7266584634780884, 0.900087833404541, 0.7791637778282166, 0.5991547703742981, 0.29112523794174194, 0.1513952612876892, 0.33517464995384216, 0.6575517654418945,
|
||||
|
||||
0.07334254682064056, 0.055006396025419235, 0.32319480180740356, 0.5904818177223206, 0.8538985848426819, 0.2870624363422394, 0.17306722700595856, 0.13402120769023895,
|
||||
0.9946538209915161, 0.1794978678226471, 0.3175468146800995, 0.568291425704956, 0.009348574094474316, 0.9006485939025879, 0.9772414565086365, 0.5568946599960327,
|
||||
0.08477384597063065, 0.3330024778842926, 0.7284286618232727, 0.14243537187576294, 0.5524689555168152, 0.2730432450771332, 0.9744951128959656, 0.6677868962287903,
|
||||
0.2556532919406891, 0.1083114966750145, 0.7761807441711426, 0.7824779748916626, 0.7616038918495178, 0.9144031405448914, 0.6586228013038635, 0.568367600440979,
|
||||
0.20175568759441376, 0.6982963681221008, 0.952195405960083, 0.8899632692337036, 0.9935673475265503, 0.8187035322189331, 0.5451221466064453, 0.45125406980514526,
|
||||
0.8905571699142456, 0.9732648134231567, 0.5934113264083862, 0.36607450246810913, 0.3230946958065033, 0.8714232444763184, 0.2156340628862381, 0.7349451780319214,
|
||||
0.36561909317970276, 0.8016026020050049, 0.7827355861663818, 0.7013553977012634, 0.6227765679359436, 0.4936826527118683, 0.8405377268791199, 0.7120969891548157,
|
||||
0.4439089894294739, 0.031034860759973526, 0.36323976516723633, 0.7307217717170715, 0.475566565990448, 0.3444169759750366, 0.6408804059028625, 0.12620532512664795};
|
||||
std::vector<int64_t> x_dims = {1, 3, 8, 8};
|
||||
int x_size = 1 * 3 * 8 * 8;
|
||||
std::vector<int64_t> expected_dims = {1, 3, 1, 1};
|
||||
std::vector<float> expected_vals = {0.9920814633369446, 0.9820047616958618, 0.9946538209915161};
|
||||
|
||||
std::vector<MLFloat16> f_X(x_size);
|
||||
std::vector<MLFloat16> f_Y(3);
|
||||
ConvertFloatToMLFloat16(x_vals.data(), f_X.data(), x_size);
|
||||
ConvertFloatToMLFloat16(expected_vals.data(), f_Y.data(), 3);
|
||||
|
||||
test.AddInput<MLFloat16>("X", x_dims, f_X);
|
||||
test.AddOutput<MLFloat16>("Y", expected_dims, f_Y);
|
||||
test.Run();
|
||||
}
|
||||
#endif
|
||||
|
||||
static void MaxPool_8_WithIndexTest(bool has_index, int64_t storage_order=0) {
|
||||
OpTester test("MaxPool", 8);
|
||||
|
||||
|
|
|
|||
|
|
@ -68,6 +68,29 @@ void Check<float>(const OpTester::Data& expected_data, const Tensor& output_tens
|
|||
}
|
||||
}
|
||||
|
||||
template <>
|
||||
void Check<MLFloat16>(const OpTester::Data& expected_data, const Tensor& output_tensor, const std::string& provider_type) {
|
||||
auto& expected_tensor = expected_data.data_.Get<Tensor>();
|
||||
auto* expected = expected_tensor.template Data<MLFloat16>();
|
||||
auto* output = output_tensor.template Data<MLFloat16>();
|
||||
auto size = output_tensor.Shape().Size();
|
||||
|
||||
std::vector<float> f_expected(size);
|
||||
std::vector<float> f_output(size);
|
||||
ConvertMLFloat16ToFloat(expected, f_expected.data(), static_cast<int>(size));
|
||||
ConvertMLFloat16ToFloat(output, f_output.data(), static_cast<int>(size));
|
||||
|
||||
float threshold = 0.001f;
|
||||
for (int i = 0; i < size; ++i) {
|
||||
if (std::isinf(f_expected[i])) // Test infinity for equality
|
||||
EXPECT_EQ(f_expected[i], f_output[i]);
|
||||
else {
|
||||
// the default for existing tests
|
||||
EXPECT_NEAR(f_expected[i], f_output[i], threshold) << "provider_type: " << provider_type;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename Type>
|
||||
void CheckDispatch(MLDataType type, const OpTester::Data& expected_data, const Tensor& output_tensor, const std::string& provider_type) {
|
||||
if (type == DataTypeImpl::GetType<Type>())
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@
|
|||
#include "gmock/gmock.h"
|
||||
#include "gtest/gtest.h"
|
||||
#include <gsl/gsl_byte>
|
||||
#include "core/util/math_cpuonly.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
class InferenceSession;
|
||||
|
|
@ -347,5 +348,20 @@ void DebugTrap();
|
|||
|
||||
void Check(const OpTester::Data& expected_data, const Tensor& output_tensor, const std::string& provider_type);
|
||||
|
||||
// Only used for CUDA test since no toher kernel has float 16 support
|
||||
#ifdef USE_CUDA
|
||||
inline void ConvertFloatToMLFloat16(const float* f_datat, MLFloat16* h_data, int input_size) {
|
||||
auto in_vector = ConstEigenVectorMap<float>(f_datat, input_size);
|
||||
auto output_vector = EigenVectorMap<Eigen::half>(static_cast<Eigen::half*>(static_cast<void*>(h_data)), input_size);
|
||||
output_vector = in_vector.template cast<Eigen::half>();
|
||||
}
|
||||
#endif
|
||||
|
||||
inline void ConvertMLFloat16ToFloat(const MLFloat16* h_data, float* f_data, int input_size) {
|
||||
auto in_vector = ConstEigenVectorMap<Eigen::half>(static_cast<const Eigen::half*>(static_cast<const void*>(h_data)), input_size);
|
||||
auto output_vector = EigenVectorMap<float>(f_data, input_size);
|
||||
output_vector = in_vector.template cast<float>();
|
||||
}
|
||||
|
||||
} // namespace test
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
Loading…
Reference in a new issue