diff --git a/docs/OperatorKernels.md b/docs/OperatorKernels.md
index e5123f172c..e1b5253548 100644
--- a/docs/OperatorKernels.md
+++ b/docs/OperatorKernels.md
@@ -601,6 +601,8 @@ Do not modify directly.*
|Min|*in* data_0:**T**
*out* min:**T**|13+|**T** = tensor(bfloat16), tensor(double), tensor(float), tensor(float16), tensor(int32), tensor(int64), tensor(uint32), tensor(uint64)|
|||12|**T** = tensor(bfloat16), tensor(double), tensor(float), tensor(float16), tensor(int32), tensor(int64), tensor(uint32), tensor(uint64)|
|||[6, 11]|**T** = tensor(bfloat16), tensor(double), tensor(float), tensor(float16)|
+|Mod|*in* A:**T**
*in* B:**T**
*out* C:**T**|13+|**T** = tensor(bfloat16), tensor(double), tensor(float), tensor(float16), tensor(int32), tensor(int64), tensor(uint32), tensor(uint64)|
+|||[10, 12]|**T** = tensor(double), tensor(float), tensor(float16), tensor(int32), tensor(int64), tensor(uint32), tensor(uint64)|
|Mul|*in* A:**T**
*in* B:**T**
*out* C:**T**|14+|**T** = tensor(bfloat16), tensor(double), tensor(float), tensor(float16), tensor(int32), tensor(int64), tensor(uint32), tensor(uint64)|
|||13|**T** = tensor(bfloat16), tensor(double), tensor(float), tensor(float16), tensor(int32), tensor(int64), tensor(uint32), tensor(uint64)|
|||[7, 12]|**T** = tensor(double), tensor(float), tensor(float16), tensor(int32), tensor(int64), tensor(uint32), tensor(uint64)|
diff --git a/onnxruntime/core/providers/cuda/cu_inc/common.cuh b/onnxruntime/core/providers/cuda/cu_inc/common.cuh
index 86968d6f03..7722ca61e6 100644
--- a/onnxruntime/core/providers/cuda/cu_inc/common.cuh
+++ b/onnxruntime/core/providers/cuda/cu_inc/common.cuh
@@ -359,6 +359,41 @@ __device__ __inline__ T _Gelu(T a) {
return a * _Normcdf(a);
}
+template
+__device__ __inline__ T _Mod(T a, T b) {
+ T r = a % b;
+ T zero = T(0);
+ if ((r > zero && b < zero) || (r < zero && b > zero)) {
+ r += b;
+ }
+ return r;
+}
+
+template
+__device__ __inline__ T _Fmod(T a, T b) {
+ return a % b;
+}
+
+template <>
+__device__ __inline__ float _Fmod(float a, float b) {
+ return fmodf(a, b);
+}
+
+template <>
+__device__ __inline__ double _Fmod(double a, double b) {
+ return fmod(a, b);
+}
+
+template <>
+__device__ __inline__ half _Fmod(half a, half b) {
+ return fmodf((float)a, (float)b);
+}
+
+template <>
+__device__ __inline__ BFloat16 _Fmod(BFloat16 a, BFloat16 b) {
+ return fmodf((float)a, (float)b);
+}
+
// We would like to use 64-bit integer to support large matrices. However, CUDA seems to support only 32-bit integer
// For now, use int32_t to ensure that both Linux and Windows see this as 32 bit integer type.
#ifndef CUDA_LONG
diff --git a/onnxruntime/core/providers/cuda/cuda_execution_provider.cc b/onnxruntime/core/providers/cuda/cuda_execution_provider.cc
index c1b5550611..57e0f1aaf9 100755
--- a/onnxruntime/core/providers/cuda/cuda_execution_provider.cc
+++ b/onnxruntime/core/providers/cuda/cuda_execution_provider.cc
@@ -833,6 +833,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain,
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 10, double, ThresholdedRelu);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 10, MLFloat16, ThresholdedRelu);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 10, 10, TopK);
+class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 10, 12, Mod);
// opset 11
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 11, float, ArgMax);
@@ -1217,6 +1218,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain,
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 13, BFloat16, Tanh);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 13, BFloat16, Gemm);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 13, BFloat16, ReduceSum);
+class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 13, Mod);
// OpSet 14
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kOnnxDomain, 14, CumSum);
@@ -1709,6 +1711,7 @@ static Status RegisterCudaKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo,
BuildKernelCreateInfo,
BuildKernelCreateInfo,
+ BuildKernelCreateInfo,
// opset 11
BuildKernelCreateInfo,
@@ -2089,6 +2092,7 @@ static Status RegisterCudaKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo,
BuildKernelCreateInfo,
BuildKernelCreateInfo,
+ BuildKernelCreateInfo,
// OpSet 14
BuildKernelCreateInfo,
diff --git a/onnxruntime/core/providers/cuda/math/binary_elementwise_ops.cc b/onnxruntime/core/providers/cuda/math/binary_elementwise_ops.cc
index 346163ecec..931f9eb137 100644
--- a/onnxruntime/core/providers/cuda/math/binary_elementwise_ops.cc
+++ b/onnxruntime/core/providers/cuda/math/binary_elementwise_ops.cc
@@ -465,6 +465,69 @@ Status Pow::ComputeInternal(OpKernelContext* context) const {
return s;
}
+ONNX_OPERATOR_VERSIONED_KERNEL_EX(
+ Mod, kOnnxDomain, 10, 12, kCudaExecutionProvider,
+ (*KernelDefBuilder::Create())
+ .TypeConstraint("T",
+ BuildKernelDefConstraints()),
+ Mod);
+
+ONNX_OPERATOR_KERNEL_EX(Mod, kOnnxDomain, 13, kCudaExecutionProvider,
+ (*KernelDefBuilder::Create())
+ .TypeConstraint("T", BuildKernelDefConstraints()),
+ Mod);
+
+Status Mod::ComputeInternal(OpKernelContext* context) const {
+ namespace on = ONNX_NAMESPACE;
+ BinaryElementwisePreparation prepare;
+ ORT_RETURN_IF_ERROR(Prepare(context, &prepare));
+ auto element_type = prepare.lhs_tensor->GetElementType();
+ ORT_ENFORCE(fmod_ || element_type == on::TensorProto_DataType_INT32 ||
+ element_type == on::TensorProto_DataType_INT64 || element_type == on::TensorProto_DataType_UINT32 ||
+ element_type == on::TensorProto_DataType_UINT64,
+ "Non-fmod can support integer types only.");
+#define CASE_MOD_ELEMENT_TYPE(name, onnx_type, data_type) \
+ case onnx_type: { \
+ Impl_##name::MappedType>( \
+ Stream(), prepare.output_rank_or_simple_broadcast, &prepare.lhs_padded_strides, \
+ reinterpret_cast::MappedType*>(prepare.lhs_tensor->Data()), \
+ &prepare.rhs_padded_strides, \
+ reinterpret_cast::MappedType*>(prepare.rhs_tensor->Data()), \
+ &prepare.fdm_output_strides, prepare.fdm_H, prepare.fdm_C, \
+ reinterpret_cast::MappedType*>( \
+ prepare.output_tensor->MutableData()), \
+ prepare.output_tensor->Shape().Size()); \
+ } break
+ if (fmod_) {
+ switch (element_type) {
+ CASE_MOD_ELEMENT_TYPE(Fmod, on::TensorProto_DataType_INT32, int32_t);
+ CASE_MOD_ELEMENT_TYPE(Fmod, on::TensorProto_DataType_INT64, int64_t);
+ CASE_MOD_ELEMENT_TYPE(Fmod, on::TensorProto_DataType_UINT32, uint32_t);
+ CASE_MOD_ELEMENT_TYPE(Fmod, on::TensorProto_DataType_UINT64, uint64_t);
+ CASE_MOD_ELEMENT_TYPE(Fmod, on::TensorProto_DataType_FLOAT, float);
+ CASE_MOD_ELEMENT_TYPE(Fmod, on::TensorProto_DataType_DOUBLE, double);
+ CASE_MOD_ELEMENT_TYPE(Fmod, on::TensorProto_DataType_FLOAT16, MLFloat16);
+ CASE_MOD_ELEMENT_TYPE(Fmod, on::TensorProto_DataType_BFLOAT16, BFloat16);
+ default:
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Unsupported element type: ", DataTypeImpl::ToString(prepare.lhs_tensor->DataType()));
+ }
+ } else {
+ switch (element_type) {
+ CASE_MOD_ELEMENT_TYPE(Mod, on::TensorProto_DataType_INT32, int32_t);
+ CASE_MOD_ELEMENT_TYPE(Mod, on::TensorProto_DataType_INT64, int64_t);
+ CASE_MOD_ELEMENT_TYPE(Mod, on::TensorProto_DataType_UINT32, uint32_t);
+ CASE_MOD_ELEMENT_TYPE(Mod, on::TensorProto_DataType_UINT64, uint64_t);
+ default:
+ return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT,
+ "Unsupported element type: ", DataTypeImpl::ToString(prepare.lhs_tensor->DataType()));
+ }
+ }
+#undef CASE_MOD_ELEMENT_TYPE
+ return Status::OK();
+}
+
//Greater op output tensor type is bool, so it cannot directly fit in the macros
//for other elementwise ops
template
diff --git a/onnxruntime/core/providers/cuda/math/binary_elementwise_ops.h b/onnxruntime/core/providers/cuda/math/binary_elementwise_ops.h
index 5cd728189d..fcab87c193 100644
--- a/onnxruntime/core/providers/cuda/math/binary_elementwise_ops.h
+++ b/onnxruntime/core/providers/cuda/math/binary_elementwise_ops.h
@@ -214,6 +214,18 @@ class PRelu final : public BinaryElementwise {
Status ComputeInternal(OpKernelContext* context) const override;
};
+class Mod final : public BinaryElementwise {
+ public:
+ Mod(const OpKernelInfo& info) : BinaryElementwise(info) {
+ int64_t fmod = info.GetAttrOrDefault("fmod", 0LL);
+ fmod_ = fmod != 0;
+ }
+ Status ComputeInternal(OpKernelContext* context) const override;
+
+ private:
+ bool fmod_{false};
+};
+
template
class CompareFunction : public BinaryElementwise {
public:
diff --git a/onnxruntime/core/providers/cuda/math/binary_elementwise_ops_impl.cu b/onnxruntime/core/providers/cuda/math/binary_elementwise_ops_impl.cu
index e408cbbc04..0a7753783f 100644
--- a/onnxruntime/core/providers/cuda/math/binary_elementwise_ops_impl.cu
+++ b/onnxruntime/core/providers/cuda/math/binary_elementwise_ops_impl.cu
@@ -89,6 +89,12 @@ namespace cuda {
SPECIALIZED_BINARY_ELEMENTWISE_IMPL(x, double) \
SPECIALIZED_BINARY_ELEMENTWISE_IMPL(x, BFloat16)
+#define SPECIALIZED_BINARY_ELEMENTWISE_IMPL_UZIL(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)
+
#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) \
@@ -141,6 +147,8 @@ SPECIALIZED_BINARY_ELEMENTWISE_IMPL(Xor, bool)
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_HFD(PRelu)
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_UZILHFD(Max)
SPECIALIZED_BINARY_ELEMENTWISE_IMPL_UZILHFD(Min)
+SPECIALIZED_BINARY_ELEMENTWISE_IMPL_UZIL(Mod)
+SPECIALIZED_BINARY_ELEMENTWISE_IMPL_UZILHFD(Fmod)
// create declarations for impl for Pow
BINARY_ELEMENTWISE_IMPL_T1(Pow)
diff --git a/onnxruntime/core/providers/cuda/math/binary_elementwise_ops_impl.h b/onnxruntime/core/providers/cuda/math/binary_elementwise_ops_impl.h
index 4fd616dc7c..26f35efe3d 100644
--- a/onnxruntime/core/providers/cuda/math/binary_elementwise_ops_impl.h
+++ b/onnxruntime/core/providers/cuda/math/binary_elementwise_ops_impl.h
@@ -26,7 +26,9 @@ namespace cuda {
BINARY_OP_NAME_EXPR(Xor, (a ^ b)) \
BINARY_OP_NAME_EXPR(PRelu, (a > (T)0 ? a : a * b)) \
BINARY_OP_NAME_EXPR(Max, _Max(a, b)) \
- BINARY_OP_NAME_EXPR(Min, _Min(a, b))
+ BINARY_OP_NAME_EXPR(Min, _Min(a, b)) \
+ BINARY_OP_NAME_EXPR(Mod, _Mod(a, b)) \
+ BINARY_OP_NAME_EXPR(Fmod, _Fmod(a, b))
// NOTE that cu files are compiled with nvcc and should not refer to any onnxruntime headers
// so struct BinaryElementwisePreparation cannot be used here
diff --git a/onnxruntime/core/providers/rocm/cu_inc/common.cuh b/onnxruntime/core/providers/rocm/cu_inc/common.cuh
index 7ae39773c0..13923d9ae9 100644
--- a/onnxruntime/core/providers/rocm/cu_inc/common.cuh
+++ b/onnxruntime/core/providers/rocm/cu_inc/common.cuh
@@ -280,6 +280,41 @@ __device__ __inline__ T _Gelu(T a) {
return a * _Normcdf(a);
}
+template
+__device__ __inline__ T _Mod(T a, T b) {
+ T r = a % b;
+ T zero = T(0);
+ if ((r > zero && b < zero) || (r < zero && b > zero)) {
+ r += b;
+ }
+ return r;
+}
+
+template
+__device__ __inline__ T _Fmod(T a, T b) {
+ return a % b;
+}
+
+template <>
+__device__ __inline__ float _Fmod(float a, float b) {
+ return fmodf(a, b);
+}
+
+template <>
+__device__ __inline__ double _Fmod(double a, double b) {
+ return fmod(a, b);
+}
+
+template <>
+__device__ __inline__ half _Fmod(half a, half b) {
+ return fmodf((float)a, (float)b);
+}
+
+template <>
+__device__ __inline__ BFloat16 _Fmod(BFloat16 a, BFloat16 b) {
+ return fmodf((float)a, (float)b);
+}
+
// We would like to use 64-bit integer to support large matrices. However, ROCM seems to support only 32-bit integer
// For now, use int32_t to ensure that both Linux and Windows see this as 32 bit integer type.
#ifndef HIP_LONG
diff --git a/onnxruntime/core/providers/rocm/rocm_execution_provider.cc b/onnxruntime/core/providers/rocm/rocm_execution_provider.cc
index 4cfa91cde5..10376f4ece 100644
--- a/onnxruntime/core/providers/rocm/rocm_execution_provider.cc
+++ b/onnxruntime/core/providers/rocm/rocm_execution_provider.cc
@@ -746,6 +746,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 10, double, ThresholdedRelu);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 10, MLFloat16, ThresholdedRelu);
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 10, 10, TopK);
+class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 10, 12, Mod);
// opset 11
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 11, float, ArgMax);
@@ -1130,6 +1131,7 @@ class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain,
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, BFloat16, Tanh);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, BFloat16, Gemm);
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, BFloat16, ReduceSum);
+class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 13, Mod);
// OpSet 14
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kOnnxDomain, 14, CumSum);
@@ -1652,6 +1654,7 @@ static Status RegisterRocmKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo,
BuildKernelCreateInfo,
BuildKernelCreateInfo,
+ BuildKernelCreateInfo,
// opset 11
BuildKernelCreateInfo,
@@ -2046,6 +2049,7 @@ static Status RegisterRocmKernels(KernelRegistry& kernel_registry) {
BuildKernelCreateInfo,
BuildKernelCreateInfo,
BuildKernelCreateInfo,
+ BuildKernelCreateInfo,
// OpSet 14
BuildKernelCreateInfo,
diff --git a/onnxruntime/test/providers/cpu/math/element_wise_ops_test.cc b/onnxruntime/test/providers/cpu/math/element_wise_ops_test.cc
index 43e9ead580..d6e3f0a461 100644
--- a/onnxruntime/test/providers/cpu/math/element_wise_ops_test.cc
+++ b/onnxruntime/test/providers/cpu/math/element_wise_ops_test.cc
@@ -33,7 +33,7 @@ void TestFloat16(const char* op_name, const std::vector& lhs_dim,
execution_providers.push_back(DefaultCudaExecutionProvider());
#elif USE_ROCM
execution_providers.push_back(DefaultRocmExecutionProvider());
-#endif
+#endif
tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
}
@@ -47,7 +47,7 @@ void TestFloat16(const char* op_name, const std::vector& lhs_dim,
execution_providers.push_back(DefaultCudaExecutionProvider());
#elif USE_ROCM
execution_providers.push_back(DefaultRocmExecutionProvider());
-#endif
+#endif
tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
}
}
@@ -2293,6 +2293,15 @@ TEST(ModOpTest, Fmod_float_mixed_sign) {
test.Run();
}
+TEST(ModOpTest, Fmod_double_mixed_sign) {
+ OpTester test("Mod", ModOp_ver);
+ test.AddAttribute("fmod", 1);
+ test.AddInput("X", {6}, {-4.3, 7.2, 5.0, 4.3, -7.2, 8.0});
+ test.AddInput("Y", {6}, {2.1f, -3.4, 8.0, -2.1, 3.4, 5.0});
+ test.AddOutput("Z", {6}, {-0.1, 0.4, 5., 0.1, -0.4, 3.});
+ test.Run();
+}
+
TEST(ModOpTest, Fmod_float16_mixed_sign) {
OpTester test("Mod", ModOp_ver);
test.AddAttribute("fmod", 1);
@@ -2305,6 +2314,24 @@ TEST(ModOpTest, Fmod_float16_mixed_sign) {
test.Run();
}
+#if defined(USE_CUDA) || defined(USE_ROCM)
+TEST(ModOpTest, Fmod_bfloat16_mixed_sign) {
+ OpTester test("Mod", 13);
+ test.AddAttribute("fmod", 1);
+ // Due to BFloat16's precision, if the result is too small, it's not easy get pass for both CUDA and ROCm.
+ test.AddInput("X", {4}, MakeBFloat16({8.0f, 5.0f, -8.0f, 8.0f}));
+ test.AddInput("Y", {4}, MakeBFloat16({-3.4f, 8.0f, 3.4f, 5.0f}));
+ test.AddOutput("Z", {4}, MakeBFloat16({1.2f, 5.f, -1.2f, 3.f}));
+ std::vector> execution_providers;
+#ifdef USE_CUDA
+ execution_providers.push_back(DefaultCudaExecutionProvider());
+#elif USE_ROCM
+ execution_providers.push_back(DefaultRocmExecutionProvider());
+#endif
+ test.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
+}
+#endif
+
TEST(ModOpTest, Int8_mixed_sign) {
OpTester test("Mod", ModOp_ver);
test.AddInput("X", {6}, {-4, 7, 5, 4, -7, 8});