| 179 | self.assertEqual(npop(x, y).item(), mlxop(x, y).item()) |
| 180 | |
| 181 | def test_add(self): |
| 182 | x = mx.array(1) |
| 183 | y = mx.array(1) |
| 184 | z = mx.add(x, y) |
| 185 | self.assertEqual(z.item(), 2) |
| 186 | |
| 187 | x = mx.array(False, mx.bool_) |
| 188 | z = x + 1 |
| 189 | self.assertEqual(z.dtype, mx.int32) |
| 190 | self.assertEqual(z.item(), 1) |
| 191 | z = 2 + x |
| 192 | self.assertEqual(z.dtype, mx.int32) |
| 193 | self.assertEqual(z.item(), 2) |
| 194 | |
| 195 | x = mx.array(1, mx.uint32) |
| 196 | z = x + 3 |
| 197 | self.assertEqual(z.dtype, mx.uint32) |
| 198 | self.assertEqual(z.item(), 4) |
| 199 | |
| 200 | z = 3 + x |
| 201 | self.assertEqual(z.dtype, mx.uint32) |
| 202 | self.assertEqual(z.item(), 4) |
| 203 | |
| 204 | z = x + 3.0 |
| 205 | self.assertEqual(z.dtype, mx.float32) |
| 206 | self.assertEqual(z.item(), 4.0) |
| 207 | |
| 208 | z = 3.0 + x |
| 209 | self.assertEqual(z.dtype, mx.float32) |
| 210 | self.assertEqual(z.item(), 4.0) |
| 211 | |
| 212 | x = mx.array(1, mx.int64) |
| 213 | z = x + 3 |
| 214 | self.assertEqual(z.dtype, mx.int64) |
| 215 | self.assertEqual(z.item(), 4) |
| 216 | z = 3 + x |
| 217 | self.assertEqual(z.dtype, mx.int64) |
| 218 | self.assertEqual(z.item(), 4) |
| 219 | z = x + 3.0 |
| 220 | self.assertEqual(z.dtype, mx.float32) |
| 221 | self.assertEqual(z.item(), 4.0) |
| 222 | z = 3.0 + x |
| 223 | self.assertEqual(z.dtype, mx.float32) |
| 224 | self.assertEqual(z.item(), 4.0) |
| 225 | |
| 226 | x = mx.array(1, mx.float32) |
| 227 | z = x + 3 |
| 228 | self.assertEqual(z.dtype, mx.float32) |
| 229 | self.assertEqual(z.item(), 4) |
| 230 | z = 3 + x |
| 231 | self.assertEqual(z.dtype, mx.float32) |
| 232 | self.assertEqual(z.item(), 4) |
| 233 | |
| 234 | def test_subtract(self): |
| 235 | x = mx.array(4.0) |