feat: add logicalAnd and logicalOR (#386)
* feat: add logicalAnd and logicalOR * run pre-commit * Refactor logical_and and logical_or functions * Add acknowledgement * Add logical AND and logical OR operators * Refactor logical_and and logical_or functions * Add support for logical operators on bool arrays * Update mlx/ops.cpp Co-authored-by: Awni Hannun <[email protected]> * Update mlx/ops.cpp Co-authored-by: Awni Hannun <[email protected]> * Add logical AND and OR operators for arrays and scalars * Refactor vjp and jvp methods in primitives.cpp * Add overloaded operators for logical AND and OR * format --------- Co-authored-by: Awni Hannun <[email protected]> Co-authored-by: Awni Hannun <[email protected]>
This commit is contained in:
co-authored by
Awni Hannun
Awni Hannun
parent
022a944367
commit
73321b8097
@@ -610,6 +610,28 @@ class TestOps(mlx_tests.MLXTestCase):
|
||||
expected = np.logical_not(a)
|
||||
self.assertTrue(np.array_equal(result, expected))
|
||||
|
||||
def test_logical_and(self):
|
||||
a = mx.array([True, False, True, False])
|
||||
b = mx.array([True, True, False, False])
|
||||
result = mx.logical_and(a, b)
|
||||
expected = np.logical_and(a, b)
|
||||
self.assertTrue(np.array_equal(result, expected))
|
||||
|
||||
# test overloaded operator
|
||||
result = a & b
|
||||
self.assertTrue(np.array_equal(result, expected))
|
||||
|
||||
def test_logical_or(self):
|
||||
a = mx.array([True, False, True, False])
|
||||
b = mx.array([True, True, False, False])
|
||||
result = mx.logical_or(a, b)
|
||||
expected = np.logical_or(a, b)
|
||||
self.assertTrue(np.array_equal(result, expected))
|
||||
|
||||
# test overloaded operator
|
||||
result = a | b
|
||||
self.assertTrue(np.array_equal(result, expected))
|
||||
|
||||
def test_square(self):
|
||||
a = mx.array([0.1, 0.5, 1.0, 10.0])
|
||||
result = mx.square(a)
|
||||
|
||||
Reference in New Issue
Block a user