mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-27 20:02:15 +00:00
Support int8 for operator Split (#8615)
* Support int8 for operator Split
This commit is contained in:
parent
3a742f2910
commit
064a385b59
3 changed files with 24 additions and 7 deletions
|
|
@ -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)|
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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>();
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in a new issue