mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-23 19:32:23 +00:00
Add workaround to remove ROCm-specific binary-elementwise files.
This commit is contained in:
parent
1059bfaf75
commit
fa851bff66
3 changed files with 0 additions and 643 deletions
|
|
@ -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<ShouldNotBroadcast>::Prepare(OpKernelContext* context, BinaryElementwisePreparation* p) const {
|
||||
p->lhs_tensor = context->Input<Tensor>(0);
|
||||
p->rhs_tensor = context->Input<Tensor>(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<int32_t>(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<int64_t> 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<ShouldBroadcast>::Prepare(OpKernelContext* context, BinaryElementwisePreparation* p) const {
|
||||
auto lhs_tensor = context->Input<Tensor>(0);
|
||||
auto rhs_tensor = context->Input<Tensor>(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<T>()), \
|
||||
class_name<T>);
|
||||
|
||||
#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<T>()).TypeConstraint("T1", DataTypeImpl::GetTensorType<bool>()), \
|
||||
x<T>);
|
||||
|
||||
#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<T>()), \
|
||||
x<T>);
|
||||
|
||||
#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<T>()), \
|
||||
class_name<T>);
|
||||
|
||||
#define BINARY_ELEMENTWISE_COMPUTE(x, T) \
|
||||
template <> \
|
||||
Status x<T>::ComputeInternal(OpKernelContext* context) const { \
|
||||
BinaryElementwisePreparation prepare; \
|
||||
ORT_RETURN_IF_ERROR(Prepare(context, &prepare)); \
|
||||
Impl_##x<typename ToHipType<T>::MappedType>( \
|
||||
prepare.output_rank_or_simple_broadcast, \
|
||||
&prepare.lhs_padded_strides, \
|
||||
reinterpret_cast<const typename ToHipType<T>::MappedType*>(prepare.lhs_tensor->template Data<T>()), \
|
||||
&prepare.rhs_padded_strides, \
|
||||
reinterpret_cast<const typename ToHipType<T>::MappedType*>(prepare.rhs_tensor->template Data<T>()), \
|
||||
&prepare.fdm_output_strides, \
|
||||
prepare.fdm_H, \
|
||||
prepare.fdm_C, \
|
||||
reinterpret_cast<typename ToHipType<T>::MappedType*>(prepare.output_tensor->template MutableData<T>()), \
|
||||
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<int32_t, int64_t, float, double>()).TypeConstraint("T1", BuildKernelDefConstraints<int32_t, int64_t, float, double>()),
|
||||
Pow);
|
||||
|
||||
ONNX_OPERATOR_KERNEL_EX(
|
||||
Pow,
|
||||
kOnnxDomain,
|
||||
13,
|
||||
kRocmExecutionProvider,
|
||||
KernelDefBuilder().TypeConstraint("T", BuildKernelDefConstraints<int32_t, int64_t, float, double, MLFloat16>()).TypeConstraint("T1", BuildKernelDefConstraints<int32_t, int64_t, float, double, MLFloat16>()),
|
||||
Pow);
|
||||
|
||||
namespace pow12_internal {
|
||||
template <class T>
|
||||
Status DispatchOnFirstArg(const BinaryElementwisePreparation& prepare) {
|
||||
namespace on = ONNX_NAMESPACE;
|
||||
Status s;
|
||||
switch (prepare.rhs_tensor->GetElementType()) {
|
||||
case on::TensorProto_DataType_INT32:
|
||||
ImplT1_Pow<typename ToHipType<T>::MappedType, typename ToHipType<int32_t>::MappedType>(
|
||||
prepare.output_rank_or_simple_broadcast,
|
||||
&prepare.lhs_padded_strides,
|
||||
reinterpret_cast<const typename ToHipType<T>::MappedType*>(prepare.lhs_tensor->template Data<T>()),
|
||||
&prepare.rhs_padded_strides,
|
||||
reinterpret_cast<const typename ToHipType<int32_t>::MappedType*>(prepare.rhs_tensor->template Data<int32_t>()),
|
||||
&prepare.fdm_output_strides,
|
||||
prepare.fdm_H,
|
||||
prepare.fdm_C,
|
||||
reinterpret_cast<typename ToHipType<T>::MappedType*>(prepare.output_tensor->template MutableData<T>()),
|
||||
prepare.output_tensor->Shape().Size());
|
||||
break;
|
||||
case on::TensorProto_DataType_INT64:
|
||||
ImplT1_Pow<typename ToHipType<T>::MappedType, typename ToHipType<int64_t>::MappedType>(
|
||||
prepare.output_rank_or_simple_broadcast,
|
||||
&prepare.lhs_padded_strides,
|
||||
reinterpret_cast<const typename ToHipType<T>::MappedType*>(prepare.lhs_tensor->template Data<T>()),
|
||||
&prepare.rhs_padded_strides,
|
||||
reinterpret_cast<const typename ToHipType<int64_t>::MappedType*>(prepare.rhs_tensor->template Data<int64_t>()),
|
||||
&prepare.fdm_output_strides,
|
||||
prepare.fdm_H,
|
||||
prepare.fdm_C,
|
||||
reinterpret_cast<typename ToHipType<T>::MappedType*>(prepare.output_tensor->template MutableData<T>()),
|
||||
prepare.output_tensor->Shape().Size());
|
||||
break;
|
||||
case on::TensorProto_DataType_FLOAT:
|
||||
ImplT1_Pow<typename ToHipType<T>::MappedType, typename ToHipType<float>::MappedType>(
|
||||
prepare.output_rank_or_simple_broadcast,
|
||||
&prepare.lhs_padded_strides,
|
||||
reinterpret_cast<const typename ToHipType<T>::MappedType*>(prepare.lhs_tensor->template Data<T>()),
|
||||
&prepare.rhs_padded_strides,
|
||||
reinterpret_cast<const typename ToHipType<float>::MappedType*>(prepare.rhs_tensor->template Data<float>()),
|
||||
&prepare.fdm_output_strides,
|
||||
prepare.fdm_H,
|
||||
prepare.fdm_C,
|
||||
reinterpret_cast<typename ToHipType<T>::MappedType*>(prepare.output_tensor->template MutableData<T>()),
|
||||
prepare.output_tensor->Shape().Size());
|
||||
break;
|
||||
case on::TensorProto_DataType_DOUBLE:
|
||||
ImplT1_Pow<typename ToHipType<T>::MappedType, typename ToHipType<double>::MappedType>(
|
||||
prepare.output_rank_or_simple_broadcast,
|
||||
&prepare.lhs_padded_strides,
|
||||
reinterpret_cast<const typename ToHipType<T>::MappedType*>(prepare.lhs_tensor->template Data<T>()),
|
||||
&prepare.rhs_padded_strides,
|
||||
reinterpret_cast<const typename ToHipType<double>::MappedType*>(prepare.rhs_tensor->template Data<double>()),
|
||||
&prepare.fdm_output_strides,
|
||||
prepare.fdm_H,
|
||||
prepare.fdm_C,
|
||||
reinterpret_cast<typename ToHipType<T>::MappedType*>(prepare.output_tensor->template MutableData<T>()),
|
||||
prepare.output_tensor->Shape().Size());
|
||||
break;
|
||||
case on::TensorProto_DataType_FLOAT16:
|
||||
ImplT1_Pow<typename ToHipType<T>::MappedType, typename ToHipType<MLFloat16>::MappedType>(
|
||||
prepare.output_rank_or_simple_broadcast,
|
||||
&prepare.lhs_padded_strides,
|
||||
reinterpret_cast<const typename ToHipType<T>::MappedType*>(prepare.lhs_tensor->template Data<T>()),
|
||||
&prepare.rhs_padded_strides,
|
||||
reinterpret_cast<const typename ToHipType<MLFloat16>::MappedType*>(prepare.rhs_tensor->template Data<MLFloat16>()),
|
||||
&prepare.fdm_output_strides,
|
||||
prepare.fdm_H,
|
||||
prepare.fdm_C,
|
||||
reinterpret_cast<typename ToHipType<T>::MappedType*>(prepare.output_tensor->template MutableData<T>()),
|
||||
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<int32_t>(prepare);
|
||||
break;
|
||||
case on::TensorProto_DataType_INT64:
|
||||
s = DispatchOnFirstArg<int64_t>(prepare);
|
||||
break;
|
||||
case on::TensorProto_DataType_FLOAT:
|
||||
s = DispatchOnFirstArg<float>(prepare);
|
||||
break;
|
||||
case on::TensorProto_DataType_DOUBLE:
|
||||
s = DispatchOnFirstArg<double>(prepare);
|
||||
break;
|
||||
case on::TensorProto_DataType_FLOAT16:
|
||||
s = DispatchOnFirstArg<MLFloat16>(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 <typename T, typename HipT>
|
||||
Status CompareFunction<T, HipT>::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<const HipT*>(prepare.lhs_tensor->template Data<T>()),
|
||||
&prepare.rhs_padded_strides,
|
||||
reinterpret_cast<const HipT*>(prepare.rhs_tensor->template Data<T>()),
|
||||
&prepare.fdm_output_strides,
|
||||
prepare.fdm_H,
|
||||
prepare.fdm_C,
|
||||
reinterpret_cast<ToHipType<bool>::MappedType*>(prepare.output_tensor->template MutableData<bool>()),
|
||||
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 <typename T>
|
||||
Status Greater<T>::ComputeInternal(OpKernelContext* context) const {
|
||||
this->CompareMethod(context, &ImplT2_Greater);
|
||||
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
Status Equal<T>::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 <typename T>
|
||||
Status Less<T>::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
|
||||
|
|
@ -1,171 +0,0 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include <hip/hip_runtime.h>
|
||||
#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<T, T, T>(), \
|
||||
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<T, T, T1>(), \
|
||||
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<T, T1, T2>(), \
|
||||
count); \
|
||||
}
|
||||
|
||||
#define SPECIALIZED_BINARY_ELEMENTWISE_IMPL(x, T) \
|
||||
template void Impl_##x<T>(int32_t output_rank, \
|
||||
const TArray<int64_t>* lhs_padded_strides, const T* lhs_data, \
|
||||
const TArray<int64_t>* rhs_padded_strides, const T* rhs_data, \
|
||||
const TArray<fast_divmod>* 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<T, T1>(int32_t output_rank, \
|
||||
const TArray<int64_t>* lhs_padded_strides, const T* lhs_data, \
|
||||
const TArray<int64_t>* rhs_padded_strides, const T1* rhs_data, \
|
||||
const TArray<fast_divmod>* 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<T, T1, T2>(int32_t output_rank, \
|
||||
const TArray<int64_t>* lhs_padded_strides, const T1* lhs_data, \
|
||||
const TArray<int64_t>* rhs_padded_strides, const T2* rhs_data, \
|
||||
const TArray<fast_divmod>* 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
|
||||
|
|
@ -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',
|
||||
|
|
|
|||
Loading…
Reference in a new issue