mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
Implement Pow(12) for cpu and cuda (#3727)
* Implement Pow(12) cpu and cuda.
This commit is contained in:
parent
027a364922
commit
f68a326bd9
12 changed files with 507 additions and 107 deletions
|
|
@ -81,8 +81,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain,
|
|||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, float, Reciprocal);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, float, Sqrt);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, double, Sqrt);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, float, Pow);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, double, Pow);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 11, Pow);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, float, Exp);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, double, Exp);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, float, Log);
|
||||
|
|
@ -423,6 +422,8 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, Mi
|
|||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, Max);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, MaxPool);
|
||||
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, Pow);
|
||||
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, float, ReduceMax);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, int32_t, ReduceMax);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, int64_t, ReduceMax);
|
||||
|
|
@ -513,8 +514,7 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, float, Sqrt)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, double,
|
||||
Sqrt)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, float, Pow)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, double, Pow)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 7, 11, Pow)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, float, Exp)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, double, Exp)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 6, float, Log)>,
|
||||
|
|
@ -1059,6 +1059,8 @@ Status RegisterOnnxOperatorKernels(KernelRegistry& kernel_registry) {
|
|||
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, Max)>,
|
||||
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, Pow)>,
|
||||
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, MaxPool)>,
|
||||
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCpuExecutionProvider, kOnnxDomain, 12, float,
|
||||
|
|
|
|||
|
|
@ -47,23 +47,23 @@ namespace onnxruntime {
|
|||
.TypeConstraint("T1", DataTypeImpl::GetTensorType<bool>()), \
|
||||
KERNEL_CLASS<TYPE>);
|
||||
|
||||
// var args are type constraints
|
||||
// var args are type constraints for T and T1
|
||||
#define REG_ELEMENTWISE_KERNEL_NONT(OP_TYPE, VERSION, KERNEL_CLASS, ...) \
|
||||
ONNX_CPU_OPERATOR_KERNEL( \
|
||||
OP_TYPE, \
|
||||
VERSION, \
|
||||
KernelDefBuilder() \
|
||||
ONNX_CPU_OPERATOR_KERNEL( \
|
||||
OP_TYPE, \
|
||||
VERSION, \
|
||||
KernelDefBuilder() \
|
||||
.TypeConstraint("T", BuildKernelDefConstraints<__VA_ARGS__>()) \
|
||||
.TypeConstraint("T1", BuildKernelDefConstraints<__VA_ARGS__>()), \
|
||||
KERNEL_CLASS);
|
||||
|
||||
// var args are type constraints
|
||||
// var args are type constraints for T and T1
|
||||
#define REG_ELEMENTWISE_VERSIONED_KERNEL_NONT(OP_TYPE, VERSION_FROM, VERSION_TO, KERNEL_CLASS, ...) \
|
||||
ONNX_CPU_OPERATOR_VERSIONED_KERNEL( \
|
||||
OP_TYPE, \
|
||||
VERSION_FROM, \
|
||||
VERSION_TO, \
|
||||
KernelDefBuilder() \
|
||||
ONNX_CPU_OPERATOR_VERSIONED_KERNEL( \
|
||||
OP_TYPE, \
|
||||
VERSION_FROM, \
|
||||
VERSION_TO, \
|
||||
KernelDefBuilder() \
|
||||
.TypeConstraint("T", BuildKernelDefConstraints<__VA_ARGS__>()) \
|
||||
.TypeConstraint("T1", BuildKernelDefConstraints<__VA_ARGS__>()), \
|
||||
KERNEL_CLASS);
|
||||
|
|
@ -114,8 +114,10 @@ REG_ELEMENTWISE_TYPED_KERNEL(Reciprocal, 6, float, Reciprocal);
|
|||
REG_ELEMENTWISE_TYPED_KERNEL(Sqrt, 6, float, Sqrt);
|
||||
REG_ELEMENTWISE_TYPED_KERNEL(Sqrt, 6, double, Sqrt);
|
||||
|
||||
REG_ELEMENTWISE_TYPED_KERNEL(Pow, 7, float, Pow);
|
||||
REG_ELEMENTWISE_TYPED_KERNEL(Pow, 7, double, Pow);
|
||||
REG_ELEMENTWISE_VERSIONED_KERNEL_NONT(Pow, 7, 11, Pow, float, double);
|
||||
// To reduce templetization we choose to support the below types for both
|
||||
// base and the exponent. This gives us 16 permutations
|
||||
REG_ELEMENTWISE_KERNEL_NONT(Pow, 12, Pow, int32_t, int64_t, float, double);
|
||||
|
||||
REG_ELEMENTWISE_TYPED_KERNEL(Exp, 6, float, Exp);
|
||||
REG_ELEMENTWISE_TYPED_KERNEL(Exp, 6, double, Exp);
|
||||
|
|
@ -273,25 +275,115 @@ Status Sqrt<T>::Compute(OpKernelContext* ctx) const {
|
|||
return Status::OK();
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
Status Pow<T>::Compute(OpKernelContext* context) const {
|
||||
const Tensor& Y = *context->Input<Tensor>(1);
|
||||
std::function<void(EigenVectorMap<T>, ConstEigenVectorMap<T>, T)> input1scalar =
|
||||
[](EigenVectorMap<T> output, ConstEigenVectorMap<T> input0, T input1) { output = Eigen::pow(input0.array(), input1); };
|
||||
namespace pow_internal {
|
||||
|
||||
template <typename T, typename E>
|
||||
void PowImpl(OpKernelContext* context, const Tensor& X, const Tensor& Y) {
|
||||
TBroadcaster<T, E> bc{X, Y};
|
||||
Tensor* const output_tensor = context->Output(0, bc.GetOutputShape());
|
||||
TBroadcastOutput<T> output{bc.GetSpanSize(), *output_tensor};
|
||||
|
||||
// Scalar base
|
||||
auto input0scalar = [](gsl::span<T> output, T X, gsl::span<const E> Y) {
|
||||
std::transform(Y.cbegin(), Y.cend(), output.begin(),
|
||||
[X](E y) {
|
||||
return static_cast<T>(std::pow(X, y));
|
||||
});
|
||||
};
|
||||
|
||||
// Scalar exponent switch to possibly available optimizations
|
||||
std::function<void(gsl::span<T>, gsl::span<const T> X, E Y)> input1scalar =
|
||||
[](gsl::span<T> output, gsl::span<const T> X, E Y) {
|
||||
std::transform(X.cbegin(), X.cend(), output.begin(),
|
||||
[Y](T x) {
|
||||
return static_cast<T>(std::pow(x, Y));
|
||||
});
|
||||
};
|
||||
|
||||
if (Y.Shape().Size() == 1) {
|
||||
T value = *Y.Data<T>();
|
||||
if (value == 2.0) {
|
||||
input1scalar = [](EigenVectorMap<T> output, ConstEigenVectorMap<T> input0, T) { output = Eigen::square(input0.array()); };
|
||||
} else if (value == 3.0) {
|
||||
input1scalar = [](EigenVectorMap<T> output, ConstEigenVectorMap<T> input0, T) { output = Eigen::cube(input0.array()); };
|
||||
auto exp = *Y.template Data<E>();
|
||||
if (exp == E{2}) {
|
||||
input1scalar = [](gsl::span<T> output, gsl::span<const T> X, E) {
|
||||
std::transform(X.cbegin(), X.cend(), output.begin(),
|
||||
[](T x) {
|
||||
return static_cast<T>(x * x);
|
||||
});
|
||||
};
|
||||
} else if (exp == E{3}) {
|
||||
input1scalar = [](gsl::span<T> output, gsl::span<const T> X, E) {
|
||||
std::transform(X.cbegin(), X.cend(), output.begin(),
|
||||
[](T x) {
|
||||
return static_cast<T>(x * x * x);
|
||||
});
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
return BroadcastTwo<T, T>(
|
||||
*context,
|
||||
[](EigenVectorMap<T> output, T input0, ConstEigenVectorMap<T> input1) { output = Eigen::pow(input0, input1.array()); },
|
||||
input1scalar,
|
||||
[](EigenVectorMap<T> output, ConstEigenVectorMap<T> input0, ConstEigenVectorMap<T> input1) { output = Eigen::pow(input0.array(), input1.array()); });
|
||||
auto general = [](gsl::span<T> output, gsl::span<const T> X, gsl::span<const E> Y) {
|
||||
std::transform(
|
||||
X.cbegin(), X.cend(), Y.cbegin(), output.begin(),
|
||||
[](T x, E y) {
|
||||
return static_cast<T>(std::pow(x, y));
|
||||
});
|
||||
};
|
||||
|
||||
BroadcastLoopSpan(bc, output, input0scalar, input1scalar, general);
|
||||
}
|
||||
|
||||
template <typename B>
|
||||
Status DispatchOnBase(OpKernelContext* context, const Tensor& X, const Tensor& Y) {
|
||||
namespace on = ONNX_NAMESPACE;
|
||||
Status s;
|
||||
switch (Y.GetElementType()) {
|
||||
case on::TensorProto_DataType_INT32:
|
||||
PowImpl<B, int32_t>(context, X, Y);
|
||||
break;
|
||||
case on::TensorProto_DataType_INT64:
|
||||
PowImpl<B, int64_t>(context, X, Y);
|
||||
break;
|
||||
case on::TensorProto_DataType_FLOAT:
|
||||
PowImpl<B, float>(context, X, Y);
|
||||
break;
|
||||
case on::TensorProto_DataType_DOUBLE:
|
||||
PowImpl<B, double>(context, X, Y);
|
||||
break;
|
||||
default:
|
||||
s = ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Unsupported Y type: ",
|
||||
DataTypeImpl::ToString(Y.DataType()));
|
||||
}
|
||||
return s;
|
||||
}
|
||||
|
||||
} // namespace pow_internal
|
||||
|
||||
Status
|
||||
Pow::Compute(OpKernelContext* context) const {
|
||||
const Tensor& X = *context->Input<Tensor>(0);
|
||||
const Tensor& Y = *context->Input<Tensor>(1);
|
||||
|
||||
namespace on = ONNX_NAMESPACE;
|
||||
using namespace pow_internal;
|
||||
|
||||
Status s;
|
||||
// Switch on base type first
|
||||
switch (X.GetElementType()) {
|
||||
case on::TensorProto_DataType_INT32:
|
||||
s = DispatchOnBase<int32_t>(context, X, Y);
|
||||
break;
|
||||
case on::TensorProto_DataType_INT64:
|
||||
s = DispatchOnBase<int64_t>(context, X, Y);
|
||||
break;
|
||||
case on::TensorProto_DataType_FLOAT:
|
||||
s = DispatchOnBase<float>(context, X, Y);
|
||||
break;
|
||||
case on::TensorProto_DataType_DOUBLE:
|
||||
s = DispatchOnBase<double>(context, X, Y);
|
||||
break;
|
||||
default:
|
||||
s = ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Unsupported X type: ",
|
||||
DataTypeImpl::ToString(X.DataType()));
|
||||
}
|
||||
return s;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
|
|
|
|||
|
|
@ -111,7 +111,6 @@ class Sqrt final : public OpKernel {
|
|||
Status Compute(OpKernelContext* context) const override;
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
class Pow final : public OpKernel {
|
||||
public:
|
||||
Pow(const OpKernelInfo& info) : OpKernel(info) {
|
||||
|
|
|
|||
|
|
@ -10,20 +10,20 @@ namespace onnxruntime {
|
|||
namespace cuda {
|
||||
|
||||
// broadcast by computing output coordinate from offset, using fast_divmod
|
||||
template <typename T, typename FuncT, bool lhs_need_compute, bool rhs_need_compute, int NumThreadsPerBlock, int NumElementsPerThread>
|
||||
template <typename T, typename T1, typename FuncT, bool lhs_need_compute, bool rhs_need_compute, int NumThreadsPerBlock, int NumElementsPerThread>
|
||||
__global__ void _BinaryElementWise(
|
||||
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 T1* rhs_data,
|
||||
const TArray<fast_divmod> fdm_output_strides,
|
||||
T* output_data,
|
||||
const FuncT& functor,
|
||||
CUDA_LONG N) {
|
||||
CUDA_LONG start = NumElementsPerThread * NumThreadsPerBlock * blockIdx.x + threadIdx.x;
|
||||
T lvalue[NumElementsPerThread];
|
||||
T rvalue[NumElementsPerThread];
|
||||
T1 rvalue[NumElementsPerThread];
|
||||
|
||||
CUDA_LONG id = start;
|
||||
#pragma unroll
|
||||
|
|
@ -68,16 +68,16 @@ __global__ void _BinaryElementWise(
|
|||
}
|
||||
|
||||
// for scalar broadcast or non-broadcast case
|
||||
template <bool IncL, bool IncR, typename T, typename FuncT, int NumThreadsPerBlock, int NumElementsPerThread>
|
||||
template <bool IncL, bool IncR, typename T, typename T1, typename FuncT, int NumThreadsPerBlock, int NumElementsPerThread>
|
||||
__global__ void _BinaryElementWiseSimple(
|
||||
const T* lhs_data,
|
||||
const T* rhs_data,
|
||||
const T1* rhs_data,
|
||||
T* output_data,
|
||||
const FuncT& func,
|
||||
CUDA_LONG N) {
|
||||
CUDA_LONG start = NumElementsPerThread * NumThreadsPerBlock * blockIdx.x + threadIdx.x;
|
||||
T lvalue[NumElementsPerThread];
|
||||
T rvalue[NumElementsPerThread];
|
||||
T1 rvalue[NumElementsPerThread];
|
||||
|
||||
CUDA_LONG id = start;
|
||||
#pragma unroll
|
||||
|
|
@ -102,17 +102,17 @@ __global__ void _BinaryElementWiseSimple(
|
|||
}
|
||||
|
||||
// for rhs per-channel broadcast case
|
||||
template <typename T, typename FuncT, int NumThreadsPerBlock, int NumElementsPerThread>
|
||||
template <typename T, typename T1, typename FuncT, int NumThreadsPerBlock, int NumElementsPerThread>
|
||||
__global__ void _BinaryElementWiseRhsPerChannelBatch1(
|
||||
const T* lhs_data,
|
||||
const T* rhs_data,
|
||||
const T1* rhs_data,
|
||||
const fast_divmod fdm_H,
|
||||
T* output_data,
|
||||
FuncT func,
|
||||
CUDA_LONG N) {
|
||||
CUDA_LONG start = NumElementsPerThread * NumThreadsPerBlock * blockIdx.x + threadIdx.x;
|
||||
T lvalue[NumElementsPerThread];
|
||||
T rvalue[NumElementsPerThread];
|
||||
T1 rvalue[NumElementsPerThread];
|
||||
|
||||
CUDA_LONG id = start;
|
||||
#pragma unroll
|
||||
|
|
@ -137,10 +137,10 @@ __global__ void _BinaryElementWiseRhsPerChannelBatch1(
|
|||
}
|
||||
}
|
||||
|
||||
template <typename T, typename FuncT, int NumThreadsPerBlock, int NumElementsPerThread>
|
||||
template <typename T, typename T1, typename FuncT, int NumThreadsPerBlock, int NumElementsPerThread>
|
||||
__global__ void _BinaryElementWiseRhsPerChannelBatchN(
|
||||
const T* lhs_data,
|
||||
const T* rhs_data,
|
||||
const T1* rhs_data,
|
||||
const fast_divmod fdm_H,
|
||||
const fast_divmod fdm_C,
|
||||
T* output_data,
|
||||
|
|
@ -148,7 +148,7 @@ __global__ void _BinaryElementWiseRhsPerChannelBatchN(
|
|||
CUDA_LONG N) {
|
||||
CUDA_LONG start = NumElementsPerThread * NumThreadsPerBlock * blockIdx.x + threadIdx.x;
|
||||
T lvalue[NumElementsPerThread];
|
||||
T rvalue[NumElementsPerThread];
|
||||
T1 rvalue[NumElementsPerThread];
|
||||
|
||||
CUDA_LONG id = start;
|
||||
#pragma unroll
|
||||
|
|
@ -177,10 +177,10 @@ __global__ void _BinaryElementWiseRhsPerChannelBatchN(
|
|||
}
|
||||
}
|
||||
|
||||
template <typename T, typename FuncT>
|
||||
template <typename T, typename T1, typename FuncT>
|
||||
void BinaryElementWiseNoBroadcastImpl(
|
||||
const T* lhs_data,
|
||||
const T* rhs_data,
|
||||
const T1* rhs_data,
|
||||
T* output_data,
|
||||
const FuncT& func,
|
||||
size_t count) {
|
||||
|
|
@ -189,7 +189,7 @@ void BinaryElementWiseNoBroadcastImpl(
|
|||
|
||||
int blocksPerGrid = static_cast<int>(CeilDiv(count, GridDim::maxThreadsPerBlock * GridDim::maxElementsPerThread));
|
||||
CUDA_LONG N = static_cast<CUDA_LONG>(count);
|
||||
_BinaryElementWiseSimple<true, true, T, FuncT, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
|
||||
_BinaryElementWiseSimple<true, true, T, T1, FuncT, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
|
||||
lhs_data,
|
||||
rhs_data,
|
||||
output_data,
|
||||
|
|
@ -197,13 +197,13 @@ void BinaryElementWiseNoBroadcastImpl(
|
|||
N);
|
||||
}
|
||||
|
||||
template <typename T, typename FuncT>
|
||||
template <typename T, typename T1, typename FuncT>
|
||||
void BinaryElementWiseImpl(
|
||||
int32_t output_rank_or_simple_broadcast,
|
||||
const TArray<int64_t>* lhs_padded_strides,
|
||||
const T* lhs_data,
|
||||
const TArray<int64_t>* rhs_padded_strides,
|
||||
const T* rhs_data,
|
||||
const T1* rhs_data,
|
||||
const TArray<fast_divmod>* fdm_output_strides,
|
||||
const fast_divmod& fdm_H,
|
||||
const fast_divmod& fdm_C,
|
||||
|
|
@ -216,21 +216,21 @@ void BinaryElementWiseImpl(
|
|||
int blocksPerGrid = static_cast<int>(CeilDiv(count, GridDim::maxThreadsPerBlock * GridDim::maxElementsPerThread));
|
||||
CUDA_LONG N = static_cast<CUDA_LONG>(count);
|
||||
if (output_rank_or_simple_broadcast == static_cast<int32_t>(SimpleBroadcast::NoBroadcast)) {
|
||||
_BinaryElementWiseSimple<true, true, T, FuncT, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
|
||||
_BinaryElementWiseSimple<true, true, T, T1, FuncT, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
|
||||
lhs_data,
|
||||
rhs_data,
|
||||
output_data,
|
||||
func,
|
||||
N);
|
||||
} else if (output_rank_or_simple_broadcast == static_cast<int32_t>(SimpleBroadcast::LeftScalar)) {
|
||||
_BinaryElementWiseSimple<false, true, T, FuncT, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
|
||||
_BinaryElementWiseSimple<false, true, T, T1, FuncT, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
|
||||
lhs_data,
|
||||
rhs_data,
|
||||
output_data,
|
||||
func,
|
||||
N);
|
||||
} else if (output_rank_or_simple_broadcast == static_cast<int32_t>(SimpleBroadcast::RightScalar)) {
|
||||
_BinaryElementWiseSimple<true, false, T, FuncT, GridDim::maxThreadsPerBlock,
|
||||
_BinaryElementWiseSimple<true, false, T, T1, FuncT, GridDim::maxThreadsPerBlock,
|
||||
GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
|
||||
lhs_data,
|
||||
rhs_data,
|
||||
|
|
@ -238,7 +238,7 @@ void BinaryElementWiseImpl(
|
|||
func,
|
||||
N);
|
||||
} else if (output_rank_or_simple_broadcast == static_cast<int32_t>(SimpleBroadcast::RightPerChannelBatch1)) {
|
||||
_BinaryElementWiseRhsPerChannelBatch1<T, FuncT, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
|
||||
_BinaryElementWiseRhsPerChannelBatch1<T, T1, FuncT, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
|
||||
lhs_data,
|
||||
rhs_data,
|
||||
fdm_H,
|
||||
|
|
@ -246,7 +246,7 @@ void BinaryElementWiseImpl(
|
|||
func,
|
||||
N);
|
||||
} else if (output_rank_or_simple_broadcast == static_cast<int32_t>(SimpleBroadcast::RightPerChannelBatchN)) {
|
||||
_BinaryElementWiseRhsPerChannelBatchN<T, FuncT, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
|
||||
_BinaryElementWiseRhsPerChannelBatchN<T, T1, FuncT, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
|
||||
lhs_data,
|
||||
rhs_data,
|
||||
fdm_H,
|
||||
|
|
@ -256,7 +256,7 @@ void BinaryElementWiseImpl(
|
|||
N);
|
||||
} else {
|
||||
if (lhs_padded_strides && rhs_padded_strides && lhs_padded_strides->size_ && rhs_padded_strides->size_)
|
||||
_BinaryElementWise<T, FuncT, true, true, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
|
||||
_BinaryElementWise<T, T1, FuncT, true, true, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
|
||||
output_rank_or_simple_broadcast,
|
||||
*lhs_padded_strides,
|
||||
lhs_data,
|
||||
|
|
@ -267,7 +267,7 @@ void BinaryElementWiseImpl(
|
|||
func,
|
||||
N);
|
||||
else if (lhs_padded_strides && lhs_padded_strides->size_)
|
||||
_BinaryElementWise<T, FuncT, true, false, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
|
||||
_BinaryElementWise<T, T1, FuncT, true, false, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
|
||||
output_rank_or_simple_broadcast,
|
||||
*lhs_padded_strides,
|
||||
lhs_data,
|
||||
|
|
@ -278,7 +278,7 @@ void BinaryElementWiseImpl(
|
|||
func,
|
||||
N);
|
||||
else
|
||||
_BinaryElementWise<T, FuncT, false, true, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
|
||||
_BinaryElementWise<T, T1, FuncT, false, true, GridDim::maxThreadsPerBlock, GridDim::maxElementsPerThread><<<blocksPerGrid, GridDim::maxThreadsPerBlock, 0>>>(
|
||||
output_rank_or_simple_broadcast,
|
||||
*lhs_padded_strides,
|
||||
lhs_data,
|
||||
|
|
|
|||
|
|
@ -181,8 +181,11 @@ __device__ __inline__ half2 _Tanh(half2 a) {
|
|||
return __float22half2_rn(tmp);
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
__device__ __inline__ T _Pow(T a, T b);
|
||||
// Capture permutations of int32/64/float/double
|
||||
template <typename T, typename T1>
|
||||
__device__ __inline__ T _Pow(T a, T1 b) {
|
||||
return static_cast<T>(pow(static_cast<double>(a), static_cast<double>(b)));
|
||||
}
|
||||
|
||||
template <>
|
||||
__device__ __inline__ float _Pow(float a, float b) { return powf(a, b); }
|
||||
|
|
|
|||
|
|
@ -313,9 +313,9 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain,
|
|||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, 10, float, Softmax);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, 10, double, Softmax);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, 10, MLFloat16, Softmax);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 7, float, Pow);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 7, double, Pow);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 7, MLFloat16, Pow);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 7, 11, float, Pow);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 7, 11, double, Pow);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 7, 11, MLFloat16, Pow);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 7, float, PRelu);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 7, double, PRelu);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 7, MLFloat16, PRelu);
|
||||
|
|
@ -755,6 +755,8 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain,
|
|||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, int8_t, MaxPool);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, uint8_t, MaxPool);
|
||||
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, Pow);
|
||||
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, float, ReduceMax);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, double, ReduceMax);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, MLFloat16, ReduceMax);
|
||||
|
|
@ -829,9 +831,9 @@ static Status RegisterCudaKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, 10, float, Softmax)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, 10, double, Softmax)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 1, 10, MLFloat16, Softmax)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 7, float, Pow)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 7, double, Pow)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 7, MLFloat16, Pow)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 7, 11, float, Pow)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 7, 11, double, Pow)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 7, 11, MLFloat16, Pow)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 7, float, PRelu)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 7, double, PRelu)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 7, MLFloat16, PRelu)>,
|
||||
|
|
@ -1270,6 +1272,8 @@ static Status RegisterCudaKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, int8_t, MaxPool)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, uint8_t, MaxPool)>,
|
||||
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, Pow)>,
|
||||
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, float, ReduceMax)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, double, ReduceMax)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 12, MLFloat16, ReduceMax)>,
|
||||
|
|
|
|||
|
|
@ -85,7 +85,7 @@ Status BinaryElementwise<ShouldBroadcast>::Prepare(OpKernelContext* context, Bin
|
|||
return Status::OK();
|
||||
}
|
||||
|
||||
#define BINARY_ELEMENTWISE_REGISTER_KERNEL_TYPED(x, ver, T) \
|
||||
#define BINARY_ELEMENTWISE_REGISTER_KERNEL_TYPED_V(x, class_name, ver, T) \
|
||||
ONNX_OPERATOR_TYPED_KERNEL_EX( \
|
||||
x, \
|
||||
kOnnxDomain, \
|
||||
|
|
@ -93,7 +93,19 @@ Status BinaryElementwise<ShouldBroadcast>::Prepare(OpKernelContext* context, Bin
|
|||
T, \
|
||||
kCudaExecutionProvider, \
|
||||
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
x<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, \
|
||||
kCudaExecutionProvider, \
|
||||
KernelDefBuilder().TypeConstraint("T", BuildKernelDefConstraints<>(__VAR_ARGS__)), \
|
||||
class_name);
|
||||
|
||||
#define BINARY_ELEMENTWISE_LOGICALOP_REGISTER_KERNEL_TYPED(x, ver, T) \
|
||||
ONNX_OPERATOR_TYPED_KERNEL_EX( \
|
||||
|
|
@ -116,6 +128,17 @@ Status BinaryElementwise<ShouldBroadcast>::Prepare(OpKernelContext* context, Bin
|
|||
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, \
|
||||
kCudaExecutionProvider, \
|
||||
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
class_name<T>);
|
||||
|
||||
#define BINARY_ELEMENTWISE_COMPUTE(x, T) \
|
||||
template <> \
|
||||
Status x<T>::ComputeInternal(OpKernelContext* context) const { \
|
||||
|
|
@ -139,6 +162,10 @@ Status BinaryElementwise<ShouldBroadcast>::Prepare(OpKernelContext* context, Bin
|
|||
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)
|
||||
|
|
@ -211,6 +238,11 @@ Status BinaryElementwise<ShouldBroadcast>::Prepare(OpKernelContext* context, Bin
|
|||
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) \
|
||||
|
|
@ -222,12 +254,116 @@ BINARY_OP_UZILHFD(Add, 7)
|
|||
BINARY_OP_UZILHFD(Sub, 7)
|
||||
BINARY_OP_UZILHFD(Mul, 7)
|
||||
BINARY_OP_UZILHFD(Div, 7)
|
||||
BINARY_OP_HFD(Pow, 7)
|
||||
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_HFD(PRelu, 7)
|
||||
|
||||
// Pow version 12
|
||||
ONNX_OPERATOR_KERNEL_EX(
|
||||
Pow,
|
||||
kOnnxDomain,
|
||||
12,
|
||||
kCudaExecutionProvider,
|
||||
KernelDefBuilder().TypeConstraint("T", BuildKernelDefConstraints<int32_t, int64_t, float, double>())
|
||||
.TypeConstraint("T1", BuildKernelDefConstraints<int32_t, int64_t, float, double>()),
|
||||
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 ToCudaType<T>::MappedType, typename ToCudaType<int32_t>::MappedType>(
|
||||
prepare.output_rank_or_simple_broadcast,
|
||||
&prepare.lhs_padded_strides,
|
||||
reinterpret_cast<const typename ToCudaType<T>::MappedType*>(prepare.lhs_tensor->template Data<T>()),
|
||||
&prepare.rhs_padded_strides,
|
||||
reinterpret_cast<const typename ToCudaType<int32_t>::MappedType*>(prepare.rhs_tensor->template Data<int32_t>()),
|
||||
&prepare.fdm_output_strides,
|
||||
prepare.fdm_H,
|
||||
prepare.fdm_C,
|
||||
reinterpret_cast<typename ToCudaType<T>::MappedType*>(prepare.output_tensor->template MutableData<T>()),
|
||||
prepare.output_tensor->Shape().Size());
|
||||
break;
|
||||
case on::TensorProto_DataType_INT64:
|
||||
ImplT1_Pow<typename ToCudaType<T>::MappedType, typename ToCudaType<int64_t>::MappedType>(
|
||||
prepare.output_rank_or_simple_broadcast,
|
||||
&prepare.lhs_padded_strides,
|
||||
reinterpret_cast<const typename ToCudaType<T>::MappedType*>(prepare.lhs_tensor->template Data<T>()),
|
||||
&prepare.rhs_padded_strides,
|
||||
reinterpret_cast<const typename ToCudaType<int64_t>::MappedType*>(prepare.rhs_tensor->template Data<int64_t>()),
|
||||
&prepare.fdm_output_strides,
|
||||
prepare.fdm_H,
|
||||
prepare.fdm_C,
|
||||
reinterpret_cast<typename ToCudaType<T>::MappedType*>(prepare.output_tensor->template MutableData<T>()),
|
||||
prepare.output_tensor->Shape().Size());
|
||||
break;
|
||||
case on::TensorProto_DataType_FLOAT:
|
||||
ImplT1_Pow<typename ToCudaType<T>::MappedType, typename ToCudaType<float>::MappedType>(
|
||||
prepare.output_rank_or_simple_broadcast,
|
||||
&prepare.lhs_padded_strides,
|
||||
reinterpret_cast<const typename ToCudaType<T>::MappedType*>(prepare.lhs_tensor->template Data<T>()),
|
||||
&prepare.rhs_padded_strides,
|
||||
reinterpret_cast<const typename ToCudaType<float>::MappedType*>(prepare.rhs_tensor->template Data<float>()),
|
||||
&prepare.fdm_output_strides,
|
||||
prepare.fdm_H,
|
||||
prepare.fdm_C,
|
||||
reinterpret_cast<typename ToCudaType<T>::MappedType*>(prepare.output_tensor->template MutableData<T>()),
|
||||
prepare.output_tensor->Shape().Size());
|
||||
break;
|
||||
case on::TensorProto_DataType_DOUBLE:
|
||||
ImplT1_Pow<typename ToCudaType<T>::MappedType, typename ToCudaType<double>::MappedType>(
|
||||
prepare.output_rank_or_simple_broadcast,
|
||||
&prepare.lhs_padded_strides,
|
||||
reinterpret_cast<const typename ToCudaType<T>::MappedType*>(prepare.lhs_tensor->template Data<T>()),
|
||||
&prepare.rhs_padded_strides,
|
||||
reinterpret_cast<const typename ToCudaType<double>::MappedType*>(prepare.rhs_tensor->template Data<double>()),
|
||||
&prepare.fdm_output_strides,
|
||||
prepare.fdm_H,
|
||||
prepare.fdm_C,
|
||||
reinterpret_cast<typename ToCudaType<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;
|
||||
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;
|
||||
default:
|
||||
s = ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "Unsupported X type: ",
|
||||
DataTypeImpl::ToString(prepare.lhs_tensor->DataType()));
|
||||
}
|
||||
return s;
|
||||
}
|
||||
|
||||
template <typename T, typename CudaT>
|
||||
Status VariadicInputBase<T, CudaT>::ComputeMethod(OpKernelContext* context, ImplCompute Impl_Compute) const {
|
||||
const auto& node = Node();
|
||||
|
|
|
|||
|
|
@ -164,6 +164,13 @@ class Div final : public BinaryElementwise<ShouldBroadcast> {
|
|||
};
|
||||
|
||||
template <typename T>
|
||||
class Pow_7 final : public BinaryElementwise<ShouldBroadcast> {
|
||||
public:
|
||||
Pow_7(const OpKernelInfo& info) : BinaryElementwise(info) {}
|
||||
Status ComputeInternal(OpKernelContext* context) const override;
|
||||
};
|
||||
|
||||
// Since version 12
|
||||
class Pow final : public BinaryElementwise<ShouldBroadcast> {
|
||||
public:
|
||||
Pow(const OpKernelInfo& info) : BinaryElementwise(info) {}
|
||||
|
|
|
|||
|
|
@ -9,12 +9,12 @@
|
|||
namespace onnxruntime {
|
||||
namespace cuda {
|
||||
|
||||
#define OP(name, expr) \
|
||||
template <class T> \
|
||||
struct OP_##name { \
|
||||
__device__ __inline__ T operator()(T a, T b) const { \
|
||||
return (expr); \
|
||||
} \
|
||||
#define OP(name, expr) \
|
||||
template <class T, class T1> \
|
||||
struct OP_##name { \
|
||||
__device__ __inline__ T operator()(T a, T1 b) const { \
|
||||
return (expr); \
|
||||
} \
|
||||
};
|
||||
|
||||
#define BINARY_ELEMENTWISE_IMPL(name) \
|
||||
|
|
@ -28,15 +28,36 @@ namespace cuda {
|
|||
fdm_H, \
|
||||
fdm_C, \
|
||||
output_data, \
|
||||
OP_##name<T>(), \
|
||||
OP_##name<T, T>(), \
|
||||
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 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, T1>(), \
|
||||
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_UZILHFD(x) \
|
||||
SPECIALIZED_BINARY_ELEMENTWISE_IMPL(x, uint32_t) \
|
||||
|
|
@ -85,7 +106,7 @@ 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)
|
||||
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)
|
||||
|
|
@ -96,5 +117,29 @@ SPECIALIZED_BINARY_ELEMENTWISE_IMPL_UZILHFD(Max)
|
|||
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_UZILHFD(Min)
|
||||
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_UZILHFD(Less)
|
||||
|
||||
// create declarations for op and impl for Pow
|
||||
OP(Pow, _Pow(a, b))
|
||||
BINARY_ELEMENTWISE_IMPL_T1(Pow)
|
||||
|
||||
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(Pow, int32_t, int32_t)
|
||||
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(Pow, int32_t, int64_t)
|
||||
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(Pow, int32_t, float)
|
||||
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(Pow, int32_t, double)
|
||||
|
||||
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(Pow, int64_t, int32_t)
|
||||
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(Pow, int64_t, int64_t)
|
||||
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(Pow, int64_t, float)
|
||||
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(Pow, int64_t, double)
|
||||
|
||||
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(Pow, float, int32_t)
|
||||
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(Pow, float, int64_t)
|
||||
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(Pow, float, float)
|
||||
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(Pow, float, double)
|
||||
|
||||
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(Pow, double, int32_t)
|
||||
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(Pow, double, int64_t)
|
||||
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(Pow, double, float)
|
||||
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_T1(Pow, double, double)
|
||||
|
||||
} // namespace cuda
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
|
|
@ -20,7 +20,7 @@ namespace cuda {
|
|||
BINARY_OP_NAME_EXPR(Sub, (a - b)) \
|
||||
BINARY_OP_NAME_EXPR(Mul, (a * b)) \
|
||||
BINARY_OP_NAME_EXPR(Div, (a / b)) \
|
||||
BINARY_OP_NAME_EXPR(Pow, _Pow(a, b)) \
|
||||
BINARY_OP_NAME_EXPR(Pow_7, _Pow(a, b)) \
|
||||
BINARY_OP_NAME_EXPR(And, (a & b)) \
|
||||
BINARY_OP_NAME_EXPR(Or, (a | b)) \
|
||||
BINARY_OP_NAME_EXPR(Xor, (a ^ b)) \
|
||||
|
|
@ -34,23 +34,39 @@ namespace cuda {
|
|||
// NOTE that cu files are compiled with nvcc and should not refer to any onnxruntime headers
|
||||
// so struct BinaryElementwisePreparation cannot be used here
|
||||
|
||||
#define BINARY_ELEMENTWISE_IMPL_DECLARATION(name) \
|
||||
template <typename T> \
|
||||
void Impl_##name( \
|
||||
int32_t output_rank_or_simple_broadcast, \
|
||||
const TArray<int64_t>* lhs_padded_strides, \
|
||||
const T* lhs_data, \
|
||||
const TArray<int64_t>* rhs_padded_strides, \
|
||||
const T* rhs_data, \
|
||||
#define BINARY_ELEMENTWISE_IMPL_DECLARATION(name) \
|
||||
template <typename T> \
|
||||
void Impl_##name( \
|
||||
int32_t output_rank_or_simple_broadcast, \
|
||||
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, \
|
||||
const fast_divmod& fdm_H, \
|
||||
const fast_divmod& fdm_C, \
|
||||
T* output_data, \
|
||||
size_t count)
|
||||
|
||||
#define BINARY_OP_NAME_EXPR(name, expr) BINARY_ELEMENTWISE_IMPL_DECLARATION(name);
|
||||
BINARY_OPS()
|
||||
#undef BINARY_OP_NAME_EXPR
|
||||
|
||||
#define BINARY_ELEMENTWISE_IMPL_DECLARATION_T1(name) \
|
||||
template <typename T, typename T1> \
|
||||
void ImplT1_##name( \
|
||||
int32_t output_rank_or_simple_broadcast, \
|
||||
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)
|
||||
|
||||
BINARY_ELEMENTWISE_IMPL_DECLARATION_T1(Pow);
|
||||
|
||||
} // namespace cuda
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
|
|
@ -621,6 +621,110 @@ TEST(MathOpTest, Pow_Broadcast_Scalar1) {
|
|||
test.Run();
|
||||
}
|
||||
|
||||
TEST(MathOpTest, Pow_Float_12) {
|
||||
OpTester test("Pow", 12);
|
||||
std::vector<int64_t> dims{2, 2};
|
||||
test.AddInput<float>("X", dims,
|
||||
{2.0f, 2.0f,
|
||||
std::sqrt(2.0f), 1.0f});
|
||||
test.AddInput<float>("Y", dims,
|
||||
{0.0f, 8.0f,
|
||||
2.0f, 9.0f});
|
||||
test.AddOutput<float>("Z", dims,
|
||||
{1.0f, 256.0f,
|
||||
2.0f, 1.0f});
|
||||
test.Run();
|
||||
}
|
||||
|
||||
TEST(MathOpTest, Pow_Double_12) {
|
||||
OpTester test("Pow", 12);
|
||||
std::vector<int64_t> dims{2, 2};
|
||||
test.AddInput<double>("X", dims,
|
||||
{2.0, 2.0,
|
||||
std::sqrt(2.0), 1.0});
|
||||
test.AddInput<double>("Y", dims,
|
||||
{0.0, 8.0,
|
||||
2.0, 9.0});
|
||||
test.AddOutput<double>("Z", dims,
|
||||
{1.0, 256.0,
|
||||
2.0, 1.0});
|
||||
test.Run();
|
||||
}
|
||||
|
||||
TEST(MathOpTest, Pow_Broadcast_Scalar0_12) {
|
||||
OpTester test("Pow", 12);
|
||||
|
||||
std::vector<int64_t> dims{3};
|
||||
test.AddInput<float>("X", {}, {2.0f});
|
||||
test.AddInput<float>("Y", dims, {1.0f, 2.0f, 3.0f});
|
||||
test.AddOutput<float>("Z", dims, {2.0f, 4.0f, 8.0f});
|
||||
test.Run();
|
||||
}
|
||||
|
||||
TEST(MathOpTest, Pow_Broadcast_Scalar1_12) {
|
||||
OpTester test("Pow", 12);
|
||||
|
||||
std::vector<int64_t> dims{3};
|
||||
test.AddInput<float>("X", dims, {1.0f, 2.0f, 3.0f});
|
||||
test.AddInput<float>("Y", {}, {2.0f});
|
||||
test.AddOutput<float>("Z", dims, {1.0f, 4.0f, 9.0f});
|
||||
test.Run();
|
||||
}
|
||||
|
||||
TEST(MathOpTest, Pow_float_int64) {
|
||||
OpTester test("Pow", 12);
|
||||
std::vector<int64_t> dims{3};
|
||||
test.AddInput<float>("X", dims, {1.0f, 2.0f, 3.0f});
|
||||
test.AddInput<int64_t>("Y", dims, {4, 5, 6});
|
||||
test.AddOutput<float>("Z", dims, {1.f, 32.f, 729.f});
|
||||
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kNGraphExecutionProvider});
|
||||
}
|
||||
|
||||
TEST(MathOpTest, Pow_int64_float) {
|
||||
OpTester test("Pow", 12);
|
||||
std::vector<int64_t> dims{3};
|
||||
test.AddInput<int64_t>("X", dims, {1, 2, 3});
|
||||
test.AddInput<float>("Y", dims, {4.f, 5.f, 6.f});
|
||||
test.AddOutput<int64_t>("Z", dims, {1, 32, 729});
|
||||
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kNGraphExecutionProvider});
|
||||
}
|
||||
|
||||
TEST(MathOpTest, Pow_float_int32) {
|
||||
OpTester test("Pow", 12);
|
||||
std::vector<int64_t> dims{3};
|
||||
test.AddInput<float>("X", dims, {1.0f, 2.0f, 3.0f});
|
||||
test.AddInput<int32_t>("Y", dims, {4, 5, 6});
|
||||
test.AddOutput<float>("Z", dims, {1.f, 32.f, 729.f});
|
||||
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kNGraphExecutionProvider});
|
||||
}
|
||||
|
||||
TEST(MathOpTest, Pow_int32_float) {
|
||||
OpTester test("Pow", 12);
|
||||
std::vector<int64_t> dims{3};
|
||||
test.AddInput<int32_t>("X", dims, {1, 2, 3});
|
||||
test.AddInput<float>("Y", dims, {4.f, 5.f, 6.f});
|
||||
test.AddOutput<int32_t>("Z", dims, {1, 32, 729});
|
||||
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kNGraphExecutionProvider});
|
||||
}
|
||||
|
||||
TEST(MathOpTest, Pow_int64_double) {
|
||||
OpTester test("Pow", 12);
|
||||
std::vector<int64_t> dims{3};
|
||||
test.AddInput<int64_t>("X", dims, {1, 2, 3});
|
||||
test.AddInput<double>("Y", dims, {4.f, 5.f, 6.f});
|
||||
test.AddOutput<int64_t>("Z", dims, {1, 32, 729});
|
||||
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kNGraphExecutionProvider});
|
||||
}
|
||||
|
||||
TEST(MathOpTest, Pow_double_int64) {
|
||||
OpTester test("Pow", 12);
|
||||
std::vector<int64_t> dims{3};
|
||||
test.AddInput<double>("X", dims, {1., 2., 3.});
|
||||
test.AddInput<int64_t>("Y", dims, {4, 5, 6});
|
||||
test.AddOutput<double>("Z", dims, {1., 32., 729.});
|
||||
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kNGraphExecutionProvider});
|
||||
}
|
||||
|
||||
TEST(MathOpTest, Exp_float) {
|
||||
OpTester test("Exp");
|
||||
std::vector<int64_t> dims{2, 2};
|
||||
|
|
|
|||
|
|
@ -132,20 +132,8 @@ def create_backend_test(testname=None):
|
|||
'^test_negative_log_likelihood_loss_input_shape_is_NCd1d2_with_weight_reduction_sum_cpu',
|
||||
'^test_negative_log_likelihood_loss_input_shape_is_NCd1d2_with_weight_reduction_sum_ignore_index_cpu',
|
||||
'^test_nesterov_momentum_cpu',
|
||||
'^test_pow_bcast_array_cpu',
|
||||
'^test_pow_bcast_scalar_cpu',
|
||||
'^test_pow_cpu',
|
||||
'^test_pow_example_cpu',
|
||||
'^test_pow_types_float32_int32_cpu',
|
||||
'^test_pow_types_float32_int64_cpu',
|
||||
'^test_pow_types_float32_uint32_cpu',
|
||||
'^test_pow_types_float32_uint64_cpu',
|
||||
'^test_pow_types_float_cpu',
|
||||
'^test_pow_types_int32_float32_cpu',
|
||||
'^test_pow_types_int32_int32_cpu',
|
||||
'^test_pow_types_int64_float32_cpu',
|
||||
'^test_pow_types_int64_int64_cpu',
|
||||
'^test_pow_types_int_cpu',
|
||||
'^test_softmax_cross_entropy_mean_3d_cpu',
|
||||
'^test_softmax_cross_entropy_mean_3d_expanded_cpu',
|
||||
'^test_softmax_cross_entropy_mean_cpu',
|
||||
|
|
@ -182,7 +170,11 @@ def create_backend_test(testname=None):
|
|||
'^test_argmin_negative_axis.*', '^test_hardmax_negative_axis.*', '^test_gemm_default_no_bias_cpu',
|
||||
'^test_flatten_negative_axis.*', '^test_reduce_[a-z1-9_]*_negative_axes_.*',
|
||||
'test_squeeze_negative_axes_cpu', 'test_unsqueeze_negative_axes_cpu', 'test_constant_pad_cpu',
|
||||
'test_edge_pad_cpu', 'test_reflect_pad_cpu', '^test_split_zero_size_splits_.*','^test_argmax_keepdims_example_select_last_index_.*', '^test_argmax_no_keepdims_example_select_last_index_.*','^test_argmin_no_keepdims_example_select_last_index_.*','^test_argmin_keepdims_example_select_last_index_.*'
|
||||
'test_edge_pad_cpu', 'test_reflect_pad_cpu', '^test_split_zero_size_splits_.*',
|
||||
'^test_argmax_keepdims_example_select_last_index_cpu', '^test_argmax_no_keepdims_example_select_last_index_cpu',
|
||||
'^test_argmin_no_keepdims_example_select_last_index_cpu','^test_argmin_keepdims_example_select_last_index_cpu',
|
||||
'^test_pow_types_float32_int32_*', '^test_pow_types_float32_int64_*', '^test_pow_types_float_*',
|
||||
'^test_pow_types_int32_float32_*', '^test_pow_types_int_*', '^test_pow_types_int64_float32_*'
|
||||
]
|
||||
|
||||
if c2.supports_device('DNNL'):
|
||||
|
|
|
|||
Loading…
Reference in a new issue