Skip to content
Prev Previous commit
Next Next commit
Added more mod tests
  • Loading branch information
edwinsolisf committed Feb 18, 2025
commit 186b30636e4b4a00a73984dee6451ea3250e6e07
6 changes: 1 addition & 5 deletions src/backend/cpu/binary.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -88,11 +88,7 @@ LOGIC_CPLX_FN(double, af_or_t, ||)

template<typename T>
static T __mod(T lhs, T rhs) {
T res = lhs % rhs; // Same as other backends

// Does this break compatibility?
// return (res < 0) ? abs(rhs - res) : res;
return res;
return lhs % rhs; // Same as other backends
}

template<typename T>
Expand Down
39 changes: 26 additions & 13 deletions test/math.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -149,31 +149,44 @@ TEST(Math, Not) {

TEST(Math, Modulus) {
af::dim4 shape(2, 2);
std::vector<long long> aData{1, 1, 1, 1};
std::vector<long long> aData{3, 3, 3, 3};
std::vector<long long> bData{2, 2, 2, 2};

auto a = af::array(shape, aData.data(), afHost);
auto b = af::array(shape, bData.data(), afHost);
auto rem = a % b;
auto diff = rem - a;

auto neg_rem = -a % b;
auto neg_diff = neg_rem + a;

ASSERT_ARRAYS_EQ(af::constant(1, shape, s64), rem);
ASSERT_ARRAYS_EQ(af::constant(0, shape, s64), diff);

ASSERT_ARRAYS_EQ(af::constant(-1, shape, s64), neg_rem);
ASSERT_ARRAYS_EQ(af::constant(0, shape, s64), neg_diff);
}

Comment thread
christophe-murphy marked this conversation as resolved.
TEST(Math, ModulusHalf) {
TEST(Math, ModulusFloat) {
SUPPORTED_TYPE_CHECK(half_float::half);
auto a = af::constant(3, {2, 2}, af::dtype::f16);
auto b = af::constant(2, {2, 2}, af::dtype::f16);
auto a32 = af::constant(3, {2, 2}, af::dtype::f32);
auto b32 = af::constant(2, {2, 2}, af::dtype::f32);
auto rem32 = a32 % b32;
af::dim4 shape(2, 2);

auto a = af::constant(3, shape, af::dtype::f16);
auto b = af::constant(2, shape, af::dtype::f16);
auto a32 = af::constant(3, shape, af::dtype::f32);
auto b32 = af::constant(2, shape, af::dtype::f32);
auto a64 = af::constant(3, shape, af::dtype::f64);
auto b64 = af::constant(2, shape, af::dtype::f64);

auto rem = a % b;
auto rem32 = a32 % b32;
auto rem64 = a64 % b64;

auto neg_rem = -a % b;
auto neg_rem32 = -a32 % b32;
auto neg_rem64 = -a64 % b64;

ASSERT_ARRAYS_EQ(af::constant(1, shape, af::dtype::f16), rem);
ASSERT_ARRAYS_EQ(af::constant(1, shape, af::dtype::f32), rem32);
ASSERT_ARRAYS_EQ(af::constant(1, shape, af::dtype::f64), rem64);

ASSERT_ARRAYS_EQ(af::constant(-1, shape, af::dtype::f16), neg_rem);
ASSERT_ARRAYS_EQ(af::constant(-1, shape, af::dtype::f32), neg_rem32);
ASSERT_ARRAYS_EQ(af::constant(-1, shape, af::dtype::f64), neg_rem64);

ASSERT_ARRAYS_EQ(rem32.as(f16), rem);
}