Add TopK to ROCm EP (#9391)

* Add TopK to ROCm EP

* flake8 fix
This commit is contained in:
Jeff Daily 2021-10-20 10:39:44 -07:00 committed by GitHub
parent f8acc6d0e8
commit 89a22fb641
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
3 changed files with 20 additions and 10 deletions

View file

@ -56,14 +56,15 @@ struct NumericLimits<MLFloat16> {
template <typename T>
__global__ void BitonicTopK(const T* X, T* V, int64_t* I, const TArray<int64_t> elem_nums, size_t size, int32_t axis, int64_t K, int64_t aligned_K, int64_t largest, int64_t sorted, int64_t dimension, int64_t aligned_dimension, T type_min, T type_max) {
auto tid = threadIdx.x;
auto bid = blockIdx.x;
int64_t tid = threadIdx.x;
int64_t bid = blockIdx.x;
int64_t bdim = blockDim.x;
extern __shared__ char shared_mem[];
auto S = (KV<T>*)(shared_mem);
auto mid_dim = axis == size - 1 ? 1 : elem_nums[axis + 1];
auto left_dim = bid / mid_dim * elem_nums[axis];
auto right_dim = axis == size - 1 ? 0 : bid % elem_nums[axis + 1];
for (auto i = tid; i < aligned_dimension; i += blockDim.x) {
for (auto i = tid; i < aligned_dimension; i += bdim) {
S[i].key = i < dimension ? X[FROM(i)] : TRIVIAL;
S[i].val = i;
}

View file

@ -1551,7 +1551,7 @@ static Status RegisterRocmKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 9, 12, int32_t, NonZero)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 9, 12, int64_t, NonZero)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 9, 12, float, NonZero)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 1, 9, TopK)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 1, 9, TopK)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 8, 8, Scan)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 9, 10, Scan)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 1, 10, Loop)>,
@ -1581,7 +1581,7 @@ static Status RegisterRocmKernels(KernelRegistry& kernel_registry) {
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 10, float, ThresholdedRelu)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 10, double, ThresholdedRelu)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 10, MLFloat16, ThresholdedRelu)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 10, 10, TopK)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 10, 10, TopK)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 1, 10, If)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 10, int8_t, QuantizeLinear)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 10, uint8_t, QuantizeLinear)>,
@ -1659,7 +1659,7 @@ static Status RegisterRocmKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 11, 12, MLFloat16, LogSoftmax)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 11, 12, Split)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 11, 12, Squeeze)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 11, TopK)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 11, TopK)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 11, SequenceAt)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 11, SequenceConstruct)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 11, SequenceEmpty)>,

View file

@ -90,10 +90,6 @@ provider_excluded_files = [
'math/matmul_integer.h',
'math/softmax_impl.cu',
'math/softmax.cc',
'math/topk.cc',
'math/topk.h',
'math/topk_impl.cu',
'math/topk_impl.h',
'nn/batch_norm.cc',
'nn/batch_norm.h',
'nn/conv.cc',
@ -244,6 +240,8 @@ def hipify(src_file_path, dst_file_path):
s = s.replace('std::log', 'logf')
s = s.replace('#include <cub/device/device_radix_sort.cuh>',
'#include <hipcub/hipcub.hpp>\n#include <hipcub/backend/rocprim/device/device_radix_sort.hpp>')
s = s.replace('#include "cub/device/device_radix_sort.cuh"',
'#include <hipcub/hipcub.hpp>\n#include <hipcub/backend/rocprim/device/device_radix_sort.hpp>')
s = s.replace('#include <cub/device/device_reduce.cuh>',
'#include <hipcub/backend/rocprim/device/device_reduce.hpp>')
s = s.replace('#include <cub/device/device_run_length_encode.cuh>',
@ -254,6 +252,14 @@ def hipify(src_file_path, dst_file_path):
'#include <hipcub/backend/rocprim/iterator/counting_input_iterator.hpp>')
s = s.replace('#include <cub/iterator/discard_output_iterator.cuh>',
'#include <hipcub/backend/rocprim/iterator/discard_output_iterator.hpp>')
s = s.replace('#include <cub/util_allocator.cuh>',
'#include <hipcub/util_allocator.hpp>')
s = s.replace('#include "cub/util_allocator.cuh"',
'#include <hipcub/util_allocator.hpp>')
s = s.replace('#include <cub/util_type.cuh>',
'#include <hipcub/backend/rocprim/util_type.hpp>')
s = s.replace('#include "cub/util_type.cuh"',
'#include <hipcub/backend/rocprim/util_type.hpp>')
s = s.replace('typedef half MappedType', 'typedef __half MappedType')
# CUBLAS -> HIPBLAS
@ -301,6 +307,9 @@ def hipify(src_file_path, dst_file_path):
s = s.replace('ROCM_VERSION', 'CUDA_VERSION') # semantically different meanings, cannot hipify
s = s.replace('__ROCM_ARCH__', '__CUDA_ARCH__') # semantically different meanings, cannot hipify
# Deletions
s = s.replace('#include "device_atomic_functions.h"', '') # HIP atomics in main hip header already
do_write = True
if os.path.exists(dst_file_path):
with open(dst_file_path, 'r', encoding='utf-8') as fout_old: