mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-27 20:02:15 +00:00
reset MIN for float/double (#10284)
This commit is contained in:
parent
e365ad7f3a
commit
ab5fd42ed4
2 changed files with 46 additions and 7 deletions
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
Loading…
Reference in a new issue