mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
[ROCm] add beam search support (#15625)
add beam search support for ROCm EP.
This commit is contained in:
parent
699c9a520b
commit
0ecfe83932
4 changed files with 34 additions and 5 deletions
|
|
@ -88,8 +88,6 @@ set(contrib_ops_excluded_files
|
|||
"tensor/image_scaler.h"
|
||||
"tensor/image_scaler_impl.cu"
|
||||
"tensor/image_scaler_impl.h"
|
||||
"transformers/beam_search.cc"
|
||||
"transformers/beam_search.h"
|
||||
"transformers/greedy_search.cc"
|
||||
"transformers/greedy_search.h"
|
||||
"conv_transpose_with_dynamic_pads.cc"
|
||||
|
|
|
|||
|
|
@ -57,10 +57,12 @@ BeamSearch::BeamSearch(const OpKernelInfo& info)
|
|||
|
||||
SetConsoleDumper(&g_cuda_dumper);
|
||||
|
||||
#ifndef USE_ROCM
|
||||
cuda_device_prop_ = &reinterpret_cast<const CUDAExecutionProvider*>(info.GetExecutionProvider())->GetDeviceProp();
|
||||
|
||||
cuda_device_arch_ = static_cast<const cudaDeviceProp*>(cuda_device_prop_)->major * 100 +
|
||||
static_cast<const cudaDeviceProp*>(cuda_device_prop_)->minor * 10;
|
||||
#endif
|
||||
}
|
||||
|
||||
Status BeamSearch::ComputeInternal(OpKernelContext* context) const {
|
||||
|
|
|
|||
|
|
@ -55,6 +55,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
|
|||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 1, MLFloat16, Affine);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, Attention);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16, Attention);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, BeamSearch);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, ConvTransposeWithDynamicPads);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 1, float, Crop);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 1, double, Crop);
|
||||
|
|
@ -183,6 +184,8 @@ Status RegisterRocmContribKernels(KernelRegistry& kernel_registry) {
|
|||
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 1, MLFloat16, Affine)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, Attention)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16, Attention)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, BeamSearch)>,
|
||||
|
||||
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, ConvTransposeWithDynamicPads)>,
|
||||
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 1, float, Crop)>,
|
||||
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 1, double, Crop)>,
|
||||
|
|
|
|||
|
|
@ -73,6 +73,11 @@ TEST(BeamSearchTest, GptBeamSearchFp32) {
|
|||
Ort::ThrowOnError(OrtSessionOptionsAppendExecutionProvider_CUDA(session_options, 0));
|
||||
#endif
|
||||
|
||||
#ifdef USE_ROCM
|
||||
OrtROCMProviderOptions rocm_options;
|
||||
session_options.AppendExecutionProvider_ROCM(rocm_options);
|
||||
#endif
|
||||
|
||||
// The ONNX model is generated like the following:
|
||||
// python convert_generation.py --model_type gpt2 -m hf-internal-testing/tiny-random-gpt2
|
||||
// --output tiny_gpt2_beamsearch_fp16.onnx --use_gpu --max_length 20
|
||||
|
|
@ -151,12 +156,19 @@ TEST(BeamSearchTest, GptBeamSearchFp16) {
|
|||
const char* const output_names[] = {"sequences"};
|
||||
|
||||
constexpr int min_cuda_architecture = 530;
|
||||
if (HasCudaEnvironment(min_cuda_architecture)) {
|
||||
bool enable_cuda = HasCudaEnvironment(min_cuda_architecture);
|
||||
bool enable_rocm = (nullptr != DefaultRocmExecutionProvider().get());
|
||||
if (enable_cuda || enable_rocm) {
|
||||
Ort::SessionOptions session_options;
|
||||
#ifdef USE_CUDA
|
||||
Ort::ThrowOnError(OrtSessionOptionsAppendExecutionProvider_CUDA(session_options, 0));
|
||||
#endif
|
||||
|
||||
#ifdef USE_ROCM
|
||||
OrtROCMProviderOptions rocm_options;
|
||||
session_options.AppendExecutionProvider_ROCM(rocm_options);
|
||||
#endif
|
||||
|
||||
// The ONNX model is generated like the following:
|
||||
// python convert_generation.py --model_type gpt2 -m hf-internal-testing/tiny-random-gpt2
|
||||
// --output tiny_gpt2_beamsearch_fp16.onnx -p fp16 --use_gpu --max_length 20
|
||||
|
|
@ -237,12 +249,19 @@ TEST(BeamSearchTest, GptBeamSearchWithInitDecoderFp16) {
|
|||
const char* const output_names[] = {"sequences"};
|
||||
|
||||
constexpr int min_cuda_architecture = 530;
|
||||
if (HasCudaEnvironment(min_cuda_architecture)) {
|
||||
bool enable_cuda = HasCudaEnvironment(min_cuda_architecture);
|
||||
bool enable_rocm = (nullptr != DefaultRocmExecutionProvider().get());
|
||||
if (enable_cuda || enable_rocm) {
|
||||
Ort::SessionOptions session_options;
|
||||
#ifdef USE_CUDA
|
||||
Ort::ThrowOnError(OrtSessionOptionsAppendExecutionProvider_CUDA(session_options, 0));
|
||||
#endif
|
||||
|
||||
#ifdef USE_ROCM
|
||||
OrtROCMProviderOptions rocm_options;
|
||||
session_options.AppendExecutionProvider_ROCM(rocm_options);
|
||||
#endif
|
||||
|
||||
// The ONNX model is generated like the following:
|
||||
// python convert_generation.py --model_type gpt2 -m hf-internal-testing/tiny-random-gpt2
|
||||
// --output tiny_gpt2_beamsearch_with_init_decoder_fp16.onnx -p fp16 --use_gpu --max_length 20
|
||||
|
|
@ -322,12 +341,19 @@ TEST(BeamSearchTest, GptBeamSearchFp16_VocabPadded) {
|
|||
const char* const output_names[] = {"sequences"};
|
||||
|
||||
constexpr int min_cuda_architecture = 530;
|
||||
if (HasCudaEnvironment(min_cuda_architecture)) {
|
||||
bool enable_cuda = HasCudaEnvironment(min_cuda_architecture);
|
||||
bool enable_rocm = (nullptr != DefaultRocmExecutionProvider().get());
|
||||
if (enable_cuda || enable_rocm) {
|
||||
Ort::SessionOptions session_options;
|
||||
#ifdef USE_CUDA
|
||||
Ort::ThrowOnError(OrtSessionOptionsAppendExecutionProvider_CUDA(session_options, 0));
|
||||
#endif
|
||||
|
||||
#ifdef USE_ROCM
|
||||
OrtROCMProviderOptions rocm_options;
|
||||
session_options.AppendExecutionProvider_ROCM(rocm_options);
|
||||
#endif
|
||||
|
||||
// The following model was obtained by padding the vocabulary size in testdata/transformers/tiny_gpt2_beamsearch_fp16.onnx
|
||||
// from 1000 to 1600 (just for illustrative and testing purposes) to see if the beam search implementation can handle
|
||||
// such a scenario
|
||||
|
|
|
|||
Loading…
Reference in a new issue