[ROCm] InstanceNormalization, BatchNormalization and LRN Ops (#11972)

* [ROCm] Add InstanceNormalization Op

* Enable InstanceNormBatch1_fp16 and InstanceNormBatch2_fp16 for ROCm

* [ROCm] Add BatchNormalization for fp32 and fp16

* Enable BatchNormTest for ROCm

* [ROCm] Add LRN Op

* [ROCM] replace miCompat functions with Helper functions
This commit is contained in:
Xinya Zhang 2022-08-03 01:14:26 -05:00 committed by GitHub
parent 99d2a63e1a
commit 01f3a197d7
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
10 changed files with 248 additions and 41 deletions

View file

@ -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

View file

@ -107,7 +107,7 @@ Status BatchNorm<T>::ComputeInternal(OpKernelContext* p_op_kernel_context) const
Impl_Cast<CudaT, float>(Stream(), mean_data, f_mean.get(), C);
Impl_Cast<CudaT, float>(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<T>::ComputeInternal(OpKernelContext* p_op_kernel_context) const
auto saved_mean_data = reinterpret_cast<CudaT*>(saved_mean->MutableData<T>());
auto saved_inv_var_data = reinterpret_cast<CudaT*>(saved_var->MutableData<T>());
CUDNN_RETURN_IF_ERROR(cudnnBatchNormalizationForwardTraining(
CUDNN_RETURN_IF_ERROR(BatchNormalizationForwardTrainingHelper(
CudnnHandle(),
cudnn_batch_norm_mode_,
&alpha,
@ -156,7 +156,7 @@ Status BatchNorm<T>::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,

View file

@ -69,7 +69,7 @@ Status InstanceNorm<T>::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<T>::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<MLFloat16>::ComputeInternal(OpKernelContext* p_op_kernel_con
auto bias_data_fp32 = GetScratchBuffer<float>(C);
Impl_Cast<CudaT, float>(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<MLFloat16>::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,

View file

@ -76,7 +76,7 @@ Status LRN<T>::ComputeInternal(OpKernelContext* context) const {
const auto one = Consts<CudaT>::One;
const auto zero = Consts<CudaT>::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();
}

View file

@ -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<double>() {
return miopenDouble;
}
template<>
miopenDataType_t MiopenTensor::GetDataType<float>() {
return miopenFloat;

View file

@ -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<void*>(alpha),
const_cast<void*>(beta),
xDesc,
x,
yDesc,
y,
bnScaleBiasMeanVarDesc,
const_cast<void*>(bnScale),
const_cast<void*>(bnBias),
const_cast<void*>(estimatedMean),
const_cast<void*>(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<void*>(alpha),
const_cast<void*>(beta),
xDesc,
x,
yDesc,
y,
bnScaleBiasMeanVarDesc,
const_cast<void*>(bnScale),
const_cast<void*>(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

View file

@ -1411,15 +1411,22 @@ static Status RegisterRocmKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 9, 12, double, Erf)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 9, 12, MLFloat16, Erf)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 1, bool, Not)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 7, 8, float, BatchNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
7, 8, float, BatchNormalization)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 7, 8, double, BatchNormalization)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 7, 8, MLFloat16, BatchNormalization)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 9, 13, float, BatchNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
7, 8, MLFloat16, BatchNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
9, 13, float, BatchNormalization)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 9, 13, double, BatchNormalization)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 9, 13, MLFloat16, BatchNormalization)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 1, 12, float, LRN)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 1, 12, double, LRN)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 1, 12, MLFloat16, LRN)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
9, 13, MLFloat16, BatchNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
1, 12, float, LRN)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
1, 12, double, LRN)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
1, 12, MLFloat16, LRN)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 1, 10, float, Conv)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 1, 10, double, Conv)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 1, 10, MLFloat16, Conv)>,
@ -1520,9 +1527,12 @@ static Status RegisterRocmKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 6, 12, Tile)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, Tile)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 1, 12, Transpose)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 6, float, InstanceNormalization)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 6, double, InstanceNormalization)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 6, MLFloat16, InstanceNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
6, float, InstanceNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
6, double, InstanceNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
6, MLFloat16, InstanceNormalization)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 7, 13, float, RNN)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 7, 13, double, RNN)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 7, 13, MLFloat16, RNN)>,
@ -1977,9 +1987,12 @@ static Status RegisterRocmKernels(KernelRegistry& kernel_registry) {
// BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, If)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, Loop)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, Flatten)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, float, LRN)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, double, LRN)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, MLFloat16, LRN)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
13, float, LRN)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
13, double, LRN)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
13, MLFloat16, LRN)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, 13, Identity)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, ScatterND)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, float, Pad)>,
@ -2046,9 +2059,11 @@ static Status RegisterRocmKernels(KernelRegistry& kernel_registry) {
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 14, double, LSTM)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 14, MLFloat16, LSTM)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 14, Reshape)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 14, 14, float, BatchNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
14, 14, float, BatchNormalization)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 14, 14, double, BatchNormalization)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 14, 14, MLFloat16, BatchNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
14, 14, MLFloat16, BatchNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 14, float, ReduceMin)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 14, double, ReduceMin)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 14, MLFloat16, ReduceMin)>,
@ -2065,9 +2080,11 @@ static Status RegisterRocmKernels(KernelRegistry& kernel_registry) {
// OpSet 15
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 15, Pow)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 15, float, BatchNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
15, float, BatchNormalization)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 15, double, BatchNormalization)>,
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 15, MLFloat16, BatchNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
15, MLFloat16, BatchNormalization)>,
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 15, Shape)>,
// Opset 16

View file

@ -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<float> 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<float>("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

View file

@ -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");

View file

@ -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")