mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
rehipify ROCm EP files under orttraining (#9443)
* rehipify rocm ep files under orttraining committed to source control * fix flake8 error
This commit is contained in:
parent
ff23b9ff55
commit
66ceb6926d
11 changed files with 126 additions and 243 deletions
|
|
@ -1,136 +0,0 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include "orttraining/training_ops/rocm/collective/nccl_common.h"
|
||||
#include <mpi.h>
|
||||
|
||||
#include "orttraining/core/framework/communication/mpi/mpi_context.h"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace rocm {
|
||||
|
||||
ncclDataType_t GetNcclDataType(onnxruntime::MLDataType type) {
|
||||
if (type == DataTypeImpl::GetType<uint8_t>()) {
|
||||
return ncclUint8;
|
||||
} else if (type == DataTypeImpl::GetType<int8_t>()) {
|
||||
return ncclInt8;
|
||||
} else if (type == DataTypeImpl::GetType<int32_t>()) {
|
||||
return ncclInt32;
|
||||
} else if (type == DataTypeImpl::GetType<int64_t>()) {
|
||||
return ncclInt64;
|
||||
} else if (type == DataTypeImpl::GetType<MLFloat16>()) {
|
||||
return ncclFloat16;
|
||||
} else if (type == DataTypeImpl::GetType<float>()) {
|
||||
return ncclFloat32;
|
||||
} else if (type == DataTypeImpl::GetType<double>()) {
|
||||
return ncclFloat64;
|
||||
} else {
|
||||
// throw logfic_error("Tensor type not supported in NCCL.");
|
||||
}
|
||||
return ncclFloat32;
|
||||
}
|
||||
|
||||
#ifdef USE_MPI
|
||||
static Status CreateNcclCommunicator(MPI_Group* mpi_world_group,
|
||||
const training::WorkerGroupType worker_group_type,
|
||||
ncclComm_t* group_comm) {
|
||||
auto worker_group = training::DistributedRunContext::GetInstance().GetWorkerGroup(worker_group_type);
|
||||
if (worker_group.ranks.size() == 1) {
|
||||
LOGS_DEFAULT(WARNING) << "Target group size = 1, skip creating nccl communicator. Group info: "
|
||||
<< worker_group.ToString();
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
// Create new group
|
||||
MPI_Group mpi_group;
|
||||
MPI_CHECK(MPI_Group_incl(*mpi_world_group, worker_group.ranks.size(), worker_group.ranks.data(), &mpi_group));
|
||||
|
||||
// Create new MPI communicator
|
||||
MPI_Comm mpi_comm;
|
||||
static int32_t mpi_group_id = 0;
|
||||
MPI_CHECK(MPI_Comm_create_group(MPI_COMM_WORLD, mpi_group, ++mpi_group_id, &(mpi_comm)));
|
||||
ORT_ENFORCE(mpi_comm != MPI_COMM_NULL, "MPI communicator creation failed.");
|
||||
|
||||
// Create new NCCL communicator
|
||||
ncclUniqueId nccl_id;
|
||||
if (worker_group.rank_in_group == 0) {
|
||||
NCCL_RETURN_IF_ERROR(ncclGetUniqueId(&nccl_id));
|
||||
}
|
||||
MPI_CHECK(MPI_Bcast(&nccl_id, sizeof(nccl_id), MPI_BYTE, 0, mpi_comm));
|
||||
NCCL_RETURN_IF_ERROR(ncclCommInitRank(group_comm, worker_group.ranks.size(), nccl_id, worker_group.rank_in_group));
|
||||
|
||||
// Clean up
|
||||
MPI_CHECK(MPI_Group_free(&mpi_group));
|
||||
MPI_CHECK(MPI_Comm_free(&mpi_comm));
|
||||
return Status::OK();
|
||||
}
|
||||
#endif
|
||||
|
||||
NcclContext::NcclContext() {
|
||||
#ifdef USE_MPI
|
||||
int is_mpi_initialized = 0;
|
||||
MPI_Initialized(&is_mpi_initialized);
|
||||
if (!is_mpi_initialized) {
|
||||
int mpi_threads_provided = 0;
|
||||
MPI_Init_thread(nullptr, nullptr, MPI_THREAD_MULTIPLE, &mpi_threads_provided);
|
||||
}
|
||||
|
||||
// Get the group under MPI_COMM_WORLD
|
||||
MPI_Group mpi_world_group;
|
||||
MPI_Comm_group(MPI_COMM_WORLD, &mpi_world_group);
|
||||
|
||||
// Initialize Data Parallel Group NCCL Communicator
|
||||
auto ret = CreateNcclCommunicator(&mpi_world_group, training::WorkerGroupType::DataParallel,
|
||||
&data_group_comm_);
|
||||
ORT_ENFORCE(ret.IsOK());
|
||||
|
||||
// Initialize Horizontal Model Parallel Group NCCL Communicator
|
||||
ret = CreateNcclCommunicator(&mpi_world_group, training::WorkerGroupType::HorizontalParallel,
|
||||
&horizontal_group_comm_);
|
||||
ORT_ENFORCE(ret.IsOK());
|
||||
|
||||
MPI_Group_free(&mpi_world_group);
|
||||
#else
|
||||
ORT_THROW("ORT must be built with MPI to use NCCL.");
|
||||
#endif
|
||||
|
||||
}
|
||||
|
||||
ncclComm_t NcclContext::Comm(training::WorkerGroupType group_type) {
|
||||
if (training::WorkerGroupType::DataParallel == group_type) {
|
||||
return data_group_comm_;
|
||||
} else if (training::WorkerGroupType::HorizontalParallel == group_type) {
|
||||
return horizontal_group_comm_;
|
||||
}
|
||||
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
NcclContext::~NcclContext() {
|
||||
if (data_group_comm_ != nullptr) {
|
||||
ncclCommDestroy(data_group_comm_);
|
||||
}
|
||||
|
||||
if (horizontal_group_comm_ != nullptr) {
|
||||
ncclCommDestroy(horizontal_group_comm_);
|
||||
}
|
||||
|
||||
#ifdef USE_MPI
|
||||
int is_mpi_finalized = 0;
|
||||
MPI_Finalized(&is_mpi_finalized);
|
||||
if (!is_mpi_finalized) {
|
||||
MPI_Finalize();
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
NcclKernel::NcclKernel(const OpKernelInfo& info) : RocmKernel(info) {
|
||||
static NcclContext context;
|
||||
nccl_ = &context;
|
||||
int64_t group_type;
|
||||
info.GetAttrOrDefault("group_type", &group_type, static_cast<int64_t>(0));
|
||||
group_type_ = static_cast<training::WorkerGroupType>(group_type);
|
||||
}
|
||||
|
||||
} // namespace rocm
|
||||
} // namespace onnxruntime
|
||||
|
|
@ -96,13 +96,13 @@ Status DivGrad<T>::ComputeInternal(OpKernelContext* context) const {
|
|||
|
||||
if (da_output_tensor) {
|
||||
std::vector<int64_t> a_output_dims = prepended_dimension_1(a_shape, dy_shape.NumDimensions());
|
||||
ReduceKernelShared<T, T, MIOPEN_REDUCE_TENSOR_NO_INDICES>(
|
||||
ORT_RETURN_IF_ERROR((ReduceKernelShared<T, T, MIOPEN_REDUCE_TENSOR_NO_INDICES>(
|
||||
temp_da_data,
|
||||
dy_shape,
|
||||
da_data,
|
||||
TensorShape({}),
|
||||
MIOPEN_REDUCE_TENSOR_ADD,
|
||||
a_output_dims);
|
||||
a_output_dims)));
|
||||
}
|
||||
break;
|
||||
}
|
||||
|
|
@ -125,13 +125,13 @@ Status DivGrad<T>::ComputeInternal(OpKernelContext* context) const {
|
|||
|
||||
if (db_output_tensor) {
|
||||
std::vector<int64_t> b_output_dims = prepended_dimension_1(b_shape, dy_shape.NumDimensions());
|
||||
ReduceKernelShared<T, T, MIOPEN_REDUCE_TENSOR_NO_INDICES>(
|
||||
ORT_RETURN_IF_ERROR((ReduceKernelShared<T, T, MIOPEN_REDUCE_TENSOR_NO_INDICES>(
|
||||
temp_db_data,
|
||||
dy_shape,
|
||||
db_data,
|
||||
TensorShape({}),
|
||||
MIOPEN_REDUCE_TENSOR_ADD,
|
||||
b_output_dims);
|
||||
b_output_dims)));
|
||||
}
|
||||
break;
|
||||
}
|
||||
|
|
@ -170,13 +170,13 @@ Status DivGrad<T>::ComputeInternal(OpKernelContext* context) const {
|
|||
|
||||
if (db_output_tensor) {
|
||||
std::vector<int64_t> b_output_dims = prepended_dimension_1(b_shape, dy_shape.NumDimensions());
|
||||
ReduceKernelShared<T, T, MIOPEN_REDUCE_TENSOR_NO_INDICES>(
|
||||
ORT_RETURN_IF_ERROR((ReduceKernelShared<T, T, MIOPEN_REDUCE_TENSOR_NO_INDICES>(
|
||||
temp_db_data,
|
||||
dy_shape,
|
||||
db_data,
|
||||
b_shape,
|
||||
MIOPEN_REDUCE_TENSOR_ADD,
|
||||
b_output_dims);
|
||||
b_output_dims)));
|
||||
}
|
||||
break;
|
||||
}
|
||||
|
|
@ -217,24 +217,24 @@ Status DivGrad<T>::ComputeInternal(OpKernelContext* context) const {
|
|||
|
||||
if (need_reduce_da) {
|
||||
std::vector<int64_t> a_output_dims = prepended_dimension_1(a_shape, dy_shape.NumDimensions());
|
||||
ReduceKernelShared<T, T, MIOPEN_REDUCE_TENSOR_NO_INDICES>(
|
||||
ORT_RETURN_IF_ERROR((ReduceKernelShared<T, T, MIOPEN_REDUCE_TENSOR_NO_INDICES>(
|
||||
da_data_ref,
|
||||
dy_shape,
|
||||
da_data,
|
||||
a_shape,
|
||||
MIOPEN_REDUCE_TENSOR_ADD,
|
||||
a_output_dims);
|
||||
a_output_dims)));
|
||||
}
|
||||
|
||||
if (need_reduce_db) {
|
||||
std::vector<int64_t> b_output_dims = prepended_dimension_1(b_shape, dy_shape.NumDimensions());
|
||||
ReduceKernelShared<T, T, MIOPEN_REDUCE_TENSOR_NO_INDICES>(
|
||||
ORT_RETURN_IF_ERROR((ReduceKernelShared<T, T, MIOPEN_REDUCE_TENSOR_NO_INDICES>(
|
||||
db_data_ref,
|
||||
dy_shape,
|
||||
db_data,
|
||||
b_shape,
|
||||
MIOPEN_REDUCE_TENSOR_ADD,
|
||||
b_output_dims);
|
||||
b_output_dims)));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -69,13 +69,13 @@ Status SoftMaxGradComputeHelper(
|
|||
T, \
|
||||
kRocmExecutionProvider, \
|
||||
(*KernelDefBuilder::Create()).TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
SoftmaxGrad<T>); \
|
||||
ONNX_OPERATOR_TYPED_KERNEL_EX( \
|
||||
LogSoftmaxGrad, \
|
||||
kMSDomain, \
|
||||
1, \
|
||||
T, \
|
||||
kRocmExecutionProvider, \
|
||||
SoftmaxGrad<T>); \
|
||||
ONNX_OPERATOR_TYPED_KERNEL_EX( \
|
||||
LogSoftmaxGrad, \
|
||||
kMSDomain, \
|
||||
1, \
|
||||
T, \
|
||||
kRocmExecutionProvider, \
|
||||
(*KernelDefBuilder::Create()).TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
SoftmaxGrad<T>);
|
||||
|
||||
|
|
|
|||
|
|
@ -17,8 +17,9 @@
|
|||
/* Modifications Copyright (c) Microsoft. */
|
||||
|
||||
// The code below is mostly copied from Pytorch PersistentSoftmax.cuh
|
||||
#include "hip/hip_runtime.h"
|
||||
|
||||
#include "orttraining/training_ops/rocm/math/softmax_grad.h"
|
||||
|
||||
#include "core/providers/rocm/cu_inc/common.cuh"
|
||||
#include "core/providers/rocm/math/softmax_impl.cuh"
|
||||
|
||||
|
|
@ -191,4 +192,4 @@ SPECIALIZED_SOFTMAX_GRAD_IMPL(half, half, float)
|
|||
SPECIALIZED_SOFTMAX_GRAD_IMPL(double, double, double)
|
||||
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include "batch_norm_grad.h"
|
||||
#include "orttraining/training_ops/rocm/nn/batch_norm_grad.h"
|
||||
#include "core/providers/common.h"
|
||||
#include "core/providers/rocm/miopen_common.h"
|
||||
#include "core/providers/cpu/nn/batch_norm_helper.h"
|
||||
|
|
@ -59,6 +59,7 @@ Status BatchNormalizationGrad<T, T1, T2>::ComputeInternal(OpKernelContext* ctx)
|
|||
vector<int64_t> new_dims;
|
||||
BatchNormHelper::NormalizeDims(input_shape, new_dims);
|
||||
ORT_RETURN_IF_ERROR(input_tensor.Set(new_dims, MiopenTensor::GetDataType<HipT>()));
|
||||
// for fp16 input, `scale_bias_tensor` will have a float type; otherwise it will be the same as input type.
|
||||
ORT_RETURN_IF_ERROR(scale_bias_tensor.Set(input_tensor, miopen_batch_norm_mode_));
|
||||
|
||||
const int64_t C = new_dims[1];
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@
|
|||
#pragma once
|
||||
|
||||
#include "gsl/gsl"
|
||||
|
||||
#include "core/providers/rocm/rocm_kernel.h"
|
||||
#include "core/providers/rocm/miopen_common.h"
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include "batch_norm_internal.h"
|
||||
#include "orttraining/training_ops/rocm/nn/batch_norm_internal.h"
|
||||
#include "core/providers/common.h"
|
||||
#include "core/providers/rocm/miopen_common.h"
|
||||
#include "core/providers/cpu/nn/batch_norm_helper.h"
|
||||
|
|
@ -67,6 +67,7 @@ Status BatchNormInternal<T, T1, T2>::ComputeInternal(OpKernelContext* p_op_kerne
|
|||
vector<int64_t> new_dims;
|
||||
BatchNormHelper::NormalizeDims(x_shape, new_dims);
|
||||
ORT_RETURN_IF_ERROR(data_desc.Set(new_dims, MiopenTensor::GetDataType<HipT>()));
|
||||
// for fp16 input, `bn_tensor_desc` will have a float type; otherwise it will be the same as input type.
|
||||
ORT_RETURN_IF_ERROR(bn_tensor_desc.Set(data_desc, miopen_batch_norm_mode_));
|
||||
|
||||
auto running_mean_data = reinterpret_cast<HipT2*>(running_mean->template MutableData<T2>());
|
||||
|
|
@ -155,7 +156,6 @@ Status BatchNormInternal<T, T1, T2>::ComputeInternal(OpKernelContext* p_op_kerne
|
|||
REGISTER_KERNEL_TYPED(T, T1, T2) \
|
||||
template Status BatchNormInternal<T, T1, T2>::ComputeInternal(OpKernelContext* ctx) const;
|
||||
|
||||
|
||||
SPECIALIZED_COMPUTE(float, float, float)
|
||||
// MIOpen kernel does not support double, disable for now.
|
||||
// SPECIALIZED_COMPUTE(double, double, double)
|
||||
|
|
|
|||
|
|
@ -45,16 +45,17 @@ Status ReduceAllL2<TIn, TOut>::ComputeInternal(OpKernelContext* ctx) const {
|
|||
HipTOut* p_output = reinterpret_cast<HipTOut*>(output->template MutableData<TOut>());
|
||||
HIP_RETURN_IF_ERROR(hipMemsetAsync(p_output, 0, sizeof(HipTOut), Stream()));
|
||||
|
||||
// bool deterministic = ctx->GetUseDeterministicCompute();
|
||||
// const bool deterministic = ctx->GetUseDeterministicCompute();
|
||||
bool deterministic = true;
|
||||
|
||||
if (!deterministic) {
|
||||
typedef MultiTensorReduceL2<HipTIn, HipTOut> TFunctor;
|
||||
TFunctor functor;
|
||||
|
||||
// Check if all values are finite and write true to deviceOutput.
|
||||
// Otherwise, false will be written.
|
||||
launch_multi_tensor_functor<1, TFunctor>(
|
||||
Stream(), 2048 * 32, tensor_sizes, grouped_tensor_pointers, functor, p_output);
|
||||
launch_multi_tensor_functor<1, TFunctor>(Stream(),
|
||||
2048 * 32, tensor_sizes, grouped_tensor_pointers, functor, p_output);
|
||||
|
||||
// *p_output is the squared sum of all elements.
|
||||
// Let's take a sqrt to get the actual L2-norm.
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ namespace rocm {
|
|||
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, View);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, Group);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, PassThrough);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, SGDOptimizer);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, ReduceSumTraining);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, double, ReduceSumTraining);
|
||||
|
|
@ -52,10 +53,10 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1
|
|||
class ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 9, float, int64_t, SparseSoftmaxCrossEntropy);
|
||||
// class ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 9, float, int32_t, SparseSoftmaxCrossEntropyGrad);
|
||||
class ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 9, float, int64_t, SparseSoftmaxCrossEntropyGrad);
|
||||
class ONNX_OPERATOR_VERSIONED_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 12, 12, float, int64_t, SoftmaxCrossEntropyLoss);
|
||||
class ONNX_OPERATOR_VERSIONED_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 12, 12, MLFloat16, int64_t, SoftmaxCrossEntropyLoss);
|
||||
class ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, float, int64_t, SoftmaxCrossEntropyLoss);
|
||||
class ONNX_OPERATOR_VERSIONED_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 12, 12, float, int64_t, SoftmaxCrossEntropyLoss);
|
||||
class ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, MLFloat16, int64_t, SoftmaxCrossEntropyLoss);
|
||||
class ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, float, int64_t, SoftmaxCrossEntropyLoss);
|
||||
class ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, int64_t, SoftmaxCrossEntropyLossGrad);
|
||||
class ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16, int64_t, SoftmaxCrossEntropyLossGrad);
|
||||
class ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, int64_t, SoftmaxCrossEntropyLossInternal);
|
||||
|
|
@ -73,11 +74,9 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1
|
|||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_MLFloat16, BatchNormalizationGrad);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_float, BatchNormalizationGrad);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16_float_float, BatchNormalizationGrad);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float_float_float, BatchNormInternal);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, double_double_double, BatchNormInternal);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_MLFloat16, BatchNormInternal);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_float, BatchNormInternal);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16_float_float, BatchNormInternal);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, ConvGrad);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, double, ConvGrad);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16, ConvGrad);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, GatherGrad);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, DropoutGrad);
|
||||
|
||||
|
|
@ -134,6 +133,29 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, Gath
|
|||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16, Scale);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, Scale);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, double, Scale);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistBinarizeEncoder);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16, GistBinarizeEncoder);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, double, GistBinarizeEncoder);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistBinarizeDecoder);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16, GistBinarizeDecoder);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, double, GistBinarizeDecoder);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, bool, GistPack1Encoder);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistPack1Encoder);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, bool, GistPack1Decoder);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistPack1Decoder);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistPack8Encoder);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16, GistPack8Encoder);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistPack8Decoder);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16, GistPack8Decoder);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistPack16Encoder);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistPack16Decoder);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistPackMsfp15Encoder);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistPackMsfp15Decoder);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float_float_float, BatchNormInternal);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, double_double_double, BatchNormInternal);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_MLFloat16, BatchNormInternal);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_float, BatchNormInternal);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16_float_float, BatchNormInternal);
|
||||
|
||||
#if defined(ORT_USE_NCCL) || defined(USE_MPI)
|
||||
// P2P communication operators.
|
||||
|
|
@ -141,11 +163,20 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, Send
|
|||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, Recv);
|
||||
#endif
|
||||
|
||||
#ifdef USE_MPI
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, AdasumAllReduce);
|
||||
#endif
|
||||
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, RecordEvent);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, WaitEvent);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, YieldOp);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, ATenOp);
|
||||
|
||||
#ifdef ENABLE_TRAINING_TORCH_INTEROP
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, PythonOp);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, PythonOpGrad);
|
||||
#endif
|
||||
|
||||
#ifdef ORT_USE_NCCL
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, NcclAllReduce);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, NcclAllGather);
|
||||
|
|
@ -158,6 +189,7 @@ Status RegisterRocmTrainingKernels(KernelRegistry& kernel_registry) {
|
|||
static const BuildKernelCreateInfoFn function_table[] = {
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, View)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, Group)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, PassThrough)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, SGDOptimizer)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, ReduceSumTraining)>,
|
||||
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, double, ReduceSumTraining)>,
|
||||
|
|
@ -211,10 +243,10 @@ Status RegisterRocmTrainingKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, LogSoftmaxGrad)>,
|
||||
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, double, LogSoftmaxGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16, LogSoftmaxGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 12, 12, float, int64_t, SoftmaxCrossEntropyLoss)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 12, 12, MLFloat16, int64_t, SoftmaxCrossEntropyLoss)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, float, int64_t, SoftmaxCrossEntropyLoss)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 12, 12, float, int64_t, SoftmaxCrossEntropyLoss)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, MLFloat16, int64_t, SoftmaxCrossEntropyLoss)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, float, int64_t, SoftmaxCrossEntropyLoss)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, int64_t, SoftmaxCrossEntropyLossGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16, int64_t, SoftmaxCrossEntropyLossGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TWO_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, int64_t, SoftmaxCrossEntropyLossInternal)>,
|
||||
|
|
@ -226,11 +258,9 @@ Status RegisterRocmTrainingKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_MLFloat16, BatchNormalizationGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_float, BatchNormalizationGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16_float_float, BatchNormalizationGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float_float_float, BatchNormInternal)>,
|
||||
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, double_double_double, BatchNormInternal)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_MLFloat16, BatchNormInternal)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_float, BatchNormInternal)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16_float_float, BatchNormInternal)>,
|
||||
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, ConvGrad)>,
|
||||
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, double, ConvGrad)>,
|
||||
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16, ConvGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, GatherGrad)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, DivGrad)>,
|
||||
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, double, DivGrad)>,
|
||||
|
|
@ -276,18 +306,50 @@ Status RegisterRocmTrainingKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16, Scale)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, Scale)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, double, Scale)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistBinarizeEncoder)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16, GistBinarizeEncoder)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, double, GistBinarizeEncoder)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistBinarizeDecoder)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16, GistBinarizeDecoder)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, double, GistBinarizeDecoder)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, bool, GistPack1Encoder)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistPack1Encoder)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, bool, GistPack1Decoder)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistPack1Decoder)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistPack8Encoder)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16, GistPack8Encoder)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistPack8Decoder)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16, GistPack8Decoder)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistPack16Encoder)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistPack16Decoder)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistPackMsfp15Encoder)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float, GistPackMsfp15Decoder)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, float_float_float, BatchNormInternal)>,
|
||||
// BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, double_double_double, BatchNormInternal)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_MLFloat16, BatchNormInternal)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16_MLFloat16_float, BatchNormInternal)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, MLFloat16_float_float, BatchNormInternal)>,
|
||||
|
||||
// P2P communication operators.
|
||||
#if defined(ORT_USE_NCCL) || defined(USE_MPI)
|
||||
// BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, Send)>,
|
||||
// BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, Recv)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, Send)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, Recv)>,
|
||||
#endif
|
||||
|
||||
// BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, RecordEvent)>,
|
||||
// BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, WaitEvent)>,
|
||||
#ifdef USE_MPI
|
||||
// BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, AdasumAllReduce)>,
|
||||
#endif
|
||||
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, RecordEvent)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, WaitEvent)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, YieldOp)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, ATenOp)>,
|
||||
|
||||
#ifdef ENABLE_TRAINING_TORCH_INTEROP
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, PythonOp)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, PythonOpGrad)>,
|
||||
#endif
|
||||
|
||||
#ifdef ORT_USE_NCCL
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, NcclAllReduce)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, NcclAllGather)>,
|
||||
|
|
|
|||
|
|
@ -1,46 +0,0 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include "orttraining/training_ops/rocm/tensor/gather_nd_grad_impl.h"
|
||||
|
||||
#include "core/providers/rocm/cu_inc/common.cuh"
|
||||
#include "core/providers/rocm/atomic/common.cuh"
|
||||
|
||||
namespace onnxruntime {
|
||||
namespace rocm {
|
||||
|
||||
template <typename T>
|
||||
__global__ void _GatherNDGradKernel(
|
||||
const size_t num_slices,
|
||||
const T* update_data,
|
||||
T* output_data,
|
||||
const size_t slice_size,
|
||||
const int64_t* slice_offsets) {
|
||||
CALCULATE_ELEMENTWISE_INDEX_OR_EXIT(i, num_slices * slice_size);
|
||||
uint64_t slice_offset = slice_offsets[i / slice_size];
|
||||
size_t j = i % slice_size;
|
||||
atomic_add(output_data + slice_offset + j, update_data[i]);
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
void GatherNDGradImpl(
|
||||
hipStream_t stream,
|
||||
const size_t num_slices,
|
||||
const void* update_data,
|
||||
void* output_data,
|
||||
const size_t slice_size,
|
||||
const int64_t* input_slice_offsets_data) {
|
||||
const auto blocks_per_grid = CeilDiv(num_slices * slice_size, GridDim::maxThreadsPerBlock);
|
||||
hipLaunchKernelGGL(HIP_KERNEL_NAME(_GatherNDGradKernel<T>), dim3(blocks_per_grid), dim3(GridDim::maxThreadsPerBlock), 0, stream,
|
||||
num_slices, static_cast<const T*>(update_data), static_cast<T*>(output_data), slice_size, input_slice_offsets_data);
|
||||
}
|
||||
|
||||
#define SPECIALIZED_GRAD_IMPL(T) \
|
||||
template void GatherNDGradImpl<T>(hipStream_t stream, const size_t num_slices, const void* update_data, void* output_data, const size_t slice_size, const int64_t* input_slice_offsets_data)
|
||||
|
||||
SPECIALIZED_GRAD_IMPL(float);
|
||||
SPECIALIZED_GRAD_IMPL(half);
|
||||
SPECIALIZED_GRAD_IMPL(double);
|
||||
|
||||
} // namespace rocm
|
||||
} // namespace onnxruntime
|
||||
|
|
@ -159,28 +159,20 @@ provider_excluded_files = [
|
|||
]
|
||||
|
||||
training_ops_excluded_files = [
|
||||
'activation/gelu_grad_impl_common.cuh',
|
||||
'activation/gelu_grad_impl_common.cuh', # uses custom tanh
|
||||
'collective/adasum_kernels.cc',
|
||||
'collective/adasum_kernels.h',
|
||||
'collective/nccl_common.cc',
|
||||
'collective/ready_event.cc',
|
||||
'collective/ready_event.h',
|
||||
'controlflow/record.cc',
|
||||
'controlflow/record.h',
|
||||
'controlflow/wait.cc',
|
||||
'controlflow/wait.h',
|
||||
'math/div_grad.cc',
|
||||
'math/softmax_grad_impl.cu',
|
||||
'math/softmax_grad.cc',
|
||||
'nn/batch_norm_grad.cc',
|
||||
'nn/batch_norm_grad.h',
|
||||
'nn/batch_norm_internal.cc',
|
||||
'nn/batch_norm_internal.h',
|
||||
'math/div_grad.cc', # miopen API differs from cudnn, no double type support
|
||||
'math/softmax_grad_impl.cu', # warp size differences
|
||||
'math/softmax_grad.cc', # miopen API differs from cudnn, no double type support
|
||||
'nn/batch_norm_grad.cc', # no double type support
|
||||
'nn/batch_norm_grad.h', # miopen API differs from cudnn
|
||||
'nn/batch_norm_internal.cc', # miopen API differs from cudnn, no double type support
|
||||
'nn/batch_norm_internal.h', # miopen API differs from cudnn, no double type support
|
||||
'nn/conv_grad.cc',
|
||||
'nn/conv_grad.h',
|
||||
'reduction/reduction_all.cc',
|
||||
'reduction/reduction_ops.cc',
|
||||
'tensor/gather_nd_grad_impl.cu',
|
||||
'reduction/reduction_all.cc', # deterministic = true, ignore ctx setting
|
||||
'reduction/reduction_ops.cc', # no double type support
|
||||
'cuda_training_kernels.cc',
|
||||
'cuda_training_kernels.h',
|
||||
]
|
||||
|
|
@ -306,6 +298,8 @@ def hipify(src_file_path, dst_file_path):
|
|||
s = s.replace('RegisterHipTrainingKernels', 'RegisterRocmTrainingKernels')
|
||||
s = s.replace('ROCM_VERSION', 'CUDA_VERSION') # semantically different meanings, cannot hipify
|
||||
s = s.replace('__ROCM_ARCH__', '__CUDA_ARCH__') # semantically different meanings, cannot hipify
|
||||
# "std::log" above incorrectly changed "std::logic_error" to "logfic_error"
|
||||
s = s.replace('logfic_error', 'std::logic_error')
|
||||
|
||||
# Deletions
|
||||
s = s.replace('#include "device_atomic_functions.h"', '') # HIP atomics in main hip header already
|
||||
|
|
@ -361,3 +355,8 @@ def amd_hipify(config_build_dir):
|
|||
log.debug(result.result())
|
||||
for result in training_results:
|
||||
log.debug(result.result())
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
import sys
|
||||
amd_hipify(sys.argv[1])
|
||||
|
|
|
|||
Loading…
Reference in a new issue