diff --git a/onnxruntime/core/providers/cpu/cpu_execution_provider.cc b/onnxruntime/core/providers/cpu/cpu_execution_provider.cc index 496d96bba76dd..654173c54ea4d 100644 --- a/onnxruntime/core/providers/cpu/cpu_execution_provider.cc +++ b/onnxruntime/core/providers/cpu/cpu_execution_provider.cc @@ -416,7 +416,8 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, GatherND); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Range); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, Unique); -class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, TopK); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, float, TopK); +class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int64_t, TopK); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int64_t_int64_t_int64_t, OneHot); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, float_int64_t_int64_t, OneHot); class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 11, int64_t_string_int64_t, OneHot); @@ -892,7 +893,7 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, + Where)>, BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo::Compute(OpKernelContext* p_op_kernel_context) const { } // Opset ver - 11 -template <> -TopK<11, float>::TopK(const OpKernelInfo& op_kernel_info) : OpKernel(op_kernel_info) { + +static void TopkOpset11ConstructorCommon(const OpKernelInfo& op_kernel_info, + int& axis, bool& largest, bool& sorted) { int64_t axis_temp; ORT_ENFORCE(op_kernel_info.GetAttr("axis", &axis_temp).IsOK()); - axis_ = gsl::narrow_cast(axis_temp); + axis = gsl::narrow_cast(axis_temp); int64_t largest_temp; ORT_ENFORCE(op_kernel_info.GetAttr("largest", &largest_temp).IsOK()); - largest_ = largest_temp == 1 ? true : false; + largest = largest_temp == 1 ? true : false; int64_t sorted_temp; ORT_ENFORCE(op_kernel_info.GetAttr("sorted", &sorted_temp).IsOK()); - sorted_ = sorted_temp == 1 ? true : false; + sorted = sorted_temp == 1 ? true : false; } -// Opset ver - 11 template <> -Status TopK<11, float>::Compute(OpKernelContext* p_op_kernel_context) const { +TopK<11, float>::TopK(const OpKernelInfo& op_kernel_info) : OpKernel(op_kernel_info) { + TopkOpset11ConstructorCommon(op_kernel_info, axis_, largest_, sorted_); +} + +template <> +TopK<11, int64_t>::TopK(const OpKernelInfo& op_kernel_info) : OpKernel(op_kernel_info) { + TopkOpset11ConstructorCommon(op_kernel_info, axis_, largest_, sorted_); +} + +static Status ComputeImplOpset11(OpKernelContext* p_op_kernel_context, int axis, bool is_largest, bool is_sorted) { const auto* X = p_op_kernel_context->Input(0); const auto* Y = p_op_kernel_context->Input(1); if (X == nullptr || Y == nullptr) { @@ -312,7 +321,18 @@ Status TopK<11, float>::Compute(OpKernelContext* p_op_kernel_context) const { return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "value of k must not be negative"); } - return TopKImpl(p_op_kernel_context, X, axis_, gsl::narrow_cast(parsed_input_k), largest_, sorted_); + return TopKImpl(p_op_kernel_context, X, axis, gsl::narrow_cast(parsed_input_k), is_largest, is_sorted); +} + +// Opset ver - 11 +template <> +Status TopK<11, float>::Compute(OpKernelContext* p_op_kernel_context) const { + return ComputeImplOpset11(p_op_kernel_context, axis_, largest_, sorted_); +} + +template <> +Status TopK<11, int64_t>::Compute(OpKernelContext* p_op_kernel_context) const { + return ComputeImplOpset11(p_op_kernel_context, axis_, largest_, sorted_); } // Register necessary kernels @@ -329,10 +349,16 @@ ONNX_CPU_OPERATOR_VERSIONED_KERNEL(TopK, 10, 10, .TypeConstraint("I", DataTypeImpl::GetTensorType()), TopK<10, float>); -ONNX_CPU_OPERATOR_KERNEL(TopK, 11, - KernelDefBuilder() - .TypeConstraint("T", DataTypeImpl::GetTensorType()) - .TypeConstraint("I", DataTypeImpl::GetTensorType()), - TopK<11, float>); +#define REGISTER_TOPK_TYPED_KERNEL(OPSET, TYPE) \ + ONNX_CPU_OPERATOR_TYPED_KERNEL(TopK, \ + OPSET, \ + TYPE, \ + KernelDefBuilder() \ + .TypeConstraint("T", DataTypeImpl::GetTensorType()) \ + .TypeConstraint("I", DataTypeImpl::GetTensorType()), \ + TopK); + +REGISTER_TOPK_TYPED_KERNEL(11, float); +REGISTER_TOPK_TYPED_KERNEL(11, int64_t); } // namespace onnxruntime