reset MIN for float/double (#10284)

This commit is contained in:
RandySheriffH 2022-01-14 13:57:29 -08:00 committed by GitHub
parent e365ad7f3a
commit ab5fd42ed4
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
2 changed files with 46 additions and 7 deletions

View file

@ -26,7 +26,7 @@ struct KV {
template <typename T>
struct NumericLimits {
static T Lowest() {
static T Min() {
return std::numeric_limits<T>::lowest();
}
static T Max() {
@ -36,7 +36,7 @@ struct NumericLimits {
template <>
struct NumericLimits<MLFloat16> {
static half Lowest() {
static half Min() {
return -65504.0;
}
static half Max() {
@ -44,6 +44,26 @@ struct NumericLimits<MLFloat16> {
}
};
template <>
struct NumericLimits<float> {
static float Min() {
return -INFINITY;
}
static float Max() {
return INFINITY;
}
};
template <>
struct NumericLimits<double> {
static double Min() {
return -HUGE_VAL;
}
static double Max() {
return HUGE_VAL;
}
};
#define BT GridDim::maxThreadsPerBlock
#define ALIGN(N) static_cast<int64_t>(pow(2, ceil(log2(static_cast<double>(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<CudaT><<<N, GridDim::maxThreadsPerBlock, aligned_dimension * sizeof(KV<CudaT>), stream>>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, aligned_K, largest, sorted, dimension, aligned_dimension, NumericLimits<T>::Lowest(), NumericLimits<T>::Max());
BitonicTopK<CudaT><<<N, GridDim::maxThreadsPerBlock, aligned_dimension * sizeof(KV<CudaT>), stream>>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, aligned_K, largest, sorted, dimension, aligned_dimension, NumericLimits<T>::Min(), NumericLimits<T>::Max());
} else if (K <= BT*16 || 0 == sorted) {
auto XPT = static_cast<int64_t>(ceil(static_cast<double>(dimension) / GridDim::maxThreadsPerBlock));
if (BT*2 >= K || 0 == sorted) {
RadixTopK<CudaT, BT, 2><<<N, BT, 256 * sizeof(uint32_t), stream>>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, largest, sorted, dimension, XPT, NumericLimits<T>::Lowest(), NumericLimits<T>::Max());
RadixTopK<CudaT, BT, 2><<<N, BT, 256 * sizeof(uint32_t), stream>>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, largest, sorted, dimension, XPT, NumericLimits<T>::Min(), NumericLimits<T>::Max());
} else if (BT*4>=K) {
RadixTopK<CudaT, BT, 4><<<N, BT, 256 * sizeof(uint32_t), stream>>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, largest, sorted, dimension, XPT, NumericLimits<T>::Lowest(), NumericLimits<T>::Max());
RadixTopK<CudaT, BT, 4><<<N, BT, 256 * sizeof(uint32_t), stream>>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, largest, sorted, dimension, XPT, NumericLimits<T>::Min(), NumericLimits<T>::Max());
} else if (BT*8>=K) {
RadixTopK<CudaT, BT, 8><<<N, BT, 256 * sizeof(uint32_t), stream>>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, largest, sorted, dimension, XPT, NumericLimits<T>::Lowest(), NumericLimits<T>::Max());
RadixTopK<CudaT, BT, 8><<<N, BT, 256 * sizeof(uint32_t), stream>>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, largest, sorted, dimension, XPT, NumericLimits<T>::Min(), NumericLimits<T>::Max());
} else {
RadixTopK<CudaT, BT, 16><<<N, BT, 256 * sizeof(uint32_t), stream>>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, largest, sorted, dimension, XPT, NumericLimits<T>::Lowest(), NumericLimits<T>::Max());
RadixTopK<CudaT, BT, 16><<<N, BT, 256 * sizeof(uint32_t), stream>>>(input_x_ptr, output_v_ptr, output_i, elem_nums, size, axis, K, largest, sorted, dimension, XPT, NumericLimits<T>::Min(), NumericLimits<T>::Max());
}
} else {
auto input_key_buffer = kernel->GetScratchBuffer<CudaT>(dimension);

View file

@ -582,6 +582,25 @@ TEST(TopKOperator, Top3ExplicitAxisSmallestElements) {
top_3_explicit_axis_smallest<double>(11, 0); //unsorted
}
template<typename T>
static void top_3_explicit_aix_infinity(int opset_version, bool positive) {
T inf = positive ? std::numeric_limits<T>::infinity() : -std::numeric_limits<T>::infinity();
std::vector<T> input_vals = {inf, inf, inf, inf, inf, inf, inf, inf};
std::vector<int64_t> input_dimensions = {4, 2};
std::vector<T> expected_vals = {inf, inf, inf, inf, inf, inf};
std::vector<int64_t> expected_indices = {0, 0, 1, 1, 2, 2};
std::vector<int64_t> 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<float>(11, true);
top_3_explicit_aix_infinity<float>(11, false);
top_3_explicit_aix_infinity<double>(11, true);
top_3_explicit_aix_infinity<double>(11, false);
}
template <typename T>
static void top_1_explicit_axis_MultiD_input_smallest(int opset_version, int64_t sorted = 1) {
std::vector<T> input_vals = {1, 2, 3, 4, 5, 6, 7, 8};