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