| 443 | self.assertEqual(out.shape, (2, 3, 1, 2)) |
| 444 | |
| 445 | def test_vmap_scatter(self): |
| 446 | def scatter(a): |
| 447 | a[mx.array(0)] = mx.array(0.0) |
| 448 | return a |
| 449 | |
| 450 | a = mx.array([[1.0, 2.0, 3.0], [2.0, 3.0, 4.0]]) |
| 451 | out = mx.vmap(scatter)(a) |
| 452 | expected = mx.array([[0.0, 2.0, 3.0], [0.0, 3.0, 4.0]]) |
| 453 | self.assertTrue(mx.allclose(out, expected)) |
| 454 | |
| 455 | out = mx.vmap(scatter, in_axes=(1,), out_axes=1)(a) |
| 456 | expected = mx.array([[0.0, 0.0, 0.0], [2.0, 3.0, 4.0]]) |
| 457 | self.assertTrue(mx.allclose(out, expected)) |
| 458 | |
| 459 | def scatter_add(a): |
| 460 | return a.at[mx.array(0)].add(mx.array(1.0)) |
| 461 | |
| 462 | a = mx.array([[1.0, 2.0, 3.0], [2.0, 3.0, 4.0]]) |
| 463 | out = mx.vmap(scatter_add)(a) |
| 464 | expected = mx.array([[2.0, 2.0, 3.0], [3.0, 3.0, 4.0]]) |
| 465 | self.assertTrue(mx.allclose(out, expected)) |
| 466 | |
| 467 | out = mx.vmap(scatter_add, in_axes=(1,), out_axes=1)(a) |
| 468 | expected = mx.array([[2.0, 3.0, 4.0], [2.0, 3.0, 4.0]]) |
| 469 | self.assertTrue(mx.allclose(out, expected)) |
| 470 | |
| 471 | # Multiple indices |
| 472 | def scatter(a): |
| 473 | a[mx.array([0, 1]), mx.array([0, 1])] = mx.array((1.0, 1.0)) |
| 474 | return a |
| 475 | |
| 476 | a = mx.zeros((3, 3, 3)) |
| 477 | |
| 478 | expected = mx.repeat(scatter(mx.zeros((3, 3)))[None], 3, axis=0) |
| 479 | out = mx.vmap(scatter, in_axes=(0,), out_axes=0)(a) |
| 480 | self.assertTrue(mx.allclose(out, expected)) |
| 481 | |
| 482 | expected = mx.zeros((3, 3, 3)) |
| 483 | expected[0, :, 0] = 1 |
| 484 | expected[1, :, 1] = 1 |
| 485 | out = mx.vmap(scatter, in_axes=(1,), out_axes=1)(a) |
| 486 | self.assertTrue(mx.allclose(out, expected)) |
| 487 | |
| 488 | expected = mx.zeros((3, 3, 3)) |
| 489 | expected[0, 0, :] = 1 |
| 490 | expected[1, 1, :] = 1 |
| 491 | out = mx.vmap(scatter, in_axes=(2,), out_axes=2)(a) |
| 492 | self.assertTrue(mx.allclose(out, expected)) |
| 493 | |
| 494 | # vmap over src and indices |
| 495 | def scatter(a, idx): |
| 496 | a[idx] = mx.array(1.0) |
| 497 | return a |
| 498 | |
| 499 | a = mx.zeros((3, 4)) |
| 500 | idx = mx.array([0, 1, 2]) |
| 501 | out = mx.vmap(scatter, in_axes=(0, 0), out_axes=0)(a, idx) |
| 502 | self.assertTrue(mx.allclose(out, mx.eye(n=3, m=4))) |