diff --git a/onnxruntime/core/providers/cuda/math/topk_impl.cu b/onnxruntime/core/providers/cuda/math/topk_impl.cu index f24a0f9ffb..0254813042 100644 --- a/onnxruntime/core/providers/cuda/math/topk_impl.cu +++ b/onnxruntime/core/providers/cuda/math/topk_impl.cu @@ -26,7 +26,7 @@ struct KV { template struct NumericLimits { - static T Lowest() { + static T Min() { return std::numeric_limits::lowest(); } static T Max() { @@ -36,7 +36,7 @@ struct NumericLimits { template <> struct NumericLimits { - static half Lowest() { + static half Min() { return -65504.0; } static half Max() { @@ -44,6 +44,26 @@ struct NumericLimits { } }; +template <> +struct NumericLimits { + static float Min() { + return -INFINITY; + } + static float Max() { + return INFINITY; + } +}; + +template <> +struct NumericLimits { + static double Min() { + return -HUGE_VAL; + } + static double Max() { + return HUGE_VAL; + } +}; + #define BT GridDim::maxThreadsPerBlock #define ALIGN(N) static_cast(pow(2, ceil(log2(static_cast(N))))) #define FROM(idx) (left_dim + (idx)*mid_dim + right_dim) @@ -426,17 +446,17 @@ Status TopKImpl(const CudaKernel* kernel, const T* input_x, T* output_v, int64_t auto aligned_K = ALIGN(K); auto aligned_dimension = ALIGN(dimension); if (aligned_dimension <= GridDim::maxThreadsPerBlock) { - BitonicTopK<<), stream>>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, aligned_K, largest, sorted, dimension, aligned_dimension, NumericLimits::Lowest(), NumericLimits::Max()); + BitonicTopK<<), stream>>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, aligned_K, largest, sorted, dimension, aligned_dimension, NumericLimits::Min(), NumericLimits::Max()); } else if (K <= BT*16 || 0 == sorted) { auto XPT = static_cast(ceil(static_cast(dimension) / GridDim::maxThreadsPerBlock)); if (BT*2 >= K || 0 == sorted) { - RadixTopK<<>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, largest, sorted, dimension, XPT, NumericLimits::Lowest(), NumericLimits::Max()); + RadixTopK<<>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, largest, sorted, dimension, XPT, NumericLimits::Min(), NumericLimits::Max()); } else if (BT*4>=K) { - RadixTopK<<>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, largest, sorted, dimension, XPT, NumericLimits::Lowest(), NumericLimits::Max()); + RadixTopK<<>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, largest, sorted, dimension, XPT, NumericLimits::Min(), NumericLimits::Max()); } else if (BT*8>=K) { - RadixTopK<<>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, largest, sorted, dimension, XPT, NumericLimits::Lowest(), NumericLimits::Max()); + RadixTopK<<>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, largest, sorted, dimension, XPT, NumericLimits::Min(), NumericLimits::Max()); } else { - RadixTopK<<>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, largest, sorted, dimension, XPT, NumericLimits::Lowest(), NumericLimits::Max()); + RadixTopK<<>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, largest, sorted, dimension, XPT, NumericLimits::Min(), NumericLimits::Max()); } } else { auto input_key_buffer = kernel->GetScratchBuffer(dimension); diff --git a/onnxruntime/test/providers/cpu/math/topk_op_test.cc b/onnxruntime/test/providers/cpu/math/topk_op_test.cc index 5b3f9e9219..39d041e724 100644 --- a/onnxruntime/test/providers/cpu/math/topk_op_test.cc +++ b/onnxruntime/test/providers/cpu/math/topk_op_test.cc @@ -582,6 +582,25 @@ TEST(TopKOperator, Top3ExplicitAxisSmallestElements) { top_3_explicit_axis_smallest(11, 0); //unsorted } +template +static void top_3_explicit_aix_infinity(int opset_version, bool positive) { + T inf = positive ? std::numeric_limits::infinity() : -std::numeric_limits::infinity(); + std::vector input_vals = {inf, inf, inf, inf, inf, inf, inf, inf}; + std::vector input_dimensions = {4, 2}; + std::vector expected_vals = {inf, inf, inf, inf, inf, inf}; + std::vector expected_indices = {0, 0, 1, 1, 2, 2}; + std::vector expected_dimensions = {3, 2}; + int64_t axis = 0; + RunTest(opset_version, 3, input_vals, input_dimensions, expected_vals, expected_indices, expected_dimensions, true, axis, 0, 1); +} + +TEST(TopKOperator, Top3ExplicitAxisInfinity) { + top_3_explicit_aix_infinity(11, true); + top_3_explicit_aix_infinity(11, false); + top_3_explicit_aix_infinity(11, true); + top_3_explicit_aix_infinity(11, false); +} + template static void top_1_explicit_axis_MultiD_input_smallest(int opset_version, int64_t sorted = 1) { std::vector input_vals = {1, 2, 3, 4, 5, 6, 7, 8};