diff --git a/onnxruntime/core/providers/rocm/math/binary_elementwise_ops.cc b/onnxruntime/core/providers/rocm/math/binary_elementwise_ops.cc deleted file mode 100644 index f84d1e65dc..0000000000 --- a/onnxruntime/core/providers/rocm/math/binary_elementwise_ops.cc +++ /dev/null @@ -1,470 +0,0 @@ -// Copyright (c) Microsoft Corporation. All rights reserved. -// Licensed under the MIT License. - -#include "core/providers/rocm/math/binary_elementwise_ops.h" -#include "core/providers/rocm/math/binary_elementwise_ops_impl.h" -#include "core/providers/rocm/math/unary_elementwise_ops_impl.h" - -using namespace onnxruntime::common; -namespace onnxruntime { -namespace rocm { - -template <> -Status BinaryElementwise::Prepare(OpKernelContext* context, BinaryElementwisePreparation* p) const { - p->lhs_tensor = context->Input(0); - p->rhs_tensor = context->Input(1); - if (!(p->lhs_tensor->Shape() == p->rhs_tensor->Shape())) - return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, Node().Name(), ": mismatching input shapes: ", - p->lhs_tensor->Shape().ToString(), " != ", p->rhs_tensor->Shape().ToString()); - p->output_tensor = context->Output(0, p->lhs_tensor->Shape()); - p->output_rank_or_simple_broadcast = static_cast(SimpleBroadcast::NoBroadcast); - return Status::OK(); -} - -Status ComputeOutputShape(const std::string& node_name, const TensorShape& lhs_shape, const TensorShape& rhs_shape, TensorShape& out_shape) { - size_t lhs_rank = lhs_shape.NumDimensions(); - size_t rhs_rank = rhs_shape.NumDimensions(); - size_t out_rank = std::max(lhs_rank, rhs_rank); - - std::vector output_dims(out_rank, 0); - for (size_t i = 0; i < out_rank; ++i) { - int64_t lhs_dim = 1; - if (i < lhs_rank) - lhs_dim = lhs_shape[lhs_rank - 1 - i]; - int64_t rhs_dim = 1; - if (i < rhs_rank) - rhs_dim = rhs_shape[rhs_rank - 1 - i]; - int64_t max = std::max(lhs_dim, rhs_dim); - int64_t min = std::min(lhs_dim, rhs_dim); - int64_t out_dim = (min == 0 ? min : max); // special case a dim value of 0. - if (lhs_dim != out_dim && lhs_dim != 1) - return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, node_name, ": left operand cannot broadcast on dim ", lhs_rank - 1 - i, - " LeftShape: ", lhs_shape.ToString(), ", RightShape: ", rhs_shape.ToString()); - if (rhs_dim != out_dim && rhs_dim != 1) - return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, node_name, ": right operand cannot broadcast on dim ", rhs_rank - 1 - i, - " LeftShape: ", lhs_shape.ToString(), ", RightShape: ", rhs_shape.ToString()); - output_dims[out_rank - 1 - i] = out_dim; - } - out_shape = TensorShape(output_dims); - return Status::OK(); -} - -Status BinaryElementwiseBroadcastPrepare( - const Tensor* lhs_tensor, - const Tensor* rhs_tensor, - Tensor* output_tensor, - BinaryElementwisePreparation* p, - const TensorShape* override_lhs_shape, - const TensorShape* override_rhs_shape) { - p->lhs_tensor = lhs_tensor; - p->rhs_tensor = rhs_tensor; - const auto& lhs_shape = override_lhs_shape ? *override_lhs_shape : lhs_tensor->Shape(); - const auto& rhs_shape = override_rhs_shape ? *override_rhs_shape : rhs_tensor->Shape(); - - p->output_tensor = output_tensor; - const auto& output_shape = output_tensor->Shape(); - - ORT_RETURN_IF_ERROR(p->BinaryElementwiseBroadcastPrepareHelper(lhs_shape, rhs_shape, output_shape)); - - return Status::OK(); -} - -template <> -Status BinaryElementwise::Prepare(OpKernelContext* context, BinaryElementwisePreparation* p) const { - auto lhs_tensor = context->Input(0); - auto rhs_tensor = context->Input(1); - const auto& lhs_shape = lhs_tensor->Shape(); - const auto& rhs_shape = rhs_tensor->Shape(); - - TensorShape output_shape; - ORT_RETURN_IF_ERROR(ComputeOutputShape(Node().Name(), lhs_shape, rhs_shape, output_shape)); - auto output_tensor = context->Output(0, output_shape); - - ORT_RETURN_IF_ERROR(BinaryElementwiseBroadcastPrepare(lhs_tensor, rhs_tensor, output_tensor, p)); - - return Status::OK(); -} - -#define BINARY_ELEMENTWISE_REGISTER_KERNEL_TYPED_V(x, class_name, ver, T) \ - ONNX_OPERATOR_TYPED_KERNEL_EX( \ - x, \ - kOnnxDomain, \ - ver, \ - T, \ - kRocmExecutionProvider, \ - KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), \ - class_name); - -#define BINARY_ELEMENTWISE_REGISTER_KERNEL_TYPED(x, ver, T) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_TYPED_V(x, x, ver, T) - -#define BINARY_ELEMENTWISE_REGISTER_KERNEL_NONTEMP(x, class_name, ver, ...) \ - ONNX_OPERATOR_TYPED_KERNEL_EX( \ - x, \ - kOnnxDomain, \ - ver, \ - kRocmExecutionProvider, \ - KernelDefBuilder().TypeConstraint("T", BuildKernelDefConstraints<>(__VAR_ARGS__)), \ - class_name); - -#define BINARY_ELEMENTWISE_LOGICALOP_REGISTER_KERNEL_TYPED(x, ver, T) \ - ONNX_OPERATOR_TYPED_KERNEL_EX( \ - x, \ - kOnnxDomain, \ - ver, \ - T, \ - kRocmExecutionProvider, \ - KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()).TypeConstraint("T1", DataTypeImpl::GetTensorType()), \ - x); - -#define BINARY_ELEMENTWISE_REGISTER_KERNEL_VERSIONED_TYPED(x, startver, endver, T) \ - ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_EX( \ - x, \ - kOnnxDomain, \ - startver, \ - endver, \ - T, \ - kRocmExecutionProvider, \ - KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), \ - x); - -#define BINARY_ELEMENTWISE_REGISTER_KERNEL_VERSIONED_TYPED_CLASS(x, class_name, startver, endver, T) \ - ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_EX( \ - x, \ - kOnnxDomain, \ - startver, \ - endver, \ - T, \ - kRocmExecutionProvider, \ - KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType()), \ - class_name); - -#define BINARY_ELEMENTWISE_COMPUTE(x, T) \ - template <> \ - Status x::ComputeInternal(OpKernelContext* context) const { \ - BinaryElementwisePreparation prepare; \ - ORT_RETURN_IF_ERROR(Prepare(context, &prepare)); \ - Impl_##x::MappedType>( \ - prepare.output_rank_or_simple_broadcast, \ - &prepare.lhs_padded_strides, \ - reinterpret_cast::MappedType*>(prepare.lhs_tensor->template Data()), \ - &prepare.rhs_padded_strides, \ - reinterpret_cast::MappedType*>(prepare.rhs_tensor->template Data()), \ - &prepare.fdm_output_strides, \ - prepare.fdm_H, \ - prepare.fdm_C, \ - reinterpret_cast::MappedType*>(prepare.output_tensor->template MutableData()), \ - prepare.output_tensor->Shape().Size()); \ - return Status::OK(); \ - } - -#define BINARY_OP_VERSIONED_TYPED(name, startver, endver, T) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_VERSIONED_TYPED(name, startver, endver, T) - -#define BINARY_OP_TYPED(name, ver, T) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_TYPED(name, ver, T) \ - BINARY_ELEMENTWISE_COMPUTE(name, T) - -#define BINARY_OP_TYPED_VERSIONED_V(name, class_name, startver, endver, T) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_VERSIONED_TYPED_CLASS(name, class_name, startver, endver, T) \ - BINARY_ELEMENTWISE_COMPUTE(class_name, T) - -#define BINARY_LOGICALOP_TYPED(name, ver, T) \ - BINARY_ELEMENTWISE_LOGICALOP_REGISTER_KERNEL_TYPED(name, ver, T) \ - BINARY_ELEMENTWISE_COMPUTE(name, T) - -// since different ops has different types, we cannot use BINARY_OPS() directly -// the postfix of means the types supported by the op: -// B: uint8_t -// W: uint16_t -// U: uint32_t -// Z: uint64_t -// C: int8_t -// S: int16_t -// I: int32_t -// L: int64_t -// H: float16 -// F: float -// D: double -// O: bool - -#define BINARY_OP_VERSIONED_HFD(name, startver, endver) \ - BINARY_OP_VERSIONED_TYPED(name, startver, endver, MLFloat16) \ - BINARY_OP_VERSIONED_TYPED(name, startver, endver, float) \ - BINARY_OP_VERSIONED_TYPED(name, startver, endver, double) - -#define BINARY_OP_VERSIONED_UZILHFD(name, startver, endver) \ - BINARY_OP_VERSIONED_TYPED(name, startver, endver, uint32_t) \ - BINARY_OP_VERSIONED_TYPED(name, startver, endver, uint64_t) \ - BINARY_OP_VERSIONED_TYPED(name, startver, endver, int32_t) \ - BINARY_OP_VERSIONED_TYPED(name, startver, endver, int64_t) \ - BINARY_OP_VERSIONED_HFD(name, startver, endver) - -#define BINARY_OP_HFD(name, ver) \ - BINARY_OP_TYPED(name, ver, MLFloat16) \ - BINARY_OP_TYPED(name, ver, float) \ - BINARY_OP_TYPED(name, ver, double) - -#define BINARY_OP_UZILHFD(name, ver) \ - BINARY_OP_TYPED(name, ver, uint32_t) \ - BINARY_OP_TYPED(name, ver, uint64_t) \ - BINARY_OP_TYPED(name, ver, int32_t) \ - BINARY_OP_TYPED(name, ver, int64_t) \ - BINARY_OP_HFD(name, ver) - -#define BINARY_OP_REGISTER_OIL(name, ver) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_TYPED(name, ver, bool) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_TYPED(name, ver, int32_t) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_TYPED(name, ver, int64_t) - -#define BINARY_OP_REGISTER_VERSIONED_OIL(name, startver, endver) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_VERSIONED_TYPED(name, startver, endver, bool) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_VERSIONED_TYPED(name, startver, endver, int32_t) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_VERSIONED_TYPED(name, startver, endver, int64_t) - -#define BINARY_LOGICALOP_REGISTER_OIL(name, ver) \ - BINARY_ELEMENTWISE_LOGICALOP_REGISTER_KERNEL_TYPED(name, ver, bool) \ - BINARY_ELEMENTWISE_LOGICALOP_REGISTER_KERNEL_TYPED(name, ver, int32_t) \ - BINARY_ELEMENTWISE_LOGICALOP_REGISTER_KERNEL_TYPED(name, ver, int64_t) - -#define BINARY_OP_REGISTER_HFD(name, ver) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_TYPED(name, ver, MLFloat16) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_TYPED(name, ver, float) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_TYPED(name, ver, double) - -#define BINARY_OP_REGISTER_UZILHFD(name, ver) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_TYPED(name, ver, uint32_t) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_TYPED(name, ver, uint64_t) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_TYPED(name, ver, int32_t) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_TYPED(name, ver, int64_t) \ - BINARY_OP_REGISTER_HFD(name, ver) - -#define BINARY_LOGICALOP_REGISTER_UZILHFD(name, ver) \ - BINARY_ELEMENTWISE_LOGICALOP_REGISTER_KERNEL_TYPED(name, ver, uint32_t) \ - BINARY_ELEMENTWISE_LOGICALOP_REGISTER_KERNEL_TYPED(name, ver, uint64_t) \ - BINARY_ELEMENTWISE_LOGICALOP_REGISTER_KERNEL_TYPED(name, ver, int32_t) \ - BINARY_ELEMENTWISE_LOGICALOP_REGISTER_KERNEL_TYPED(name, ver, int64_t) \ - BINARY_ELEMENTWISE_LOGICALOP_REGISTER_KERNEL_TYPED(name, ver, MLFloat16) \ - BINARY_ELEMENTWISE_LOGICALOP_REGISTER_KERNEL_TYPED(name, ver, float) \ - BINARY_ELEMENTWISE_LOGICALOP_REGISTER_KERNEL_TYPED(name, ver, double) - -#define BINARY_OP_REGISTER_VERSIONED_HFD(name, startver, endver) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_VERSIONED_TYPED(name, startver, endver, MLFloat16) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_VERSIONED_TYPED(name, startver, endver, float) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_VERSIONED_TYPED(name, startver, endver, double) - -#define BINARY_OP_REGISTER_VERSIONED_CLASS_HFD(name, class_name, startver, endver) \ - BINARY_OP_TYPED_VERSIONED_V(name, class_name, startver, endver, MLFloat16) \ - BINARY_OP_TYPED_VERSIONED_V(name, class_name, startver, endver, float) \ - BINARY_OP_TYPED_VERSIONED_V(name, class_name, startver, endver, double) - -#define BINARY_OP_REGISTER_VERSIONED_UZILHFD(name, startver, endver) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_VERSIONED_TYPED(name, startver, endver, uint32_t) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_VERSIONED_TYPED(name, startver, endver, uint64_t) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_VERSIONED_TYPED(name, startver, endver, int32_t) \ - BINARY_ELEMENTWISE_REGISTER_KERNEL_VERSIONED_TYPED(name, startver, endver, int64_t) \ - BINARY_OP_REGISTER_VERSIONED_HFD(name, startver, endver) - -BINARY_OP_VERSIONED_UZILHFD(Add, 7, 12) -BINARY_OP_VERSIONED_UZILHFD(Sub, 7, 12) -BINARY_OP_VERSIONED_UZILHFD(Mul, 7, 12) -BINARY_OP_VERSIONED_UZILHFD(Div, 7, 12) - -BINARY_OP_UZILHFD(Add, 13) -BINARY_OP_UZILHFD(Sub, 13) -BINARY_OP_UZILHFD(Mul, 13) -BINARY_OP_UZILHFD(Div, 13) - -BINARY_OP_REGISTER_VERSIONED_CLASS_HFD(Pow, Pow_7, 7, 11) -BINARY_LOGICALOP_TYPED(And, 7, bool) -BINARY_LOGICALOP_TYPED(Or, 7, bool) -BINARY_LOGICALOP_TYPED(Xor, 7, bool) -BINARY_OP_VERSIONED_HFD(PRelu, 7, 8) -BINARY_OP_HFD(PRelu, 9) - -// Pow since version 12 -ONNX_OPERATOR_VERSIONED_KERNEL_EX( - Pow, - kOnnxDomain, - 12, 12, - kRocmExecutionProvider, - KernelDefBuilder().TypeConstraint("T", BuildKernelDefConstraints()).TypeConstraint("T1", BuildKernelDefConstraints()), - Pow); - -ONNX_OPERATOR_KERNEL_EX( - Pow, - kOnnxDomain, - 13, - kRocmExecutionProvider, - KernelDefBuilder().TypeConstraint("T", BuildKernelDefConstraints()).TypeConstraint("T1", BuildKernelDefConstraints()), - Pow); - -namespace pow12_internal { -template -Status DispatchOnFirstArg(const BinaryElementwisePreparation& prepare) { - namespace on = ONNX_NAMESPACE; - Status s; - switch (prepare.rhs_tensor->GetElementType()) { - case on::TensorProto_DataType_INT32: - ImplT1_Pow::MappedType, typename ToHipType::MappedType>( - prepare.output_rank_or_simple_broadcast, - &prepare.lhs_padded_strides, - reinterpret_cast::MappedType*>(prepare.lhs_tensor->template Data()), - &prepare.rhs_padded_strides, - reinterpret_cast::MappedType*>(prepare.rhs_tensor->template Data()), - &prepare.fdm_output_strides, - prepare.fdm_H, - prepare.fdm_C, - reinterpret_cast::MappedType*>(prepare.output_tensor->template MutableData()), - prepare.output_tensor->Shape().Size()); - break; - case on::TensorProto_DataType_INT64: - ImplT1_Pow::MappedType, typename ToHipType::MappedType>( - prepare.output_rank_or_simple_broadcast, - &prepare.lhs_padded_strides, - reinterpret_cast::MappedType*>(prepare.lhs_tensor->template Data()), - &prepare.rhs_padded_strides, - reinterpret_cast::MappedType*>(prepare.rhs_tensor->template Data()), - &prepare.fdm_output_strides, - prepare.fdm_H, - prepare.fdm_C, - reinterpret_cast::MappedType*>(prepare.output_tensor->template MutableData()), - prepare.output_tensor->Shape().Size()); - break; - case on::TensorProto_DataType_FLOAT: - ImplT1_Pow::MappedType, typename ToHipType::MappedType>( - prepare.output_rank_or_simple_broadcast, - &prepare.lhs_padded_strides, - reinterpret_cast::MappedType*>(prepare.lhs_tensor->template Data()), - &prepare.rhs_padded_strides, - reinterpret_cast::MappedType*>(prepare.rhs_tensor->template Data()), - &prepare.fdm_output_strides, - prepare.fdm_H, - prepare.fdm_C, - reinterpret_cast::MappedType*>(prepare.output_tensor->template MutableData()), - prepare.output_tensor->Shape().Size()); - break; - case on::TensorProto_DataType_DOUBLE: - ImplT1_Pow::MappedType, typename ToHipType::MappedType>( - prepare.output_rank_or_simple_broadcast, - &prepare.lhs_padded_strides, - reinterpret_cast::MappedType*>(prepare.lhs_tensor->template Data()), - &prepare.rhs_padded_strides, - reinterpret_cast::MappedType*>(prepare.rhs_tensor->template Data()), - &prepare.fdm_output_strides, - prepare.fdm_H, - prepare.fdm_C, - reinterpret_cast::MappedType*>(prepare.output_tensor->template MutableData()), - prepare.output_tensor->Shape().Size()); - break; - case on::TensorProto_DataType_FLOAT16: - ImplT1_Pow::MappedType, typename ToHipType::MappedType>( - prepare.output_rank_or_simple_broadcast, - &prepare.lhs_padded_strides, - reinterpret_cast::MappedType*>(prepare.lhs_tensor->template Data()), - &prepare.rhs_padded_strides, - reinterpret_cast::MappedType*>(prepare.rhs_tensor->template Data()), - &prepare.fdm_output_strides, - prepare.fdm_H, - prepare.fdm_C, - reinterpret_cast::MappedType*>(prepare.output_tensor->template MutableData()), - prepare.output_tensor->Shape().Size()); - break; - default: - s = ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Unsupported Y type: ", - DataTypeImpl::ToString(prepare.rhs_tensor->DataType())); - } - return s; -} -} // namespace pow12_internal - -Status Pow::ComputeInternal(OpKernelContext* context) const { - BinaryElementwisePreparation prepare; - ORT_RETURN_IF_ERROR(Prepare(context, &prepare)); - namespace on = ONNX_NAMESPACE; - using namespace pow12_internal; - - Status s; - - switch (prepare.lhs_tensor->GetElementType()) { - case on::TensorProto_DataType_INT32: - s = DispatchOnFirstArg(prepare); - break; - case on::TensorProto_DataType_INT64: - s = DispatchOnFirstArg(prepare); - break; - case on::TensorProto_DataType_FLOAT: - s = DispatchOnFirstArg(prepare); - break; - case on::TensorProto_DataType_DOUBLE: - s = DispatchOnFirstArg(prepare); - break; - case on::TensorProto_DataType_FLOAT16: - s = DispatchOnFirstArg(prepare); - break; - default: - s = ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Unsupported X type: ", - DataTypeImpl::ToString(prepare.lhs_tensor->DataType())); - } - return s; -} - -//Greater op output tensor type is bool, so it cannot directly fit in the macros -//for other elementwise ops -template -Status CompareFunction::CompareMethod(OpKernelContext* context, ImplCompare Impl_Compare) const { - BinaryElementwisePreparation prepare; - ORT_RETURN_IF_ERROR(Prepare(context, &prepare)); - - Impl_Compare( - prepare.output_rank_or_simple_broadcast, - &prepare.lhs_padded_strides, - reinterpret_cast(prepare.lhs_tensor->template Data()), - &prepare.rhs_padded_strides, - reinterpret_cast(prepare.rhs_tensor->template Data()), - &prepare.fdm_output_strides, - prepare.fdm_H, - prepare.fdm_C, - reinterpret_cast::MappedType*>(prepare.output_tensor->template MutableData()), - prepare.output_tensor->Shape().Size()); - - return Status::OK(); -} - -//Greater op output tensor type is bool, so it cannot directly fit in the macros -//for other elementwise ops -template -Status Greater::ComputeInternal(OpKernelContext* context) const { - this->CompareMethod(context, &ImplT2_Greater); - - return Status::OK(); -} - -template -Status Equal::ComputeInternal(OpKernelContext* context) const { - this->CompareMethod(context, &ImplT2_Equal); - - return Status::OK(); -} - -//Less op output tensor type is bool, so it cannot directly fit in the macros -//for other elementwise ops -template -Status Less::ComputeInternal(OpKernelContext* context) const { - this->CompareMethod(context, &ImplT2_Less); - - return Status::OK(); -} - -BINARY_OP_REGISTER_OIL(Equal, 13) -BINARY_OP_REGISTER_VERSIONED_OIL(Equal, 11, 12) -BINARY_OP_REGISTER_VERSIONED_OIL(Equal, 7, 10) -BINARY_LOGICALOP_REGISTER_UZILHFD(Greater, 13) -BINARY_OP_REGISTER_VERSIONED_UZILHFD(Greater, 9, 12) -BINARY_OP_REGISTER_VERSIONED_HFD(Greater, 7, 8) -BINARY_LOGICALOP_REGISTER_UZILHFD(Less, 13) -BINARY_OP_REGISTER_VERSIONED_UZILHFD(Less, 9, 12) -BINARY_OP_REGISTER_VERSIONED_HFD(Less, 7, 8) - -} // namespace rocm -} // namespace onnxruntime diff --git a/onnxruntime/core/providers/rocm/math/binary_elementwise_ops_impl.cu b/onnxruntime/core/providers/rocm/math/binary_elementwise_ops_impl.cu deleted file mode 100644 index ec305b04d0..0000000000 --- a/onnxruntime/core/providers/rocm/math/binary_elementwise_ops_impl.cu +++ /dev/null @@ -1,171 +0,0 @@ -// Copyright (c) Microsoft Corporation. All rights reserved. -// Licensed under the MIT License. - -#include -#include "core/providers/rocm/math/binary_elementwise_ops_impl.h" -#include "core/providers/rocm/cu_inc/common.cuh" -#include "core/providers/rocm/cu_inc/binary_elementwise_impl.cuh" -#include "core/providers/rocm/math/binary_elementwise_ops_impl_functors.cuh" - -namespace onnxruntime { -namespace rocm { - -#define BINARY_ELEMENTWISE_IMPL(name) \ - BINARY_ELEMENTWISE_IMPL_DECLARATION(name) { \ - BinaryElementWiseImpl(output_rank_or_simple_broadcast, \ - lhs_padded_strides, \ - lhs_data, \ - rhs_padded_strides, \ - rhs_data, \ - fdm_output_strides, \ - fdm_H, \ - fdm_C, \ - output_data, \ - OP_##name(), \ - count); \ - } - -#define BINARY_ELEMENTWISE_IMPL_T1(name) \ - BINARY_ELEMENTWISE_IMPL_DECLARATION_T1(name) { \ - BinaryElementWiseImpl(output_rank_or_simple_broadcast, \ - lhs_padded_strides, \ - lhs_data, \ - rhs_padded_strides, \ - rhs_data, \ - fdm_output_strides, \ - fdm_H, \ - fdm_C, \ - output_data, \ - OP_##name(), \ - count); \ - } - -#define BINARY_ELEMENTWISE_IMPL_T2(name) \ - BINARY_ELEMENTWISE_IMPL_DECLARATION_T2(name) { \ - BinaryElementWiseImpl(output_rank_or_simple_broadcast, \ - lhs_padded_strides, \ - lhs_data, \ - rhs_padded_strides, \ - rhs_data, \ - fdm_output_strides, \ - fdm_H, \ - fdm_C, \ - output_data, \ - OP_##name(), \ - count); \ - } - -#define SPECIALIZED_BINARY_ELEMENTWISE_IMPL(x, T) \ - template void Impl_##x(int32_t output_rank, \ - const TArray* lhs_padded_strides, const T* lhs_data, \ - const TArray* rhs_padded_strides, const T* rhs_data, \ - const TArray* fdm_output_strides, const fast_divmod& fdm_H, const fast_divmod& fdm_C, T* output_data, size_t count); - -#define SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(x, T, T1) \ - template void ImplT1_##x(int32_t output_rank, \ - const TArray* lhs_padded_strides, const T* lhs_data, \ - const TArray* rhs_padded_strides, const T1* rhs_data, \ - const TArray* fdm_output_strides, const fast_divmod& fdm_H, const fast_divmod& fdm_C, T* output_data, size_t count); - -#define SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T2(x, T, T1, T2) \ - template void ImplT2_##x(int32_t output_rank, \ - const TArray* lhs_padded_strides, const T1* lhs_data, \ - const TArray* rhs_padded_strides, const T2* rhs_data, \ - const TArray* fdm_output_strides, const fast_divmod& fdm_H, const fast_divmod& fdm_C, T* output_data, size_t count); - -#define SPECIALIZED_BINARY_ELEMENTWISE_IMPL_UZILHFD(x) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL(x, uint32_t) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL(x, uint64_t) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL(x, int32_t) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL(x, int64_t) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL(x, half) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL(x, float) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL(x, double) - -#define SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1_ILHFD(x, T) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(x, T, int32_t) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(x, T, int64_t) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(x, T, half) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(x, T, float) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(x, T, double) - -#define SPECIALIZED_BINARY_ELEMENTWISE_IMPL_OIL(x) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL(x, bool) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL(x, int32_t) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL(x, int64_t) - -#define SPECIALIZED_BINARY_ELEMENTWISE_IMPL_HFD(x) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL(x, half) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL(x, float) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL(x, double) - -// create declarations for impl -#define BINARY_OP_NAME_EXPR(name, expr) \ - BINARY_ELEMENTWISE_IMPL(name) - -BINARY_OPS() -#undef BINARY_OP_NAME_EXPR - -// create specialized impl -// the postfix of means the types supported by the op: -// B: uint8_t -// W: uint16_t -// U: uint32_t -// Z: uint64_t -// C: int8_t -// S: int16_t -// I: int32_t -// L: int64_t -// H: float16 -// F: float -// D: double -// O: bool - -SPECIALIZED_BINARY_ELEMENTWISE_IMPL_UZILHFD(Add) -SPECIALIZED_BINARY_ELEMENTWISE_IMPL(Add, bool) -SPECIALIZED_BINARY_ELEMENTWISE_IMPL_UZILHFD(Sub) -SPECIALIZED_BINARY_ELEMENTWISE_IMPL_UZILHFD(Mul) -SPECIALIZED_BINARY_ELEMENTWISE_IMPL_UZILHFD(Div) -SPECIALIZED_BINARY_ELEMENTWISE_IMPL_HFD(Pow_7) -SPECIALIZED_BINARY_ELEMENTWISE_IMPL(And, bool) -SPECIALIZED_BINARY_ELEMENTWISE_IMPL(Or, bool) -SPECIALIZED_BINARY_ELEMENTWISE_IMPL(Xor, bool) -SPECIALIZED_BINARY_ELEMENTWISE_IMPL_HFD(PRelu) -SPECIALIZED_BINARY_ELEMENTWISE_IMPL_UZILHFD(Max) -SPECIALIZED_BINARY_ELEMENTWISE_IMPL_UZILHFD(Min) - -// create declarations for impl for Pow -BINARY_ELEMENTWISE_IMPL_T1(Pow) - -SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1_ILHFD(Pow, int32_t) -SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1_ILHFD(Pow, int64_t) -SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1_ILHFD(Pow, float) -SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1_ILHFD(Pow, double) -SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1_ILHFD(Pow, half) - -// create declarations for impl2 -#define BINARY_OP_NAME_EXPR2(name, expr) \ - BINARY_ELEMENTWISE_IMPL_T2(name) - -BINARY_OPS2() -#undef BINARY_OP_NAME_EXPR2 - -#define SPECIALIZED_BINARY_ELEMENTWISE_IMPL_UZILHFD2(name) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T2(name, bool, uint32_t, uint32_t) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T2(name, bool, uint64_t, uint64_t) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T2(name, bool, int32_t, int32_t) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T2(name, bool, int64_t, int64_t) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T2(name, bool, half, half) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T2(name, bool, float, float) \ - SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T2(name, bool, double, double) - -SPECIALIZED_BINARY_ELEMENTWISE_IMPL_UZILHFD2(Greater) - -SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T2(Equal, bool, bool, bool) -SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T2(Equal, bool, int32_t, int32_t) -SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T2(Equal, bool, int64_t, int64_t) - -SPECIALIZED_BINARY_ELEMENTWISE_IMPL_UZILHFD2(Less) - -} // namespace rocm -} // namespace onnxruntime diff --git a/tools/ci_build/amd_hipify.py b/tools/ci_build/amd_hipify.py index 40afd3a65c..4db31e18d1 100644 --- a/tools/ci_build/amd_hipify.py +++ b/tools/ci_build/amd_hipify.py @@ -84,8 +84,6 @@ provider_excluded_files = [ 'math/einsum_utils/einsum_auxiliary_ops.h', 'math/einsum_utils/einsum_auxiliary_ops_diagonal.cu', 'math/einsum_utils/einsum_auxiliary_ops_diagonal.h', - 'math/binary_elementwise_ops.cc', - 'math/binary_elementwise_ops_impl.cu', 'math/cumsum.cc', 'math/cumsum.h', 'math/cumsum_impl.cu',