diff --git a/onnxruntime/core/providers/cuda/cudnn_common.h b/onnxruntime/core/providers/cuda/cudnn_common.h index de720521be..9ea3794d5a 100644 --- a/onnxruntime/core/providers/cuda/cudnn_common.h +++ b/onnxruntime/core/providers/cuda/cudnn_common.h @@ -139,5 +139,95 @@ inline double ClampCudnnBatchNormEpsilon(double epsilon) { return epsilon; } +inline cudnnStatus_t +BatchNormalizationForwardInferenceHelper(cudnnHandle_t handle, + cudnnBatchNormMode_t mode, + const void *alpha, + const void *beta, + const cudnnTensorDescriptor_t xDesc, + const void *x, + const cudnnTensorDescriptor_t yDesc, + void *y, + const cudnnTensorDescriptor_t bnScaleBiasMeanVarDesc, + const void *bnScale, + const void *bnBias, + const void *estimatedMean, + const void *estimatedVariance, + double epsilon) { + return cudnnBatchNormalizationForwardInference(handle, + mode, + alpha, + beta, + xDesc, + x, + yDesc, + y, + bnScaleBiasMeanVarDesc, + bnScale, + bnBias, + estimatedMean, + estimatedVariance, + epsilon); +} + +inline cudnnStatus_t +BatchNormalizationForwardTrainingHelper(cudnnHandle_t handle, + cudnnBatchNormMode_t mode, + const void *alpha, + const void *beta, + const cudnnTensorDescriptor_t xDesc, + const void *x, + const cudnnTensorDescriptor_t yDesc, + void *y, + const cudnnTensorDescriptor_t bnScaleBiasMeanVarDesc, + const void *bnScale, + const void *bnBias, + double exponentialAverageFactor, + void *resultRunningMean, + void *resultRunningVariance, + double epsilon, + void *resultSaveMean, + void *resultSaveInvVariance) { + return cudnnBatchNormalizationForwardTraining(handle, + mode, + alpha, + beta, + xDesc, + x, + yDesc, + y, + bnScaleBiasMeanVarDesc, + bnScale, + bnBias, + exponentialAverageFactor, + resultRunningMean, + resultRunningVariance, + epsilon, + resultSaveMean, + resultSaveInvVariance); +} + +inline cudnnStatus_t +LRNCrossChannelForwardHelper(cudnnHandle_t handle, + cudnnLRNDescriptor_t normDesc, + cudnnLRNMode_t lrnMode, + const void *alpha, + const cudnnTensorDescriptor_t xDesc, + const void *x, + const void *beta, + const cudnnTensorDescriptor_t yDesc, + void *y) { + return cudnnLRNCrossChannelForward(handle, normDesc, lrnMode, alpha, xDesc, x, beta, yDesc, y); +} + +inline cudnnStatus_t +SetLRNDescriptorHelper(cudnnLRNDescriptor_t normDesc, + unsigned lrnN, + double lrnAlpha, + double lrnBeta, + double lrnK) { + return cudnnSetLRNDescriptor(normDesc, lrnN, lrnAlpha, lrnBeta, lrnK); +} + } // namespace cuda } // namespace onnxruntime diff --git a/onnxruntime/core/providers/cuda/nn/batch_norm.cc b/onnxruntime/core/providers/cuda/nn/batch_norm.cc index 1742078eb7..414ca06958 100644 --- a/onnxruntime/core/providers/cuda/nn/batch_norm.cc +++ b/onnxruntime/core/providers/cuda/nn/batch_norm.cc @@ -107,7 +107,7 @@ Status BatchNorm::ComputeInternal(OpKernelContext* p_op_kernel_context) const Impl_Cast(Stream(), mean_data, f_mean.get(), C); Impl_Cast(Stream(), var_data, f_var.get(), C); - CUDNN_RETURN_IF_ERROR(cudnnBatchNormalizationForwardInference( + CUDNN_RETURN_IF_ERROR(BatchNormalizationForwardInferenceHelper( CudnnHandle(), cudnn_batch_norm_mode_, &alpha, @@ -136,7 +136,7 @@ Status BatchNorm::ComputeInternal(OpKernelContext* p_op_kernel_context) const auto saved_mean_data = reinterpret_cast(saved_mean->MutableData()); auto saved_inv_var_data = reinterpret_cast(saved_var->MutableData()); - CUDNN_RETURN_IF_ERROR(cudnnBatchNormalizationForwardTraining( + CUDNN_RETURN_IF_ERROR(BatchNormalizationForwardTrainingHelper( CudnnHandle(), cudnn_batch_norm_mode_, &alpha, @@ -156,7 +156,7 @@ Status BatchNorm::ComputeInternal(OpKernelContext* p_op_kernel_context) const saved_inv_var_data)); // in BatchNorm Forward Inference mode if only Y output present } else { - CUDNN_RETURN_IF_ERROR(cudnnBatchNormalizationForwardInference( + CUDNN_RETURN_IF_ERROR(BatchNormalizationForwardInferenceHelper( CudnnHandle(), cudnn_batch_norm_mode_, &alpha, diff --git a/onnxruntime/core/providers/cuda/nn/instance_norm.cc b/onnxruntime/core/providers/cuda/nn/instance_norm.cc index d1599fff2a..05b477ca84 100644 --- a/onnxruntime/core/providers/cuda/nn/instance_norm.cc +++ b/onnxruntime/core/providers/cuda/nn/instance_norm.cc @@ -69,7 +69,7 @@ Status InstanceNorm::ComputeInternal(OpKernelContext* p_op_kernel_context) co CudnnTensor stats_desc; ORT_RETURN_IF_ERROR(stats_desc.Set(data_desc, CUDNN_BATCHNORM_SPATIAL)); - CUDNN_RETURN_IF_ERROR(cudnnBatchNormalizationForwardTraining( + CUDNN_RETURN_IF_ERROR(BatchNormalizationForwardTrainingHelper( CudnnHandle(), CUDNN_BATCHNORM_SPATIAL, &one, @@ -116,7 +116,7 @@ Status InstanceNorm::ComputeInternal(OpKernelContext* p_op_kernel_context) co CUDA_RETURN_IF_ERROR(cudaMemsetAsync(unused_bias.get(), 0, stats_byte_count, Stream())); // first, compute mean and variance per-instance per-channel using cudnnBatchNorm training - CUDNN_RETURN_IF_ERROR(cudnnBatchNormalizationForwardTraining( + CUDNN_RETURN_IF_ERROR(BatchNormalizationForwardTrainingHelper( CudnnHandle(), CUDNN_BATCHNORM_SPATIAL, &one, @@ -207,7 +207,7 @@ Status InstanceNorm::ComputeInternal(OpKernelContext* p_op_kernel_con auto bias_data_fp32 = GetScratchBuffer(C); Impl_Cast(Stream(), bias_data, bias_data_fp32.get(), C); - CUDNN_RETURN_IF_ERROR(cudnnBatchNormalizationForwardTraining( + CUDNN_RETURN_IF_ERROR(BatchNormalizationForwardTrainingHelper( CudnnHandle(), CUDNN_BATCHNORM_SPATIAL, &one, @@ -259,7 +259,7 @@ Status InstanceNorm::ComputeInternal(OpKernelContext* p_op_kernel_con CUDA_RETURN_IF_ERROR(cudaMemsetAsync(unused_bias.get(), 0, stats_byte_count, Stream())); // first, compute mean and variance per-instance per-channel using cudnnBatchNorm training - CUDNN_RETURN_IF_ERROR(cudnnBatchNormalizationForwardTraining( + CUDNN_RETURN_IF_ERROR(BatchNormalizationForwardTrainingHelper( CudnnHandle(), CUDNN_BATCHNORM_SPATIAL, &one, diff --git a/onnxruntime/core/providers/cuda/nn/lrn.cc b/onnxruntime/core/providers/cuda/nn/lrn.cc index 40930040fc..7fd763d11e 100644 --- a/onnxruntime/core/providers/cuda/nn/lrn.cc +++ b/onnxruntime/core/providers/cuda/nn/lrn.cc @@ -76,7 +76,7 @@ Status LRN::ComputeInternal(OpKernelContext* context) const { const auto one = Consts::One; const auto zero = Consts::Zero; - CUDNN_RETURN_IF_ERROR(cudnnLRNCrossChannelForward( + CUDNN_RETURN_IF_ERROR(LRNCrossChannelForwardHelper( CudnnHandle(), norm_desc_, CUDNN_LRN_CROSS_CHANNEL_DIM1, @@ -104,7 +104,7 @@ Status CudnnLRNDescriptor::Set(uint32_t N, double alpha, double beta, double K) if (!desc_) CUDNN_RETURN_IF_ERROR(cudnnCreateLRNDescriptor(&desc_)); - CUDNN_RETURN_IF_ERROR(cudnnSetLRNDescriptor(desc_, N, alpha, beta, K)); + CUDNN_RETURN_IF_ERROR(SetLRNDescriptorHelper(desc_, N, alpha, beta, K)); return Status::OK(); } diff --git a/onnxruntime/core/providers/rocm/miopen_common.cc b/onnxruntime/core/providers/rocm/miopen_common.cc index ae51641a54..6078160106 100644 --- a/onnxruntime/core/providers/rocm/miopen_common.cc +++ b/onnxruntime/core/providers/rocm/miopen_common.cc @@ -81,6 +81,11 @@ miopenDataType_t MiopenTensor::GetDataType() { ORT_THROW("miopen engine currently supports only single/half/int32/int8 precision data types."); } +template<> +miopenDataType_t MiopenTensor::GetDataType() { + return miopenDouble; +} + template<> miopenDataType_t MiopenTensor::GetDataType() { return miopenFloat; diff --git a/onnxruntime/core/providers/rocm/miopen_common.h b/onnxruntime/core/providers/rocm/miopen_common.h index b2da1ae990..6e5b9a02d9 100644 --- a/onnxruntime/core/providers/rocm/miopen_common.h +++ b/onnxruntime/core/providers/rocm/miopen_common.h @@ -102,5 +102,99 @@ inline double ClampMiopenBatchNormEpsilon(double epsilon) { return epsilon; } +inline miopenStatus_t +BatchNormalizationForwardInferenceHelper(miopenHandle_t handle, + miopenBatchNormMode_t mode, + const void *alpha, + const void *beta, + const miopenTensorDescriptor_t xDesc, + const void *x, + const miopenTensorDescriptor_t yDesc, + void *y, + const miopenTensorDescriptor_t bnScaleBiasMeanVarDesc, + const void *bnScale, + const void *bnBias, + const void *estimatedMean, + const void *estimatedVariance, + double epsilon) { + return miopenBatchNormalizationForwardInference(handle, + mode, + const_cast(alpha), + const_cast(beta), + xDesc, + x, + yDesc, + y, + bnScaleBiasMeanVarDesc, + const_cast(bnScale), + const_cast(bnBias), + const_cast(estimatedMean), + const_cast(estimatedVariance), + epsilon); +} + +inline miopenStatus_t +BatchNormalizationForwardTrainingHelper(miopenHandle_t handle, + miopenBatchNormMode_t mode, + const void *alpha, + const void *beta, + const miopenTensorDescriptor_t xDesc, + const void *x, + const miopenTensorDescriptor_t yDesc, + void *y, + const miopenTensorDescriptor_t bnScaleBiasMeanVarDesc, + const void *bnScale, + const void *bnBias, + double exponentialAverageFactor, + void *resultRunningMean, + void *resultRunningVariance, + double epsilon, + void *resultSaveMean, + void *resultSaveInvVariance) { + return miopenBatchNormalizationForwardTraining(handle, + mode, + const_cast(alpha), + const_cast(beta), + xDesc, + x, + yDesc, + y, + bnScaleBiasMeanVarDesc, + const_cast(bnScale), + const_cast(bnBias), + exponentialAverageFactor, + resultRunningMean, + resultRunningVariance, + epsilon, + resultSaveMean, + resultSaveInvVariance); +} + +inline miopenStatus_t +LRNCrossChannelForwardHelper(miopenHandle_t handle, + miopenLRNDescriptor_t normDesc, + miopenLRNMode_t lrnMode, + const void *alpha, + const miopenTensorDescriptor_t xDesc, + const void *x, + const void *beta, + const miopenTensorDescriptor_t yDesc, + void *y) { + if (lrnMode != miopenLRNCrossChannel) { + LOGS_DEFAULT(ERROR) << __func__ << " must be called with lrnMode == miopenLRNCrossChannel"; + return miopenStatusBadParm; + } + return miopenLRNForward(handle, normDesc, alpha, xDesc, x, beta, yDesc, y, false, nullptr); +} + +inline miopenStatus_t +SetLRNDescriptorHelper(miopenLRNDescriptor_t normDesc, + unsigned lrnN, + double lrnAlpha, + double lrnBeta, + double lrnK) { + return miopenSetLRNDescriptor(normDesc, miopenLRNCrossChannel, lrnN, lrnAlpha, lrnBeta, lrnK); +} + } // namespace rocm } // namespace onnxruntime diff --git a/onnxruntime/core/providers/rocm/rocm_execution_provider.cc b/onnxruntime/core/providers/rocm/rocm_execution_provider.cc index 0b7b608d1a..f867672dfe 100644 --- a/onnxruntime/core/providers/rocm/rocm_execution_provider.cc +++ b/onnxruntime/core/providers/rocm/rocm_execution_provider.cc @@ -1411,15 +1411,22 @@ static Status RegisterRocmKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, - // BuildKernelCreateInfo, + BuildKernelCreateInfo, // BuildKernelCreateInfo, - // BuildKernelCreateInfo, - // BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, // BuildKernelCreateInfo, - // BuildKernelCreateInfo, - // BuildKernelCreateInfo, - // BuildKernelCreateInfo, - // BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, // BuildKernelCreateInfo, BuildKernelCreateInfo, @@ -1520,9 +1527,12 @@ static Status RegisterRocmKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, - // BuildKernelCreateInfo, - // BuildKernelCreateInfo, - // BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, // BuildKernelCreateInfo, // BuildKernelCreateInfo, // BuildKernelCreateInfo, @@ -1977,9 +1987,12 @@ static Status RegisterRocmKernels(KernelRegistry& kernel_registry) { // BuildKernelCreateInfo, // BuildKernelCreateInfo, BuildKernelCreateInfo, - // BuildKernelCreateInfo, - // BuildKernelCreateInfo, - // BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, @@ -2046,9 +2059,11 @@ static Status RegisterRocmKernels(KernelRegistry& kernel_registry) { // BuildKernelCreateInfo, // BuildKernelCreateInfo, BuildKernelCreateInfo, - // BuildKernelCreateInfo, + BuildKernelCreateInfo, // BuildKernelCreateInfo, - // BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, // BuildKernelCreateInfo, BuildKernelCreateInfo, @@ -2065,9 +2080,11 @@ static Status RegisterRocmKernels(KernelRegistry& kernel_registry) { // OpSet 15 BuildKernelCreateInfo, - // BuildKernelCreateInfo, + BuildKernelCreateInfo, // BuildKernelCreateInfo, - // BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, // Opset 16 diff --git a/onnxruntime/test/providers/cpu/nn/batch_norm_op_test.cc b/onnxruntime/test/providers/cpu/nn/batch_norm_op_test.cc index 88b6ca2518..5394ac2042 100644 --- a/onnxruntime/test/providers/cpu/nn/batch_norm_op_test.cc +++ b/onnxruntime/test/providers/cpu/nn/batch_norm_op_test.cc @@ -644,8 +644,8 @@ TEST(BatchNormTest, NonSpatial_Complicated) { 8); // opset-8 } -// Only CUDA kernel has float 16 support -#ifdef USE_CUDA +// Only CUDA and ROCm kernels have float 16 support +#if defined(USE_CUDA) || defined(USE_ROCM) TEST(BatchNormTest, BatchNorm2d_fp16) { vector X{-0.91221f, -0.283559f, 0.937637f, 2.09818f, -0.100199f, -0.608113f, 0.444562f, -1.07505f, 0.940591f, -0.922262f, 0.0931303f, 0.69611f, 1.55187f, 0.159808f, 0.914874f, -1.24856f, -1.98928f, -0.331621f, @@ -763,7 +763,9 @@ TEST(BatchNormTest, ForwardTrainingTestWithSavedOutputsOpset9) { // exclude CUDA Execution Provider due to flakiness // exclude TRT and OpenVINO for same reasons as seen in TestBatchNorm() - test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kCudaExecutionProvider, kTensorrtExecutionProvider, kOpenVINOExecutionProvider, kDnnlExecutionProvider}); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", + {kCudaExecutionProvider, kRocmExecutionProvider, + kTensorrtExecutionProvider, kOpenVINOExecutionProvider, kDnnlExecutionProvider}); } TEST(BatchNormTest, ForwardTrainingTestOpset14) { @@ -789,7 +791,9 @@ TEST(BatchNormTest, ForwardTrainingTestOpset14) { // exclude CUDA Execution Provider due to flakiness // exclude TRT and OpenVINO for same reasons as seen in TestBatchNorm() - test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kCudaExecutionProvider, kTensorrtExecutionProvider, kOpenVINOExecutionProvider, kDnnlExecutionProvider}); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", + {kCudaExecutionProvider, kRocmExecutionProvider, + kTensorrtExecutionProvider, kOpenVINOExecutionProvider, kDnnlExecutionProvider}); } TEST(BatchNormTest, ForwardTrainingTestOpset15) { @@ -814,7 +818,9 @@ TEST(BatchNormTest, ForwardTrainingTestOpset15) { test.AddOutput("running_var", channel_dims, {0.696052f, 1.41316f}); // Same exclusions as the opset 14 test - test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kCudaExecutionProvider, kTensorrtExecutionProvider, kOpenVINOExecutionProvider, kDnnlExecutionProvider}); + test.Run(OpTester::ExpectResult::kExpectSuccess, "", + {kCudaExecutionProvider, kRocmExecutionProvider, + kTensorrtExecutionProvider, kOpenVINOExecutionProvider, kDnnlExecutionProvider}); } #endif // BATCHNORM_INCLUDE_TRAINING_SUPPORT diff --git a/onnxruntime/test/providers/cpu/nn/instance_norm_op_test.cc b/onnxruntime/test/providers/cpu/nn/instance_norm_op_test.cc index 45c8ed74f6..7e8486c9e3 100644 --- a/onnxruntime/test/providers/cpu/nn/instance_norm_op_test.cc +++ b/onnxruntime/test/providers/cpu/nn/instance_norm_op_test.cc @@ -116,8 +116,8 @@ TEST(InstanceNormalizationOpTest, InstanceNormBatch2) { #endif } -// Only CUDA kernel has float 16 support -#ifdef USE_CUDA +// Only CUDA and ROCm kernels have float 16 support +#if defined(USE_CUDA) || defined(USE_ROCM) TEST(InstanceNormalizationOpTest, InstanceNormBatch1_fp16) { OpTester test("InstanceNormalization"); diff --git a/tools/ci_build/amd_hipify.py b/tools/ci_build/amd_hipify.py index 7d1d7cdda6..558a2d863a 100644 --- a/tools/ci_build/amd_hipify.py +++ b/tools/ci_build/amd_hipify.py @@ -109,18 +109,10 @@ provider_excluded_files = [ "math/softmax_impl.cu", "math/softmax_warpwise_impl.cuh", "math/softmax.cc", - "nn/batch_norm.cc", - "nn/batch_norm.h", "nn/conv.cc", "nn/conv.h", "nn/conv_transpose.cc", "nn/conv_transpose.h", - "nn/instance_norm.cc", - "nn/instance_norm.h", - "nn/instance_norm_impl.cu", - "nn/instance_norm_impl.h", - "nn/lrn.cc", - "nn/lrn.h", "nn/max_pool_with_index.cu", "nn/max_pool_with_index.h", "nn/pool.cc", @@ -317,6 +309,9 @@ def hipify(src_file_path, dst_file_path): s = s.replace("hipdnn", "miopen") s = s.replace("HIPDNN_STATUS_SUCCESS", "miopenStatusSuccess") s = s.replace("HIPDNN", "MIOPEN") + s = s.replace("MIOPEN_BATCHNORM_SPATIAL", "miopenBNSpatial") + s = s.replace("MIOPEN_BATCHNORM_PER_ACTIVATION", "miopenBNPerActivation") + s = s.replace("MIOPEN_LRN_CROSS_CHANNEL", "miopenLRNCrossChannel") # CUSPARSE -> HIPSPARSE s = s.replace("CUSPARSE", "HIPSPARSE")