mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
Add uint8 support for BitShift operator (#2214)
* Add uint8 support for BitShift operator * Remove more tests from exclusion * Updates
This commit is contained in:
parent
91122a2cf5
commit
5eb42f4452
5 changed files with 42 additions and 26 deletions
|
|
@ -404,6 +404,7 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Sp
|
|||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, ScatterND);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Gemm);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, GatherElements);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, uint8_t, BitShift);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, uint32_t, BitShift);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, uint64_t, BitShift);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Pad);
|
||||
|
|
@ -828,25 +829,25 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Shrink)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, float, Erf)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
int64_t_int64_t_int64_t, OneHot)>,
|
||||
int64_t_int64_t_int64_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
float_int64_t_int64_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
int64_t_string_int64_t, OneHot)>,
|
||||
float_int64_t_int64_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
float_string_int64_t, OneHot)>,
|
||||
int64_t_string_int64_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
float_float_float, OneHot)>,
|
||||
float_string_int64_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
int64_t_int32_t_float, OneHot)>,
|
||||
float_float_float, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
int64_t_float_int64_t, OneHot)>,
|
||||
int64_t_int32_t_float, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
int32_t_float_int32_t, OneHot)>,
|
||||
int64_t_float_int64_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
int32_t_float_float, OneHot)>,
|
||||
int32_t_float_int32_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
int64_t_float_float, OneHot)>,
|
||||
int32_t_float_float, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
int64_t_float_float, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10,
|
||||
MaxUnpool)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Sinh)>,
|
||||
|
|
@ -911,11 +912,11 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
|
|||
AveragePool)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, Mod)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, float,
|
||||
Resize)>,
|
||||
Resize)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, int32_t,
|
||||
Resize)>,
|
||||
Resize)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, uint8_t,
|
||||
Resize)>,
|
||||
Resize)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, ThresholdedRelu)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, uint8_t,
|
||||
DequantizeLinear)>,
|
||||
|
|
@ -1055,6 +1056,7 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, ScatterND)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Gemm)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, GatherElements)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, uint8_t, BitShift)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, uint32_t, BitShift)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, uint64_t, BitShift)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Pad)>,
|
||||
|
|
@ -1080,13 +1082,13 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
|
|||
int32_t_float_int32_t, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11,
|
||||
int32_t_float_float, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11,
|
||||
int64_t_float_float, OneHot)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11,
|
||||
float, Resize)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11,
|
||||
int32_t, Resize)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11,
|
||||
uint8_t, Resize)>,
|
||||
};
|
||||
|
||||
|
|
@ -1259,7 +1261,7 @@ static Status RegisterCPUKernels(KernelRegistry& kernel_registry) {
|
|||
return Status::OK();
|
||||
}
|
||||
|
||||
struct KernelRegistryAndStatus{
|
||||
struct KernelRegistryAndStatus {
|
||||
std::shared_ptr<KernelRegistry> kernel_registry = std::make_shared<KernelRegistry>();
|
||||
Status st;
|
||||
};
|
||||
|
|
|
|||
|
|
@ -129,7 +129,7 @@ REG_ELEMENTWISE_LOGICALOP_TYPED_KERNEL(Equal, 11, float, Equal);
|
|||
REG_ELEMENTWISE_VERSIONED_TYPED_KERNEL(Mean, 6, 7, float, Mean_6);
|
||||
REG_ELEMENTWISE_TYPED_KERNEL(Mean, 8, float, Mean_8);
|
||||
|
||||
//REG_ELEMENTWISE_TYPED_KERNEL(BitShift, 11, uint8_t, BitShift);
|
||||
REG_ELEMENTWISE_TYPED_KERNEL(BitShift, 11, uint8_t, BitShift);
|
||||
//REG_ELEMENTWISE_TYPED_KERNEL(BitShift, 11, uint16_t, BitShift);
|
||||
REG_ELEMENTWISE_TYPED_KERNEL(BitShift, 11, uint32_t, BitShift);
|
||||
REG_ELEMENTWISE_TYPED_KERNEL(BitShift, 11, uint64_t, BitShift);
|
||||
|
|
|
|||
|
|
@ -433,9 +433,7 @@ int real_main(int argc, char* argv[], Ort::Env& env) {
|
|||
{"resize_upsample_sizes_nearest_ceil_half_pixel", "Bad onnx test output. Needs test fix."},
|
||||
{"resize_upsample_sizes_nearest_floor_align_corners", "Bad onnx test output. Needs test fix."},
|
||||
{"resize_upsample_sizes_nearest_round_prefer_ceil_asymmetric", "Bad onnx test output. Needs test fix."},
|
||||
{"bitshift_right_uint8", "BitShift(11) uint8 support not enabled currently"},
|
||||
{"bitshift_right_uint16", "BitShift(11) uint16 support not enabled currently"},
|
||||
{"bitshift_left_uint8", "BitShift(11) uint8 support not enabled currently"},
|
||||
{"bitshift_left_uint16", "BitShift(11) uint16 support not enabled currently"},
|
||||
{"reflect_pad", "Pad(11) int32 support not enabled currently"},
|
||||
{"edge_pad", "Pad(11) int32 support not enabled currently"},
|
||||
|
|
|
|||
|
|
@ -401,7 +401,7 @@ TEST(MathOpTest, Abs_int8) {
|
|||
std::vector<int64_t> dims{4};
|
||||
test.AddInput<int8_t>("X", dims, {1, 2, -1, -5});
|
||||
test.AddOutput<int8_t>("Y", dims, {1, 2, 1, 5});
|
||||
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider}); //TensorRT: INT8, Assertion `regionRanges != nullptr' failed
|
||||
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider}); //TensorRT: INT8, Assertion `regionRanges != nullptr' failed
|
||||
}
|
||||
|
||||
TEST(MathOpTest, Abs_int32) {
|
||||
|
|
@ -429,7 +429,7 @@ TEST(MathOpTest, Neg_int8) {
|
|||
std::vector<int64_t> dims{4};
|
||||
test.AddInput<int8_t>("X", dims, {1, -2, 0, -10});
|
||||
test.AddOutput<int8_t>("Y", dims, {-1, 2, 0, 10});
|
||||
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider}); //TensorRT: INT8 is not supported
|
||||
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider}); //TensorRT: INT8 is not supported
|
||||
}
|
||||
|
||||
TEST(MathOpTest, Neg_int32) {
|
||||
|
|
@ -1628,5 +1628,23 @@ TEST(BitShiftOpTest, BroadcastXRight) {
|
|||
test.Run();
|
||||
}
|
||||
|
||||
TEST(BitShiftOpTest, BroadcastYLeft_Uint8) {
|
||||
OpTester test("BitShift", 11);
|
||||
test.AddAttribute("direction", "LEFT");
|
||||
test.AddInput<uint8_t>("X", {3, 2}, {1, 2, 3, 4, 5, 6});
|
||||
test.AddInput<uint8_t>("Y", {2}, {1, 2});
|
||||
test.AddOutput<uint8_t>("Z", {3, 2}, {2, 8, 6, 16, 10, 24});
|
||||
test.Run();
|
||||
}
|
||||
|
||||
TEST(BitShiftOpTest, BroadcastXRight_Uint8) {
|
||||
OpTester test("BitShift", 11);
|
||||
test.AddAttribute("direction", "RIGHT");
|
||||
test.AddInput<uint8_t>("X", {2}, {64, 32});
|
||||
test.AddInput<uint8_t>("Y", {3, 2}, {1, 2, 3, 4, 5, 6});
|
||||
test.AddOutput<uint8_t>("Z", {3, 2}, {32, 8, 8, 2, 2, 0});
|
||||
test.Run();
|
||||
}
|
||||
|
||||
} // namespace test
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
|
|
@ -93,9 +93,7 @@ def other_tests_failing_permanently_filters():
|
|||
|
||||
def test_with_types_disabled_due_to_binary_size_concerns_filters():
|
||||
filters = ['^test_bitshift_right_uint16_cpu',
|
||||
'^test_bitshift_right_uint8_cpu',
|
||||
'^test_bitshift_left_uint16_cpu',
|
||||
'^test_bitshift_left_uint8_cpu',
|
||||
'^test_edge_pad_cpu',
|
||||
'^test_reflect_pad_cpu']
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue