Skip to content

Commit baf89e2

Browse files
fangchenlitadejaclaude
committed
GH-46901: [C++][Compute] Add remainder and modulo kernels
Add the `remainder`/`remainder_checked` (truncated, sign follows the dividend) and `modulo`/`modulo_checked` (floored, sign follows the divisor) scalar arithmetic kernels for integer, floating-point and decimal inputs. Decimal arguments are promoted to a common scale like `add`, but the result type is resolved by a dedicated resolver: a remainder is always smaller in magnitude than the divisor, so no extra digit is needed for a carry and the result is `precision = max(p1, p2)`, `scale = s1`. This keeps maximum-precision inputs (decimal128(38, 0), decimal256(76, 0)) from overflowing the decimal precision range. Co-authored-by: tadeja <864005+tadeja@users.noreply.github.com> Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01NpB1vVJngVjyyWvWn4AUm7
1 parent a769c29 commit baf89e2

9 files changed

Lines changed: 722 additions & 48 deletions

File tree

‎cpp/src/arrow/compute/api_scalar.cc‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -799,8 +799,10 @@ Result<Datum> RoundToMultiple(const Datum& arg, RoundToMultipleOptions options,
799799
SCALAR_ARITHMETIC_BINARY(Add, "add", "add_checked")
800800
SCALAR_ARITHMETIC_BINARY(Divide, "divide", "divide_checked")
801801
SCALAR_ARITHMETIC_BINARY(Logb, "logb", "logb_checked")
802+
SCALAR_ARITHMETIC_BINARY(Modulo, "modulo", "modulo_checked")
802803
SCALAR_ARITHMETIC_BINARY(Multiply, "multiply", "multiply_checked")
803804
SCALAR_ARITHMETIC_BINARY(Power, "power", "power_checked")
805+
SCALAR_ARITHMETIC_BINARY(Remainder, "remainder", "remainder_checked")
804806
SCALAR_ARITHMETIC_BINARY(ShiftLeft, "shift_left", "shift_left_checked")
805807
SCALAR_ARITHMETIC_BINARY(ShiftRight, "shift_right", "shift_right_checked")
806808
SCALAR_ARITHMETIC_BINARY(Subtract, "subtract", "subtract_checked")

‎cpp/src/arrow/compute/api_scalar.h‎

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -671,6 +671,42 @@ Result<Datum> Divide(const Datum& left, const Datum& right,
671671
ArithmeticOptions options = ArithmeticOptions(),
672672
ExecContext* ctx = NULLPTR);
673673

674+
/// \brief Compute the remainder (truncated division) of two values.
675+
/// Array values must be the same length. If either argument is null the result
676+
/// will be null. For integer and decimal types, if there is a zero divisor, an
677+
/// error will be raised. For floating-point types, a zero divisor yields NaN
678+
/// unless overflow checking is enabled, in which case an error is raised.
679+
///
680+
/// The result has the same sign as the dividend (C/C++ semantics).
681+
///
682+
/// \param[in] left the dividend
683+
/// \param[in] right the divisor
684+
/// \param[in] options arithmetic options (enable/disable overflow checking), optional
685+
/// \param[in] ctx the function execution context, optional
686+
/// \return the elementwise remainder
687+
ARROW_EXPORT
688+
Result<Datum> Remainder(const Datum& left, const Datum& right,
689+
ArithmeticOptions options = ArithmeticOptions(),
690+
ExecContext* ctx = NULLPTR);
691+
692+
/// \brief Compute the modulo (floored division) of two values.
693+
/// Array values must be the same length. If either argument is null the result
694+
/// will be null. For integer and decimal types, if there is a zero divisor, an
695+
/// error will be raised. For floating-point types, a zero divisor yields NaN
696+
/// unless overflow checking is enabled, in which case an error is raised.
697+
///
698+
/// The result has the same sign as the divisor (Python semantics).
699+
///
700+
/// \param[in] left the dividend
701+
/// \param[in] right the divisor
702+
/// \param[in] options arithmetic options (enable/disable overflow checking), optional
703+
/// \param[in] ctx the function execution context, optional
704+
/// \return the elementwise modulo
705+
ARROW_EXPORT
706+
Result<Datum> Modulo(const Datum& left, const Datum& right,
707+
ArithmeticOptions options = ArithmeticOptions(),
708+
ExecContext* ctx = NULLPTR);
709+
674710
/// \brief Negate values.
675711
///
676712
/// If argument is null the result will be null.

‎cpp/src/arrow/compute/kernels/base_arithmetic_internal.h‎

Lines changed: 167 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -34,6 +34,7 @@ namespace arrow {
3434

3535
using internal::AddWithOverflow;
3636
using internal::DivideWithOverflow;
37+
using internal::ModuloWithOverflow;
3738
using internal::MultiplyWithOverflow;
3839
using internal::NegateWithOverflow;
3940
using internal::SubtractWithOverflow;
@@ -468,6 +469,172 @@ struct FloatingDivideChecked {
468469
// TODO: Add decimal
469470
};
470471

472+
// Remainder (truncated): result has same sign as dividend (C/C++ semantics)
473+
struct Remainder {
474+
template <typename T, typename Arg0, typename Arg1>
475+
static enable_if_floating_value<T> Call(KernelContext*, Arg0 left, Arg1 right,
476+
Status*) {
477+
return std::fmod(left, right);
478+
}
479+
480+
template <typename T, typename Arg0, typename Arg1>
481+
static enable_if_integer_value<T> Call(KernelContext*, Arg0 left, Arg1 right,
482+
Status* st) {
483+
T result;
484+
if (ARROW_PREDICT_FALSE(ModuloWithOverflow(left, right, &result))) {
485+
if (right == 0) {
486+
*st = Status::Invalid("divide by zero");
487+
} else {
488+
// INT_MIN % -1 overflow case, result is 0
489+
result = 0;
490+
}
491+
}
492+
return result;
493+
}
494+
495+
template <typename T, typename Arg0, typename Arg1>
496+
static enable_if_decimal_value<T> Call(KernelContext*, Arg0 left, Arg1 right,
497+
Status* st) {
498+
if (right == Arg1()) {
499+
*st = Status::Invalid("divide by zero");
500+
return T();
501+
}
502+
return left % right;
503+
}
504+
};
505+
506+
struct RemainderChecked {
507+
template <typename T, typename Arg0, typename Arg1>
508+
static enable_if_floating_value<T> Call(KernelContext*, Arg0 left, Arg1 right,
509+
Status* st) {
510+
static_assert(std::is_same<T, Arg0>::value && std::is_same<T, Arg1>::value, "");
511+
if (ARROW_PREDICT_FALSE(right == 0)) {
512+
*st = Status::Invalid("divide by zero");
513+
return 0;
514+
}
515+
return std::fmod(left, right);
516+
}
517+
518+
template <typename T, typename Arg0, typename Arg1>
519+
static enable_if_integer_value<T> Call(KernelContext*, Arg0 left, Arg1 right,
520+
Status* st) {
521+
static_assert(std::is_same<T, Arg0>::value && std::is_same<T, Arg1>::value, "");
522+
T result;
523+
if (ARROW_PREDICT_FALSE(ModuloWithOverflow(left, right, &result))) {
524+
if (right == 0) {
525+
*st = Status::Invalid("divide by zero");
526+
} else {
527+
*st = Status::Invalid("overflow");
528+
}
529+
}
530+
return result;
531+
}
532+
533+
template <typename T, typename Arg0, typename Arg1>
534+
static enable_if_decimal_value<T> Call(KernelContext* ctx, Arg0 left, Arg1 right,
535+
Status* st) {
536+
return Remainder::Call<T>(ctx, left, right, st);
537+
}
538+
};
539+
540+
// Helper: Convert truncated remainder to floored modulo for signed types.
541+
// Floored modulo has the same sign as the divisor (Python semantics).
542+
template <typename T>
543+
T AdjustRemainderToFloored(T rem, T right) {
544+
if constexpr (std::is_signed_v<T>) {
545+
if ((rem > 0 && right < 0) || (rem < 0 && right > 0)) {
546+
rem += right;
547+
}
548+
}
549+
return rem;
550+
}
551+
552+
// Modulo (floored): result has same sign as divisor (Python semantics)
553+
struct Modulo {
554+
template <typename T, typename Arg0, typename Arg1>
555+
static enable_if_floating_value<T> Call(KernelContext*, Arg0 left, Arg1 right,
556+
Status*) {
557+
T rem = std::fmod(left, right);
558+
if (rem == 0) {
559+
// Preserve the sign based on divisor for zero results
560+
return std::copysign(rem, right);
561+
}
562+
return AdjustRemainderToFloored(rem, right);
563+
}
564+
565+
template <typename T, typename Arg0, typename Arg1>
566+
static enable_if_integer_value<T> Call(KernelContext*, Arg0 left, Arg1 right,
567+
Status* st) {
568+
T result;
569+
if (ARROW_PREDICT_FALSE(ModuloWithOverflow(left, right, &result))) {
570+
if (right == 0) {
571+
*st = Status::Invalid("divide by zero");
572+
} else {
573+
// INT_MIN % -1 overflow case, result is 0
574+
result = 0;
575+
}
576+
return result;
577+
}
578+
return AdjustRemainderToFloored(result, right);
579+
}
580+
581+
template <typename T, typename Arg0, typename Arg1>
582+
static enable_if_decimal_value<T> Call(KernelContext*, Arg0 left, Arg1 right,
583+
Status* st) {
584+
static const T kZero{};
585+
if (right == kZero) {
586+
*st = Status::Invalid("divide by zero");
587+
return T();
588+
}
589+
T rem = left % right;
590+
// Convert truncated to floored: adjust if signs differ
591+
if ((rem > kZero && right < kZero) || (rem < kZero && right > kZero)) {
592+
rem = rem + right;
593+
}
594+
return rem;
595+
}
596+
};
597+
598+
struct ModuloChecked {
599+
template <typename T, typename Arg0, typename Arg1>
600+
static enable_if_floating_value<T> Call(KernelContext*, Arg0 left, Arg1 right,
601+
Status* st) {
602+
static_assert(std::is_same<T, Arg0>::value && std::is_same<T, Arg1>::value, "");
603+
if (ARROW_PREDICT_FALSE(right == 0)) {
604+
*st = Status::Invalid("divide by zero");
605+
return 0;
606+
}
607+
T rem = std::fmod(left, right);
608+
if (rem == 0) {
609+
// Preserve the sign based on divisor for zero results
610+
return std::copysign(rem, right);
611+
}
612+
return AdjustRemainderToFloored(rem, right);
613+
}
614+
615+
template <typename T, typename Arg0, typename Arg1>
616+
static enable_if_integer_value<T> Call(KernelContext*, Arg0 left, Arg1 right,
617+
Status* st) {
618+
static_assert(std::is_same<T, Arg0>::value && std::is_same<T, Arg1>::value, "");
619+
T result;
620+
if (ARROW_PREDICT_FALSE(ModuloWithOverflow(left, right, &result))) {
621+
if (right == 0) {
622+
*st = Status::Invalid("divide by zero");
623+
} else {
624+
*st = Status::Invalid("overflow");
625+
}
626+
return result;
627+
}
628+
return AdjustRemainderToFloored(result, right);
629+
}
630+
631+
template <typename T, typename Arg0, typename Arg1>
632+
static enable_if_decimal_value<T> Call(KernelContext* ctx, Arg0 left, Arg1 right,
633+
Status* st) {
634+
return Modulo::Call<T>(ctx, left, right, st);
635+
}
636+
};
637+
471638
struct Negate {
472639
template <typename T, typename Arg>
473640
static constexpr enable_if_floating_value<T> Call(KernelContext*, Arg arg, Status*) {

‎cpp/src/arrow/compute/kernels/scalar_arithmetic.cc‎

Lines changed: 75 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -613,6 +613,19 @@ Result<TypeHolder> ResolveDecimalAdditionOrSubtractionOutput(
613613
});
614614
}
615615

616+
Result<TypeHolder> ResolveDecimalRemainderOutput(KernelContext*,
617+
const std::vector<TypeHolder>& types) {
618+
return ResolveDecimalBinaryOperationOutput(
619+
types,
620+
[](int32_t p1, int32_t s1, int32_t p2,
621+
int32_t s2) -> Result<std::pair<int32_t, int32_t>> {
622+
DCHECK_EQ(s1, s2);
623+
// Unlike addition, no extra digit is needed: the magnitude of a
624+
// remainder is always less than that of the divisor.
625+
return std::make_pair(std::max(p1, p2), s1);
626+
});
627+
}
628+
616629
Result<TypeHolder> ResolveDecimalMultiplicationOutput(
617630
KernelContext*, const std::vector<TypeHolder>& types) {
618631
return ResolveDecimalBinaryOperationOutput(
@@ -674,6 +687,9 @@ void AddDecimalBinaryKernels(const std::string& name, ScalarFunction* func) {
674687
if (op == "add" || op == "subtract") {
675688
out_type = OutputType(ResolveDecimalAdditionOrSubtractionOutput);
676689
constraint = DecimalsHaveSameScale();
690+
} else if (op == "remainder" || op == "modulo") {
691+
out_type = OutputType(ResolveDecimalRemainderOutput);
692+
constraint = DecimalsHaveSameScale();
677693
} else if (op == "multiply") {
678694
out_type = OutputType(ResolveDecimalMultiplicationOutput);
679695
} else if (op == "divide") {
@@ -784,7 +800,7 @@ struct ArithmeticFunction : ScalarFunction {
784800
// "add_checked" -> "add"
785801
const auto func_name = name();
786802
const std::string op = func_name.substr(0, func_name.find("_"));
787-
if (op == "add" || op == "subtract") {
803+
if (op == "add" || op == "subtract" || op == "remainder" || op == "modulo") {
788804
return CastBinaryDecimalArgs(DecimalPromotion::kAdd, types);
789805
} else if (op == "multiply") {
790806
return CastBinaryDecimalArgs(DecimalPromotion::kMultiply, types);
@@ -1173,6 +1189,42 @@ const FunctionDoc div_checked_doc{
11731189
"integer overflow is encountered."),
11741190
{"dividend", "divisor"}};
11751191

1192+
const FunctionDoc remainder_doc{
1193+
"Compute the remainder of the arguments element-wise",
1194+
("The result has the same sign as the dividend (truncated division).\n"
1195+
"This is equivalent to the C/C++ '%' operator.\n"
1196+
"Integer and decimal division by zero returns an error, while\n"
1197+
"floating-point division by zero returns NaN.\n"
1198+
"Use function \"remainder_checked\" if you want to get an error\n"
1199+
"in all the aforementioned cases."),
1200+
{"dividend", "divisor"}};
1201+
1202+
const FunctionDoc remainder_checked_doc{
1203+
"Compute the remainder of the arguments element-wise",
1204+
("The result has the same sign as the dividend (truncated division).\n"
1205+
"This is equivalent to the C/C++ '%' operator.\n"
1206+
"An error is returned when trying to divide by zero, or when\n"
1207+
"integer overflow is encountered."),
1208+
{"dividend", "divisor"}};
1209+
1210+
const FunctionDoc modulo_doc{
1211+
"Compute the modulo of the arguments element-wise",
1212+
("The result has the same sign as the divisor (floored division).\n"
1213+
"This is equivalent to Python's '%' operator.\n"
1214+
"Integer and decimal division by zero returns an error, while\n"
1215+
"floating-point division by zero returns NaN.\n"
1216+
"Use function \"modulo_checked\" if you want to get an error\n"
1217+
"in all the aforementioned cases."),
1218+
{"dividend", "divisor"}};
1219+
1220+
const FunctionDoc modulo_checked_doc{
1221+
"Compute the modulo of the arguments element-wise",
1222+
("The result has the same sign as the divisor (floored division).\n"
1223+
"This is equivalent to Python's '%' operator.\n"
1224+
"An error is returned when trying to divide by zero, or when\n"
1225+
"integer overflow is encountered."),
1226+
{"dividend", "divisor"}};
1227+
11761228
const FunctionDoc negate_doc{"Negate the argument element-wise",
11771229
("Results will wrap around on integer overflow.\n"
11781230
"Use function \"negate_checked\" if you want overflow\n"
@@ -1724,6 +1776,28 @@ void RegisterScalarArithmetic(FunctionRegistry* registry) {
17241776

17251777
DCHECK_OK(registry->AddFunction(std::move(divide_checked)));
17261778

1779+
// ----------------------------------------------------------------------
1780+
auto remainder = MakeArithmeticFunctionNotNull<Remainder>("remainder", remainder_doc);
1781+
AddDecimalBinaryKernels<Remainder>("remainder", remainder.get());
1782+
DCHECK_OK(registry->AddFunction(std::move(remainder)));
1783+
1784+
// ----------------------------------------------------------------------
1785+
auto remainder_checked = MakeArithmeticFunctionNotNull<RemainderChecked>(
1786+
"remainder_checked", remainder_checked_doc);
1787+
AddDecimalBinaryKernels<RemainderChecked>("remainder_checked", remainder_checked.get());
1788+
DCHECK_OK(registry->AddFunction(std::move(remainder_checked)));
1789+
1790+
// ----------------------------------------------------------------------
1791+
auto modulo = MakeArithmeticFunctionNotNull<Modulo>("modulo", modulo_doc);
1792+
AddDecimalBinaryKernels<Modulo>("modulo", modulo.get());
1793+
DCHECK_OK(registry->AddFunction(std::move(modulo)));
1794+
1795+
// ----------------------------------------------------------------------
1796+
auto modulo_checked =
1797+
MakeArithmeticFunctionNotNull<ModuloChecked>("modulo_checked", modulo_checked_doc);
1798+
AddDecimalBinaryKernels<ModuloChecked>("modulo_checked", modulo_checked.get());
1799+
DCHECK_OK(registry->AddFunction(std::move(modulo_checked)));
1800+
17271801
// ----------------------------------------------------------------------
17281802
auto negate = MakeUnaryArithmeticFunction<Negate>("negate", negate_doc);
17291803
AddDecimalUnaryKernels<Negate>(negate.get());

0 commit comments

Comments
 (0)