OpSet 11 Update for Neg Axis (9 Ops) (#1893)

* OpSet 11 Update for Neg Axis:
scan, flatten, compress, concat, gather, slice, split, squeeze, unsqueeze

* fix flatten op test

* Fix flatten and Squeeze

* fix test cases

* add gather neg indices to both cpu and cuda

* Exclude  nGraph from neg axis test

* re-enable test cases

* Fix test cases
This commit is contained in:
ybrnathan 2019-09-27 13:04:08 -07:00 committed by GitHub
parent d370aad80c
commit 8df3e87b70
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
21 changed files with 326 additions and 96 deletions

View file

@ -485,11 +485,19 @@ Status ScanImpl::TransposeOutput() {
return status;
}
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(Scan,
9,
10,
KernelDefBuilder()
.TypeConstraint("I", DataTypeImpl::GetTensorType<int64_t>())
.TypeConstraint("V", DataTypeImpl::AllTensorTypes()),
Scan<9>);
// Opset 11 starts to support Neg Axis.
ONNX_CPU_OPERATOR_KERNEL(Scan,
9,
11,
KernelDefBuilder()
.TypeConstraint("I", DataTypeImpl::GetTensorType<int64_t>())
.TypeConstraint("V", DataTypeImpl::AllTensorTypes()),
Scan<9>);
} // namespace onnxruntime

View file

@ -116,7 +116,7 @@ class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDoma
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, Conv);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, ConvTranspose);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 8, Flatten);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Flatten);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, Flatten);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, InstanceNormalization);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, LpNormalization);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, LRN);
@ -172,8 +172,8 @@ class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOn
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 9, float, Cast);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 9, double, Cast);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 9, MLFloat16, Cast);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 4, Concat);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, Gather);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 4, 10, Concat);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, Gather);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 9, Dropout);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, Identity);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 2, Pad);
@ -196,11 +196,11 @@ class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOn
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 9, string, Slice);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, SpaceToDepth);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, DepthToSpace);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 2, Split);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, Squeeze);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 2, 10, Split);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, Squeeze);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, Tile);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, Transpose);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, Unsqueeze);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, Unsqueeze);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 9, float, Upsample);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 9, int32_t, Upsample);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 9, uint8_t, Upsample);
@ -221,7 +221,7 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, If)
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, Loop);
// Opset 9
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Compress);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, Compress);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, ConstantOfShape);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, MeanVarianceNormalization);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int32_t, Greater);
@ -250,7 +250,7 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Cos
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Asinh);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Acosh);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Atanh);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Scan);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, Scan);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, Scatter);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, string, TfIdfVectorizer);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int32_t, TfIdfVectorizer);
@ -284,19 +284,19 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain,
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, int8_t, MatMulInteger);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, ConvInteger);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, QLinearConv);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, bool, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, float, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, double, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, MLFloat16, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, uint8_t, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, uint16_t, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, uint32_t, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, uint64_t, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, int8_t, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, int16_t, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, int32_t, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, int64_t, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, string, Slice);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, bool, Slice);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, float, Slice);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, double, Slice);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, MLFloat16, Slice);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, uint8_t, Slice);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, uint16_t, Slice);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, uint32_t, Slice);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, uint64_t, Slice);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, int8_t, Slice);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, int16_t, Slice);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, int32_t, Slice);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, int64_t, Slice);
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, string, Slice);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, Dropout);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, NonMaxSuppression);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, IsInf);
@ -323,10 +323,32 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Lo
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Softmax);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Loop);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, DepthToSpace);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Scan);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Flatten);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Compress);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Concat);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Gather);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, bool, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, float, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, double, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, MLFloat16, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, uint8_t, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, uint16_t, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, uint32_t, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, uint64_t, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int8_t, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int16_t, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int32_t, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int64_t, Slice);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, string, Slice);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Split);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Squeeze);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Unsqueeze);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Det);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, ScatterElements);
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, NonMaxSuppression);
void RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
static const BuildKernelCreateInfoFn function_table[] = {
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 10, Clip)>,
@ -427,7 +449,7 @@ void RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, Conv)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, ConvTranspose)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 8, Flatten)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Flatten)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, Flatten)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, InstanceNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, LpNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, LRN)>,
@ -483,8 +505,8 @@ void RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 9, float, Cast)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 9, double, Cast)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, 9, MLFloat16, Cast)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 4, Concat)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, Gather)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 4, 10, Concat)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, Gather)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 9, Dropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, Identity)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 2, Pad)>,
@ -507,11 +529,11 @@ void RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 9, string, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, SpaceToDepth)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, DepthToSpace)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 2, Split)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, Squeeze)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 2, 10, Split)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, Squeeze)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, Tile)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, Transpose)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, Unsqueeze)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, Unsqueeze)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 9, float, Upsample)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 9, int32_t, Upsample)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 9, uint8_t, Upsample)>,
@ -532,7 +554,7 @@ void RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 1, 10, Loop)>,
// Opset 9
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Compress)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, Compress)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, ConstantOfShape)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, MeanVarianceNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int32_t, Greater)>,
@ -561,7 +583,7 @@ void RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Asinh)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Acosh)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Atanh)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, Scan)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, Scan)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, 10, Scatter)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, string, TfIdfVectorizer)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 9, int32_t, TfIdfVectorizer)>,
@ -595,19 +617,19 @@ void RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, int8_t, MatMulInteger)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, ConvInteger)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, QLinearConv)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, bool, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, float, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, double, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, MLFloat16, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, uint8_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, uint16_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, uint32_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, uint64_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, int8_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, int16_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, int32_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, int64_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, string, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, bool, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, float, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, double, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, MLFloat16, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, uint8_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, uint16_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, uint32_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, uint64_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, int8_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, int16_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, int32_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, int64_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, string, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, Dropout)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, 10, NonMaxSuppression)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 10, IsInf)>,
@ -634,6 +656,27 @@ void RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Softmax)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Loop)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, DepthToSpace)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Scan)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Flatten)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Compress)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Concat)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Gather)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, bool, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, float, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, double, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, MLFloat16, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, uint8_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, uint16_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, uint32_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, uint64_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int8_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int16_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int32_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int64_t, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, string, Slice)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Split)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Squeeze)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Unsqueeze)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Det)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, ScatterElements)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, NonMaxSuppression)>,

View file

@ -13,9 +13,19 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
.TypeConstraint("T", DataTypeImpl::AllTensorTypes()),
Flatten);
ONNX_CPU_OPERATOR_KERNEL(
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
Flatten,
9,
10,
KernelDefBuilder()
.Alias(0, 0)
.TypeConstraint("T", DataTypeImpl::AllTensorTypes()),
Flatten);
// Opset 11 starts to support Neg Axis.
ONNX_CPU_OPERATOR_KERNEL(
Flatten,
11,
KernelDefBuilder()
.Alias(0, 0)
.TypeConstraint("T", DataTypeImpl::AllTensorTypes()),

View file

@ -7,6 +7,7 @@
#include "core/framework/op_kernel.h"
#include "gsl/gsl_util"
#include "core/providers/cpu/tensor/utils.h"
#include "core/providers/common.h"
namespace onnxruntime {
@ -21,9 +22,17 @@ class Flatten final : public OpKernel {
if (X == nullptr) return Status(common::ONNXRUNTIME, common::FAIL, "input count mismatch");
const TensorShape& X_shape = X->Shape();
ORT_ENFORCE(gsl::narrow_cast<int64_t>(X_shape.NumDimensions()) >= axis_, "The rank of input tensor must be >= axis");
auto axis = axis_;
// Valid axis range is [-rank, rank] instead of [-rank, rank-1], add additional check to only handle neg axis case.
if (axis < 0)
{
axis = HandleNegativeAxis(axis, X_shape.NumDimensions()); // handle negative and enforce axis is valid
}
ORT_ENFORCE(gsl::narrow_cast<int64_t>(X_shape.NumDimensions()) >= axis, "The rank of input tensor must be >= axis");
Tensor* Y = context->Output(0, TensorShape({X_shape.SizeToDimension(axis_), X_shape.SizeFromDimension(axis_)}));
Tensor* Y = context->Output(0, TensorShape({X_shape.SizeToDimension(axis), X_shape.SizeFromDimension(axis)}));
CopyCpuTensor(X, Y);

View file

@ -2,13 +2,23 @@
// Licensed under the MIT License.
#include "core/providers/cpu/tensor/compress.h"
#include "core/providers/common.h"
using namespace ::onnxruntime::common;
namespace onnxruntime {
ONNX_CPU_OPERATOR_KERNEL(
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
Compress,
9,
10,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::AllTensorTypes())
.TypeConstraint("T1", DataTypeImpl::GetTensorType<bool>()),
Compress);
// Opset 11 starts to support Neg Axis.
ONNX_CPU_OPERATOR_KERNEL(
Compress,
11,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::AllTensorTypes())
.TypeConstraint("T1", DataTypeImpl::GetTensorType<bool>()),
Compress);
@ -17,8 +27,10 @@ Status Compress::Compute(OpKernelContext* ctx) const {
const auto* input_tensor = ctx->Input<Tensor>(0);
size_t rank = input_tensor->Shape().NumDimensions();
auto& input_dimensions = input_tensor->Shape().GetDims();
int64_t axis = axis_;
if (has_axis_) {
ORT_ENFORCE(axis_ < static_cast<int64_t>(rank), "axis greater than input data dimension!");
axis = HandleNegativeAxis(axis, rank); // handle negative and enforce axis is valid
ORT_ENFORCE(axis < static_cast<int64_t>(rank), "axis greater than input data dimension!");
}
const auto* condition = ctx->Input<Tensor>(1);
@ -27,7 +39,7 @@ Status Compress::Compute(OpKernelContext* ctx) const {
int64_t positive_condition_count = 0;
// if has axis, we need to compress on dimension[axis], otherwise compress on the flattened input data
int64_t compress_input_length = has_axis_ ? input_dimensions[axis_] : input_tensor->Shape().Size();
int64_t compress_input_length = has_axis_ ? input_dimensions[axis] : input_tensor->Shape().Size();
int64_t valid_condition_length = compress_input_length < condition_length ? compress_input_length : condition_length;
// Figure out output shape
@ -39,7 +51,7 @@ Status Compress::Compute(OpKernelContext* ctx) const {
std::vector<int64_t> output_dims(input_dimensions);
if (has_axis_) {
output_dims[axis_] = positive_condition_count;
output_dims[axis] = positive_condition_count;
} else {
output_dims.resize(1);
output_dims[0] = positive_condition_count;
@ -60,14 +72,14 @@ Status Compress::Compute(OpKernelContext* ctx) const {
if (has_axis_) {
int64_t axes_left_stride = 1;
int64_t axes_right_stride = 1;
for (int i = 0; i < axis_; ++i) {
for (int i = 0; i < axis; ++i) {
axes_left_stride *= input_dimensions[i];
}
for (auto i = static_cast<size_t>(axis_ + 1); i < rank; ++i) {
for (auto i = static_cast<size_t>(axis + 1); i < rank; ++i) {
axes_right_stride *= input_dimensions[i];
}
int64_t axes_included_right_stride = axes_right_stride * input_dimensions[axis_];
int64_t axes_included_right_stride = axes_right_stride * input_dimensions[axis];
int64_t axes_included_right_stride_bytes = axes_included_right_stride * element_bytes;
ORT_ENFORCE(axes_right_stride >= 0 &&
static_cast<uint64_t>(axes_right_stride) < std::numeric_limits<size_t>::max());

View file

@ -6,9 +6,17 @@
namespace onnxruntime {
ONNX_CPU_OPERATOR_KERNEL(
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
Concat,
4,
10,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::AllTensorTypes()),
Concat);
// Opset 11 starts to support Neg Axis.
ONNX_CPU_OPERATOR_KERNEL(
Concat,
11,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::AllTensorTypes()),
Concat);

View file

@ -7,9 +7,16 @@
namespace onnxruntime {
ONNX_CPU_OPERATOR_KERNEL(
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
Gather,
1,
10,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::AllTensorTypes()).TypeConstraint("Tind", std::vector<MLDataType>{DataTypeImpl::GetTensorType<int32_t>(), DataTypeImpl::GetTensorType<int64_t>()}),
Gather);
ONNX_CPU_OPERATOR_KERNEL(
Gather,
11,
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::AllTensorTypes()).TypeConstraint("Tind", std::vector<MLDataType>{DataTypeImpl::GetTensorType<int32_t>(), DataTypeImpl::GetTensorType<int64_t>()}),
Gather);
@ -49,11 +56,14 @@ Status GatherCopyData(const Tensor* indices_tensor, const uint8_t* src_base, uin
// Check the indices first in case there's a out of bound index.
// We can't merge this code in the omp loop below as omp does not allow return in the loop
auto axis_dim_limit = input_data_shape[axis];
for (int64_t i = 0; i < N; ++i) {
Tin idx = indices_data[i];
if (idx < 0 || idx >= input_data_shape[axis]) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "indices element out of data bounds, idx=", idx,
" data_dim=", input_data_shape[axis]);
if (idx < -axis_dim_limit || idx >= axis_dim_limit) {
return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
"indices element out of data bounds, idx=", idx,
" must be within the inclusive range [", -axis_dim_limit,",", axis_dim_limit - 1, "]");
}
}
@ -67,6 +77,7 @@ Status GatherCopyData(const Tensor* indices_tensor, const uint8_t* src_base, uin
const int64_t src_offset_batch = batch * data_batch_bytes;
const int64_t dst_offset_batch = batch * gathered_batch_bytes;
Tin idx = indices_data[i];
idx = idx < 0 ? idx + static_cast<Tin>(axis_dim_limit) : idx;
const int64_t src_offset = src_offset_batch + idx * block_size;
const int64_t dst_offset = dst_offset_batch + i * block_size;

View file

@ -3,6 +3,7 @@
#include "core/providers/cpu/tensor/slice.h"
#include "core/providers/cpu/tensor/utils.h"
#include "core/providers/common.h"
#include <unordered_map>
#include <limits>
@ -33,9 +34,10 @@ ADD_TYPED_SLICE_V9_OP(bool);
ADD_TYPED_SLICE_V9_OP(string);
#define ADD_TYPED_SLICE_V10_OP(data_type) \
ONNX_CPU_OPERATOR_TYPED_KERNEL( \
ONNX_CPU_OPERATOR_VERSIONED_TYPED_KERNEL( \
Slice, \
10, \
10, \
data_type, \
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<data_type>()) \
.TypeConstraint("Tind", {DataTypeImpl::GetTensorType<int32_t>(), \
@ -56,6 +58,30 @@ ADD_TYPED_SLICE_V10_OP(MLFloat16);
ADD_TYPED_SLICE_V10_OP(bool);
ADD_TYPED_SLICE_V10_OP(string);
#define ADD_TYPED_SLICE_V11_OP(data_type) \
ONNX_CPU_OPERATOR_TYPED_KERNEL( \
Slice, \
11, \
data_type, \
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<data_type>()) \
.TypeConstraint("Tind", {DataTypeImpl::GetTensorType<int32_t>(), \
DataTypeImpl::GetTensorType<int64_t>()}), \
Slice<data_type, true>);
ADD_TYPED_SLICE_V11_OP(uint8_t);
ADD_TYPED_SLICE_V11_OP(uint16_t);
ADD_TYPED_SLICE_V11_OP(uint32_t);
ADD_TYPED_SLICE_V11_OP(uint64_t);
ADD_TYPED_SLICE_V11_OP(int8_t);
ADD_TYPED_SLICE_V11_OP(int16_t);
ADD_TYPED_SLICE_V11_OP(int32_t);
ADD_TYPED_SLICE_V11_OP(int64_t);
ADD_TYPED_SLICE_V11_OP(float);
ADD_TYPED_SLICE_V11_OP(double);
ADD_TYPED_SLICE_V11_OP(MLFloat16);
ADD_TYPED_SLICE_V11_OP(bool);
ADD_TYPED_SLICE_V11_OP(string);
namespace {
// std::clamp doesn't exist until C++17 so create a local version
template <typename T>
@ -85,7 +111,7 @@ Status SliceBase::PrepareForCompute(const std::vector<int64_t>& raw_starts,
std::unordered_set<int64_t> unique_axes;
const auto& dimension_count = input_dimensions.size();
for (size_t axis_index = 0, axes_count = axes.size(); axis_index < axes_count; ++axis_index) {
auto axis = axes[axis_index] < 0 ? axes[axis_index] + static_cast<int64_t>(dimension_count) : axes[axis_index];
auto axis = HandleNegativeAxis(axes[axis_index], dimension_count); // handle negative and enforce axis is valid
if (axis >= static_cast<int64_t>(dimension_count) || axis < 0)
return Status(ONNXRUNTIME, INVALID_ARGUMENT, "'axes' has an axis outside of the tensor dimension count");
if (unique_axes.find(axis) != unique_axes.end())

View file

@ -10,9 +10,21 @@
namespace onnxruntime {
ONNX_CPU_OPERATOR_KERNEL(
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
Split,
2,
10,
KernelDefBuilder().TypeConstraint("T",
std::vector<MLDataType>{
DataTypeImpl::GetTensorType<float>(),
DataTypeImpl::GetTensorType<int32_t>(),
DataTypeImpl::GetTensorType<std::string>()}),
Split);
// Opset 11 starts to support Neg Axis.
ONNX_CPU_OPERATOR_KERNEL(
Split,
11,
KernelDefBuilder().TypeConstraint("T",
std::vector<MLDataType>{
DataTypeImpl::GetTensorType<float>(),

View file

@ -5,12 +5,21 @@
namespace onnxruntime {
ONNX_CPU_OPERATOR_KERNEL(
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
Squeeze,
1,
10,
KernelDefBuilder()
.TypeConstraint("T", DataTypeImpl::AllTensorTypes())
.Alias(0, 0),
Squeeze);
// Opset 11 starts to support Neg Axis.
ONNX_CPU_OPERATOR_KERNEL(
Squeeze,
11,
KernelDefBuilder()
.TypeConstraint("T", DataTypeImpl::AllTensorTypes())
.Alias(0, 0),
Squeeze);
} // namespace onnxruntime

View file

@ -6,6 +6,7 @@
#include "core/common/common.h"
#include "core/framework/op_kernel.h"
#include "utils.h"
#include "core/providers/common.h"
namespace onnxruntime {
@ -29,9 +30,19 @@ class SqueezeBase {
const TensorShape& axes) {
size_t j = 0;
std::vector<int64_t> output_shape;
for (size_t i = 0; i < input_shape.NumDimensions(); ++i) {
if ((j < axes.NumDimensions() && axes[j] == static_cast<int64_t>(i)) ||
(axes.NumDimensions() == 0 && input_shape[i] == 1)) {
auto num_dimensions = input_shape.NumDimensions();
// Handle negtive axis, then resort and uniq.
std::vector<int64_t> axes_corrected(axes.NumDimensions());
for (size_t i = 0; i < axes.NumDimensions(); i++) {
axes_corrected[i] = HandleNegativeAxis(axes[i], num_dimensions);
}
std::sort(axes_corrected.begin(), axes_corrected.end());
axes_corrected.erase(std::unique(axes_corrected.begin(), axes_corrected.end()), axes_corrected.end());
for (size_t i = 0; i < num_dimensions; ++i) {
if ((j < axes_corrected.size() && axes_corrected[j] == static_cast<int64_t>(i)) ||
(axes_corrected.size() == 0 && input_shape[i] == 1)) {
ORT_ENFORCE(input_shape[i] == 1, "Dimension of input ", i, " must be 1 instead of ", input_shape[i],
". shape=", input_shape);
++j;

View file

@ -3,13 +3,24 @@
#include "core/providers/cpu/tensor/unsqueeze.h"
#include "utils.h"
#include "core/providers/common.h"
using namespace ::onnxruntime::common;
namespace onnxruntime {
ONNX_CPU_OPERATOR_KERNEL(
ONNX_CPU_OPERATOR_VERSIONED_KERNEL(
Unsqueeze,
1,
10,
KernelDefBuilder()
.Alias(0, 0)
.TypeConstraint("T", DataTypeImpl::AllTensorTypes()),
Unsqueeze);
ONNX_CPU_OPERATOR_KERNEL(
Unsqueeze,
11,
KernelDefBuilder()
.Alias(0, 0)
.TypeConstraint("T", DataTypeImpl::AllTensorTypes()),
@ -26,6 +37,8 @@ Status UnsqueezeBase::PrepareCompute(OpKernelContext* ctx, Prepare& p) const {
// Set all axes_ indices to 1 in output_dims and check for duplicates
for (int64_t axis : axes_) {
// Valid axis range is [0, output_rank - 1]
axis = HandleNegativeAxis(axis, output_dims.size());
if (axis < 0 || axis >= static_cast<int64_t>(output_dims.size()))
return Status(ONNXRUNTIME, INVALID_ARGUMENT, "'axes' has an out of range axis");
if (output_dims[axis] != 0)

View file

@ -24,6 +24,7 @@ __global__ void _GatherKernel(
div_strides[1].divmod(block_offset, indices_index, offset);
int block_size = div_strides[1].d_;
int64_t idx = indices_data[indices_index];
idx = idx < 0 ? idx + indices_max : idx;
if (idx < 0 || idx >= indices_max) {
output_data[id] = 0;
return;

View file

@ -445,16 +445,6 @@ int real_main(int argc, char* argv[], Ort::Env& env) {
{"sequence_model3", "SequenceConstruct not implemented yet"},
{"sequence_model2", "SequenceConstruct not implemented yet"},
{"sequence_model1", "Sequence* not implemented yet"},
{"unsqueeze_unsorted_axes", "Unsqueeze not implemented yet"},
{"unsqueeze_two_axes", "Unsqueeze not implemented yet"},
{"unsqueeze_three_axes", "Unsqueeze not implemented yet"},
{"unsqueeze_negative_axes", "Unsqueeze not implemented yet"},
{"unsqueeze_axis_3", "Unsqueeze not implemented yet"},
{"unsqueeze_axis_2", "Unsqueeze not implemented yet"},
{"unsqueeze_axis_1", "Unsqueeze not implemented yet"},
{"unsqueeze_axis_0", "Unsqueeze not implemented yet"},
{"squeeze_negative_axes", "Squeeze(11) not implemented yet"},
{"slice_negative_axes", "Slice(11) not implemented yet"},
{"scatter_elements_with_negative_indices", "ScatterElements(11) not implemented yet"},
{"reduce_sum_square_negative_axes_keepdims_random", "ReduceSumSquare(11) not implemented yet"},
{"reduce_sum_square_negative_axes_keepdims_example", "ReduceSumSquare(11) not implemented yet"},
@ -480,20 +470,9 @@ int real_main(int argc, char* argv[], Ort::Env& env) {
{"onehot_with_axis", "OneHot(11) not implemented yet"},
{"onehot_negative_indices", "OneHot(11) not implemented yet"},
{"gather_elements_negative_indices", "GatherElements(11) not implemented yet"},
{"flatten_negative_axis4", "Flatten(11) not implemented yet"},
{"flatten_negative_axis3", "Flatten(11) not implemented yet"},
{"flatten_negative_axis2", "Flatten(11) not implemented yet"},
{"flatten_negative_axis1", "Flatten(11) not implemented yet"},
{"reflect_pad", "Pad(11) not implemented yet"},
{"edge_pad", "Pad(11) not implemented yet"},
{"constant_pad", "Pad(11) not implemented yet"},
{"concat_3d_axis_negative_3", "Concat(11) not implemented yet"},
{"concat_3d_axis_negative_2", "Concat(11) not implemented yet"},
{"concat_3d_axis_negative_1", "Concat(11) not implemented yet"},
{"concat_2d_axis_negative_2", "Concat(11) not implemented yet"},
{"concat_2d_axis_negative_1", "Concat(11) not implemented yet"},
{"concat_1d_axis_negative_1", "Concat(11) not implemented yet"},
{"compress_negative_axis", "Compress(11) not implemented yet"},
{"bitshift_right_uint8", "BitShift(11) not implemented yet"},
{"bitshift_right_uint64", "BitShift(11) not implemented yet"},
{"bitshift_right_uint32", "BitShift(11) not implemented yet"},
@ -523,6 +502,12 @@ int real_main(int argc, char* argv[], Ort::Env& env) {
broken_tests.insert({"argmin_negative_axis_keepdims_random", "not implemented yet for opset 11"});
broken_tests.insert({"gemm_default_no_bias", "not implemented yet for opset 11"});
broken_tests.insert({"hardmax_negative_axis", "not implemented yet for opset 11"});
broken_tests.insert({"flatten_negative_axis1", "not implemented yet for opset 11"});
broken_tests.insert({"flatten_negative_axis2", "not implemented yet for opset 11"});
broken_tests.insert({"flatten_negative_axis3", "not implemented yet for opset 11"});
broken_tests.insert({"flatten_negative_axis4", "not implemented yet for opset 11"});
broken_tests.insert({"squeeze_negative_axes", "not implemented yet for opset 11"});
broken_tests.insert({"unsqueeze_negative_axes", "not implemented yet for opset 11"});
#endif
#ifdef USE_MKLDNN

View file

@ -11,7 +11,7 @@ namespace test {
class FlattenOpTest : public testing::Test {
public:
FlattenOpTest() : test_("Flatten"), data0_(120, 1.0f) {}
FlattenOpTest() : test_("Flatten", 11), data0_(120, 1.0f) {}
protected:
OpTester test_;
@ -56,5 +56,13 @@ TEST_F(FlattenOpTest, Flatten_axis4) {
test_.AddOutput<float>("output", {16L, 1L}, data1_);
test_.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider});
}
TEST_F(FlattenOpTest, Flatten_neg_axis3) {
test_.AddAttribute<int64_t>("axis", -1L);
test_.AddInput<float>("data", {2L, 3L, 4L, 5L}, data0_);
test_.AddOutput<float>("output", {24L, 5L}, data0_);
test_.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider, kNGraphExecutionProvider});
}
} // namespace test
} // namespace onnxruntime

View file

@ -153,5 +153,23 @@ TEST(CompressTest, Compress_default_axis_string) {
test.Run();
}
TEST(CompressTest, Compress_3dims_neg_axis) {
OpTester test("Compress", 11);
test.AddAttribute("axis", int64_t(-2));
test.AddInput<float>("input", {2, 2, 3}, {
1.0f, 2.0f, 3.0f,
4.0f, 5.0f, 6.0f,
7.0f, 8.0f, 9.0f,
10.0f, 11.0f, 12.0f});
test.AddInput<bool>("condition", {2}, {0, 1});
test.AddOutput<float>("output", {2, 1, 3}, {
4.0f, 5.0f, 6.0f,
10.0f, 11.0f, 12.0f});
test.Run();
}
} // namespace Test
} // namespace onnxruntime

View file

@ -305,5 +305,23 @@ TEST(GatherOpTest, Gather_perf) {
test.AddOutput<int32_t>("output", {800, 1, 100}, output);
test.Run();
}
TEST(GatherOpTest, Gather_axis1_neg_indices2d_int8) {
OpTester test("Gather", 11);
test.AddAttribute<int64_t>("axis", 1LL);
test.AddInput<int8_t>("data", {3, 3},
{0, 1, 2,
10, 11, 12,
20, 21, 22});
test.AddInput<int32_t>("indices", {2, 2},
{-2, -3,
-1, -2});
test.AddOutput<int8_t>("output", {3, 2, 2},
{1, 0, 2, 1,
11, 10, 12, 11,
21, 20, 22, 21});
test.Run();
}
} // namespace test
} // namespace onnxruntime

View file

@ -528,5 +528,17 @@ TEST(SliceTest, OptionalAxesInputAloneMissing) {
testv10.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider});
}
TEST(SliceTest, Slice2D_ReverseSubsetOfNegAxes_1) {
RunSliceTest<float>({2, 2},
{1.0f, 2.0f, 3.0f, 4.0f},
{-1},
{std::numeric_limits<int64_t>::max()},
{-1}, // axis = -1 only
{-1},
{2, 2},
{2.0f, 1.0f, 4.0f, 3.0f},
true);
}
} // namespace test
} // namespace onnxruntime

View file

@ -94,5 +94,18 @@ TEST(SqueezeOpTest, BadAxes) {
// Expect failure.
test.Run(OpTester::ExpectResult::kExpectFailure, "Dimension of input 0 must be 1 instead of 3", {kTensorrtExecutionProvider});
}
TEST(SqueezeOpTest, SqueezeNegAxis_2) {
OpTester test("Squeeze", 11);
test.AddAttribute("axes", std::vector<int64_t>{0, -3, -2});
test.AddInput<float>("data", {1, 4, 1, 1, 2},
std::vector<float>{1.0f, 2.0f, 3.0f, 4.0f, 5.0f, 6.0f, 7.0f, 8.0f});
test.AddOutput<float>("squeezed", {4, 2},
std::vector<float>{1.0f, 2.0f, 3.0f, 4.0f, 5.0f, 6.0f, 7.0f, 8.0f});
// nGraph does not support neg axis.
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kNGraphExecutionProvider});
}
} // namespace test
} // namespace onnxruntime

View file

@ -63,5 +63,15 @@ TEST(TensorOpTest, Unsqueeze_OutOfRange) {
test.Run(OpTester::ExpectResult::kExpectFailure, "Mismatch between number of source and target dimensions.");
}
TEST(TensorOpTest, UnsqueezeNegAxis_3) {
OpTester test("Unsqueeze", 11);
test.AddAttribute("axes", std::vector<int64_t>{-4, 1, -6});
test.AddInput<float>("input", {2, 3, 4}, std::vector<float>(2 * 3 * 4, 1.0f));
test.AddOutput<float>("output", {1, 1, 1, 2, 3, 4}, std::vector<float>(2 * 3 * 4, 1.0f));
// nGraph does not support negative axis.
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kNGraphExecutionProvider});
}
} // namespace test
} // namespace onnxruntime

View file

@ -144,18 +144,11 @@ def create_backend_test(testname=None):
'^test_resize_upsample_sizes_nearest_round_prefer_ceil_asymmetric_cpu.*',
'^test_scatternd_cpu.*',
'^test_sequence_*',
'^test_unsqueeze_*',
'^test_squeeze_*',
'^test_slice_*',
'^test_scatter_*',
'^test_reduce_*',
'^test_onehot_*',
'^test_flatten_*',
'^test_concat_*',
'^test_compress_*',
'^test_constant_pad_cpu.*',
'^test_gemm_default_scalar_bias_cpu.*',
'^test_gather_negative_indices_cpu.*',
'^test_gemm_*',
'^test_edge_pad_cpu.*',
'^test_reflect_pad_cpu.*'