Support int8 for operator Split (#8615)

* Support int8 for operator Split
This commit is contained in:
Xavier Dupré 2021-08-10 23:04:16 +02:00 committed by GitHub
parent 3a742f2910
commit 064a385b59
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
3 changed files with 24 additions and 7 deletions

View file

@ -303,9 +303,9 @@ Do not modify directly.*
|Softsign|*in* input:**T**<br> *out* output:**T**|1+|**T** = tensor(float)|
|SpaceToDepth|*in* input:**T**<br> *out* output:**T**|13+|**T** = tensor(double), tensor(float)|
|||[1, 12]|**T** = tensor(double), tensor(float)|
|Split|*in* input:**T**<br> *in* split:**T**<br> *out* outputs...:**T**<br><br>or<br><br>*in* input:**T**<br> *in* split:**tensor(int64)**<br> *out* outputs:**T**<br><br>or<br><br>*in* input:**T**<br> *out* outputs:**T**|13+|**T** = tensor(float), tensor(int32), tensor(int64), tensor(string), tensor(uint8)|
|||[11, 12]|**T** = tensor(float), tensor(int32), tensor(int64), tensor(string), tensor(uint8)|
|||[2, 10]|**T** = tensor(float), tensor(int32), tensor(int64), tensor(string), tensor(uint8)|
|Split|*in* input:**T**<br> *in* split:**T**<br> *out* outputs...:**T**<br><br>or<br><br>*in* input:**T**<br> *in* split:**tensor(int64)**<br> *out* outputs:**T**<br><br>or<br><br>*in* input:**T**<br> *out* outputs:**T**|13+|**T** = tensor(float), tensor(int32), tensor(int64), tensor(int8), tensor(string), tensor(uint8)|
|||[11, 12]|**T** = tensor(float), tensor(int32), tensor(int64), tensor(int8), tensor(string), tensor(uint8)|
|||[2, 10]|**T** = tensor(float), tensor(int32), tensor(int64), tensor(int8), tensor(string), tensor(uint8)|
|SplitToSequence|*in* input:**T**<br> *in* split:**I**<br> *out* output_sequence:**S**|11+|**I** = tensor(int32), tensor(int64)<br/> **S** = seq(tensor(bfloat16)), seq(tensor(bool)), seq(tensor(double)), seq(tensor(float)), seq(tensor(float16)), seq(tensor(int16)), seq(tensor(int32)), seq(tensor(int64)), seq(tensor(int8)), seq(tensor(string)), seq(tensor(uint16)), seq(tensor(uint32)), seq(tensor(uint64)), seq(tensor(uint8))<br/> **T** = tensor(double), tensor(float), tensor(int32), tensor(int64), tensor(string)|
|Sqrt|*in* X:**T**<br> *out* Y:**T**|13+|**T** = tensor(double), tensor(float)|
|||[6, 12]|**T** = tensor(double), tensor(float)|

View file

@ -16,7 +16,7 @@ namespace onnxruntime {
namespace op_kernel_type_control {
ORT_SPECIFY_OP_KERNEL_ARG_DEFAULT_TYPES_ALL_OPSETS(
kCpuExecutionProvider, kOnnxDomain, Split, Input, 0,
float, int32_t, int64_t, uint8_t, std::string);
float, int8_t, int32_t, int64_t, uint8_t, std::string);
ORT_SPECIFY_OP_KERNEL_ARG_REQUIRED_TYPES_ALL_OPSETS(
kCpuExecutionProvider, kOnnxDomain, Split, Input, 0,
int32_t, int64_t);
@ -27,13 +27,18 @@ using SplitDataTypes = ORT_OP_KERNEL_ARG_DEFAULT_TYPE_LIST_ALL_OPSETS(
using EnabledSplitDataTypes = ORT_OP_KERNEL_ARG_ENABLED_TYPE_LIST_ALL_OPSETS(
kCpuExecutionProvider, kOnnxDomain, Split, Input, 0);
using OldSplitDataTypes = onnxruntime::TypeList<float, int32_t, int64_t, uint8_t, std::string>;
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
Split,
2,
10,
KernelDefBuilder().TypeConstraint("T",
BuildKernelDefConstraintsFromTypeList<SplitDataTypes>(),
BuildKernelDefConstraintsFromTypeList<EnabledSplitDataTypes>()),
BuildKernelDefConstraintsFromTypeList<EnabledSplitDataTypes>())
.FixedTypeConstraintForHash(
"T",
BuildKernelDefConstraintsFromTypeList<OldSplitDataTypes>()),
Split);
// Opset 11 starts to support Neg Axis.
@ -43,7 +48,10 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
12,
KernelDefBuilder().TypeConstraint("T",
BuildKernelDefConstraintsFromTypeList<SplitDataTypes>(),
BuildKernelDefConstraintsFromTypeList<EnabledSplitDataTypes>()),
BuildKernelDefConstraintsFromTypeList<EnabledSplitDataTypes>())
.FixedTypeConstraintForHash(
"T",
BuildKernelDefConstraintsFromTypeList<OldSplitDataTypes>()),
Split);
// Opset 13 starts to supports 'split' as optional input.
@ -52,7 +60,10 @@ ONNX_CPU_OPERATOR_KERNEL(
13,
KernelDefBuilder().TypeConstraint("T",
BuildKernelDefConstraintsFromTypeList<SplitDataTypes>(),
BuildKernelDefConstraintsFromTypeList<EnabledSplitDataTypes>()),
BuildKernelDefConstraintsFromTypeList<EnabledSplitDataTypes>())
.FixedTypeConstraintForHash(
"T",
BuildKernelDefConstraintsFromTypeList<OldSplitDataTypes>()),
Split);
Status SplitBase::PrepareForCompute(const TensorShape& input_shape, int num_outputs, int64_t& axis, int& before_dims,
@ -108,6 +119,8 @@ Status Split::Compute(OpKernelContext* context) const {
status = ComputeImpl<int64_t>(*context, input);
else if (input.IsDataType<uint8_t>())
status = ComputeImpl<uint8_t>(*context, input);
else if (input.IsDataType<int8_t>())
status = ComputeImpl<int8_t>(*context, input);
else if (input.IsDataTypeString())
status = ComputeImpl<std::string>(*context, input);
else

View file

@ -94,6 +94,10 @@ static void SplitTestInt() {
RunTest<T>(axis, {}, input, outputs, false); //TensorRT parser: Assertion failed: axis != BATCH_DIM
}
TEST(SplitOperatorTest, Axis0EqualSplitInt8) {
SplitTestInt<int8_t>();
}
TEST(SplitOperatorTest, Axis0EqualSplitInt32) {
SplitTestInt<int32_t>();
}