From 66ceb6926df6c7a81dc1f241c6bf59cfe58d3a72 Mon Sep 17 00:00:00 2001 From: Jeff Daily Date: Thu, 21 Oct 2021 13:36:21 -0700 Subject: [PATCH] rehipify ROCm EP files under orttraining (#9443) * rehipify rocm ep files under orttraining committed to source control * fix flake8 error --- .../rocm/collective/nccl_common.cc | 136 ------------------ .../training_ops/rocm/math/div_grad.cc | 20 +-- .../training_ops/rocm/math/softmax_grad.cc | 14 +- .../rocm/math/softmax_grad_impl.cu | 5 +- .../training_ops/rocm/nn/batch_norm_grad.cc | 3 +- .../training_ops/rocm/nn/batch_norm_grad.h | 1 + .../rocm/nn/batch_norm_internal.cc | 4 +- .../rocm/reduction/reduction_all.cc | 7 +- .../rocm/rocm_training_kernels.cc | 98 ++++++++++--- .../rocm/tensor/gather_nd_grad_impl.cu | 46 ------ tools/ci_build/amd_hipify.py | 35 +++-- 11 files changed, 126 insertions(+), 243 deletions(-) delete mode 100644 orttraining/orttraining/training_ops/rocm/collective/nccl_common.cc delete mode 100644 orttraining/orttraining/training_ops/rocm/tensor/gather_nd_grad_impl.cu diff --git a/orttraining/orttraining/training_ops/rocm/collective/nccl_common.cc b/orttraining/orttraining/training_ops/rocm/collective/nccl_common.cc deleted file mode 100644 index 8a12674c43..0000000000 --- a/orttraining/orttraining/training_ops/rocm/collective/nccl_common.cc +++ /dev/null @@ -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 - -#include "orttraining/core/framework/communication/mpi/mpi_context.h" - -namespace onnxruntime { -namespace rocm { - -ncclDataType_t GetNcclDataType(onnxruntime::MLDataType type) { - if (type == DataTypeImpl::GetType()) { - return ncclUint8; - } else if (type == DataTypeImpl::GetType()) { - return ncclInt8; - } else if (type == DataTypeImpl::GetType()) { - return ncclInt32; - } else if (type == DataTypeImpl::GetType()) { - return ncclInt64; - } else if (type == DataTypeImpl::GetType()) { - return ncclFloat16; - } else if (type == DataTypeImpl::GetType()) { - return ncclFloat32; - } else if (type == DataTypeImpl::GetType()) { - 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(0)); - group_type_ = static_cast(group_type); -} - -} // namespace rocm -} // namespace onnxruntime diff --git a/orttraining/orttraining/training_ops/rocm/math/div_grad.cc b/orttraining/orttraining/training_ops/rocm/math/div_grad.cc index ef1e1b93dc..16a21403db 100644 --- a/orttraining/orttraining/training_ops/rocm/math/div_grad.cc +++ b/orttraining/orttraining/training_ops/rocm/math/div_grad.cc @@ -96,13 +96,13 @@ Status DivGrad::ComputeInternal(OpKernelContext* context) const { if (da_output_tensor) { std::vector a_output_dims = prepended_dimension_1(a_shape, dy_shape.NumDimensions()); - ReduceKernelShared( + ORT_RETURN_IF_ERROR((ReduceKernelShared( 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::ComputeInternal(OpKernelContext* context) const { if (db_output_tensor) { std::vector b_output_dims = prepended_dimension_1(b_shape, dy_shape.NumDimensions()); - ReduceKernelShared( + ORT_RETURN_IF_ERROR((ReduceKernelShared( 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::ComputeInternal(OpKernelContext* context) const { if (db_output_tensor) { std::vector b_output_dims = prepended_dimension_1(b_shape, dy_shape.NumDimensions()); - ReduceKernelShared( + ORT_RETURN_IF_ERROR((ReduceKernelShared( 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::ComputeInternal(OpKernelContext* context) const { if (need_reduce_da) { std::vector a_output_dims = prepended_dimension_1(a_shape, dy_shape.NumDimensions()); - ReduceKernelShared( + ORT_RETURN_IF_ERROR((ReduceKernelShared( 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 b_output_dims = prepended_dimension_1(b_shape, dy_shape.NumDimensions()); - ReduceKernelShared( + ORT_RETURN_IF_ERROR((ReduceKernelShared( db_data_ref, dy_shape, db_data, b_shape, MIOPEN_REDUCE_TENSOR_ADD, - b_output_dims); + b_output_dims))); } } } diff --git a/orttraining/orttraining/training_ops/rocm/math/softmax_grad.cc b/orttraining/orttraining/training_ops/rocm/math/softmax_grad.cc index 6dfb5ceaad..052dc86670 100644 --- a/orttraining/orttraining/training_ops/rocm/math/softmax_grad.cc +++ b/orttraining/orttraining/training_ops/rocm/math/softmax_grad.cc @@ -69,13 +69,13 @@ Status SoftMaxGradComputeHelper( T, \ kRocmExecutionProvider, \ (*KernelDefBuilder::Create()).TypeConstraint("T", DataTypeImpl::GetTensorType()), \ - SoftmaxGrad); \ - ONNX_OPERATOR_TYPED_KERNEL_EX( \ - LogSoftmaxGrad, \ - kMSDomain, \ - 1, \ - T, \ - kRocmExecutionProvider, \ + SoftmaxGrad); \ + ONNX_OPERATOR_TYPED_KERNEL_EX( \ + LogSoftmaxGrad, \ + kMSDomain, \ + 1, \ + T, \ + kRocmExecutionProvider, \ (*KernelDefBuilder::Create()).TypeConstraint("T", DataTypeImpl::GetTensorType()), \ SoftmaxGrad); diff --git a/orttraining/orttraining/training_ops/rocm/math/softmax_grad_impl.cu b/orttraining/orttraining/training_ops/rocm/math/softmax_grad_impl.cu index 2781435170..a68c3c6df9 100644 --- a/orttraining/orttraining/training_ops/rocm/math/softmax_grad_impl.cu +++ b/orttraining/orttraining/training_ops/rocm/math/softmax_grad_impl.cu @@ -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) } -} \ No newline at end of file +} diff --git a/orttraining/orttraining/training_ops/rocm/nn/batch_norm_grad.cc b/orttraining/orttraining/training_ops/rocm/nn/batch_norm_grad.cc index 0a65545e3a..e2b6827141 100644 --- a/orttraining/orttraining/training_ops/rocm/nn/batch_norm_grad.cc +++ b/orttraining/orttraining/training_ops/rocm/nn/batch_norm_grad.cc @@ -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::ComputeInternal(OpKernelContext* ctx) vector new_dims; BatchNormHelper::NormalizeDims(input_shape, new_dims); ORT_RETURN_IF_ERROR(input_tensor.Set(new_dims, MiopenTensor::GetDataType())); + // 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]; diff --git a/orttraining/orttraining/training_ops/rocm/nn/batch_norm_grad.h b/orttraining/orttraining/training_ops/rocm/nn/batch_norm_grad.h index 4995692b1d..c94f73881b 100644 --- a/orttraining/orttraining/training_ops/rocm/nn/batch_norm_grad.h +++ b/orttraining/orttraining/training_ops/rocm/nn/batch_norm_grad.h @@ -4,6 +4,7 @@ #pragma once #include "gsl/gsl" + #include "core/providers/rocm/rocm_kernel.h" #include "core/providers/rocm/miopen_common.h" diff --git a/orttraining/orttraining/training_ops/rocm/nn/batch_norm_internal.cc b/orttraining/orttraining/training_ops/rocm/nn/batch_norm_internal.cc index 3c5c8f35d9..f45348f096 100644 --- a/orttraining/orttraining/training_ops/rocm/nn/batch_norm_internal.cc +++ b/orttraining/orttraining/training_ops/rocm/nn/batch_norm_internal.cc @@ -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::ComputeInternal(OpKernelContext* p_op_kerne vector new_dims; BatchNormHelper::NormalizeDims(x_shape, new_dims); ORT_RETURN_IF_ERROR(data_desc.Set(new_dims, MiopenTensor::GetDataType())); + // 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(running_mean->template MutableData()); @@ -155,7 +156,6 @@ Status BatchNormInternal::ComputeInternal(OpKernelContext* p_op_kerne REGISTER_KERNEL_TYPED(T, T1, T2) \ template Status BatchNormInternal::ComputeInternal(OpKernelContext* ctx) const; - SPECIALIZED_COMPUTE(float, float, float) // MIOpen kernel does not support double, disable for now. // SPECIALIZED_COMPUTE(double, double, double) diff --git a/orttraining/orttraining/training_ops/rocm/reduction/reduction_all.cc b/orttraining/orttraining/training_ops/rocm/reduction/reduction_all.cc index 8ecaa96502..d1e7985504 100644 --- a/orttraining/orttraining/training_ops/rocm/reduction/reduction_all.cc +++ b/orttraining/orttraining/training_ops/rocm/reduction/reduction_all.cc @@ -45,16 +45,17 @@ Status ReduceAllL2::ComputeInternal(OpKernelContext* ctx) const { HipTOut* p_output = reinterpret_cast(output->template MutableData()); 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 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. diff --git a/orttraining/orttraining/training_ops/rocm/rocm_training_kernels.cc b/orttraining/orttraining/training_ops/rocm/rocm_training_kernels.cc index c6a3591d84..2fce0f5ef0 100644 --- a/orttraining/orttraining/training_ops/rocm/rocm_training_kernels.cc +++ b/orttraining/orttraining/training_ops/rocm/rocm_training_kernels.cc @@ -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, BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, // BuildKernelCreateInfo, @@ -211,10 +243,10 @@ Status RegisterRocmTrainingKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, // BuildKernelCreateInfo, BuildKernelCreateInfo, - BuildKernelCreateInfo, BuildKernelCreateInfo, - BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, @@ -226,11 +258,9 @@ Status RegisterRocmTrainingKernels(KernelRegistry& kernel_registry) { BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, - BuildKernelCreateInfo, - // BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, - BuildKernelCreateInfo, + // BuildKernelCreateInfo, + // BuildKernelCreateInfo, + // BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, // BuildKernelCreateInfo, @@ -276,18 +306,50 @@ Status RegisterRocmTrainingKernels(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, + // BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, // P2P communication operators. #if defined(ORT_USE_NCCL) || defined(USE_MPI) - // BuildKernelCreateInfo, - // BuildKernelCreateInfo, + BuildKernelCreateInfo, + BuildKernelCreateInfo, #endif - // BuildKernelCreateInfo, - // BuildKernelCreateInfo, +#ifdef USE_MPI + // BuildKernelCreateInfo, +#endif + + BuildKernelCreateInfo, + BuildKernelCreateInfo, BuildKernelCreateInfo, BuildKernelCreateInfo, +#ifdef ENABLE_TRAINING_TORCH_INTEROP + BuildKernelCreateInfo, + BuildKernelCreateInfo, +#endif + #ifdef ORT_USE_NCCL BuildKernelCreateInfo, BuildKernelCreateInfo, diff --git a/orttraining/orttraining/training_ops/rocm/tensor/gather_nd_grad_impl.cu b/orttraining/orttraining/training_ops/rocm/tensor/gather_nd_grad_impl.cu deleted file mode 100644 index 8f924df170..0000000000 --- a/orttraining/orttraining/training_ops/rocm/tensor/gather_nd_grad_impl.cu +++ /dev/null @@ -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 -__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 -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), dim3(blocks_per_grid), dim3(GridDim::maxThreadsPerBlock), 0, stream, - num_slices, static_cast(update_data), static_cast(output_data), slice_size, input_slice_offsets_data); -} - -#define SPECIALIZED_GRAD_IMPL(T) \ - template 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) - -SPECIALIZED_GRAD_IMPL(float); -SPECIALIZED_GRAD_IMPL(half); -SPECIALIZED_GRAD_IMPL(double); - -} // namespace rocm -} // namespace onnxruntime diff --git a/tools/ci_build/amd_hipify.py b/tools/ci_build/amd_hipify.py index bd88a243ed..66750e8ddf 100644 --- a/tools/ci_build/amd_hipify.py +++ b/tools/ci_build/amd_hipify.py @@ -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])