mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
parent
f8acc6d0e8
commit
89a22fb641
3 changed files with 20 additions and 10 deletions
|
|
@ -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;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)>,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Reference in a new issue