Slice update with operation (#3266)
This commit is contained in:
+104
-31
@@ -300,65 +300,138 @@ class TestAutograd(mlx_tests.MLXTestCase):
|
||||
x[idx] = 2.0
|
||||
return x.sum()
|
||||
|
||||
dfdx = mx.grad(fun)(mx.array([1.0, 2.0, 3.0]), mx.array([1]))
|
||||
self.assertTrue(mx.array_equal(dfdx, mx.array([1.0, 0.0, 1.0])))
|
||||
dfdx = mx.grad(fun)(mx.array([1.0, 2.0, 3.0, 4.0]), mx.array([1, 3]))
|
||||
self.assertTrue(mx.array_equal(dfdx, mx.array([1.0, 0.0, 1.0, 0.0])))
|
||||
self.assertEqual(dfdx.dtype, mx.float32)
|
||||
|
||||
y = mx.array([0.0, 1.0, 2.0])
|
||||
y = mx.array([0.0, 1.0, 2.0, 3.0])
|
||||
|
||||
def fun(x, idx):
|
||||
y[idx] = x
|
||||
return y.sum()
|
||||
|
||||
dfdx = mx.grad(fun)(mx.array([2.0]), mx.array([1]))
|
||||
self.assertTrue(mx.array_equal(dfdx, mx.array([1.0])))
|
||||
dfdx = mx.grad(fun)(mx.array([2.0, 3.0]), mx.array([1, 3]))
|
||||
self.assertTrue(mx.array_equal(dfdx, mx.array([1.0, 1.0])))
|
||||
self.assertEqual(dfdx.dtype, mx.float32)
|
||||
|
||||
def test_scatter_add_vjp(self):
|
||||
def fun(src, updates):
|
||||
x = src.at[mx.array([1, 3])].add(updates)
|
||||
return x
|
||||
|
||||
cotan = mx.array([4.0, 5.0, 6.0, 7.0])
|
||||
updates = mx.array([1.0, 2.0])
|
||||
_, vjps = mx.vjp(fun, [mx.array([1.0, 2.0, 3.0, 4.0]), updates], [cotan])
|
||||
mx.eval(vjps)
|
||||
|
||||
self.assertTrue(mx.allclose(vjps[0], mx.array([4.0, 5.0, 6.0, 7.0])))
|
||||
self.assertTrue(mx.allclose(vjps[1], mx.array([5.0, 7.0])))
|
||||
|
||||
def test_scatter_max_vjp(self):
|
||||
def fun(src, updates):
|
||||
x = src.at[1].maximum(updates)
|
||||
x = src.at[mx.array([1, 3])].maximum(updates)
|
||||
return x
|
||||
|
||||
cotan = mx.array([4.0, 5.0, 6.0])
|
||||
_, vjps = mx.vjp(fun, [mx.array([1.0, 2.0, 3.0]), mx.array([[3.0]])], [cotan])
|
||||
cotan = mx.array([4.0, 5.0, 6.0, 7.0])
|
||||
updates = mx.array([1.0, 2.0])
|
||||
_, vjps = mx.vjp(fun, [mx.array([1.0, 2.0, 3.0, 4.0]), updates], [cotan])
|
||||
mx.eval(vjps)
|
||||
|
||||
# Update larger than value
|
||||
self.assertTrue(mx.allclose(vjps[0], mx.array([4.0, 0.0, 6.0])))
|
||||
self.assertTrue(mx.allclose(vjps[1], mx.array([5.0])))
|
||||
self.assertTrue(mx.allclose(vjps[0], mx.array([4.0, 5.0, 6.0, 7.0])))
|
||||
self.assertTrue(mx.allclose(vjps[1], mx.array([0.0, 0.0])))
|
||||
|
||||
cotan = mx.array([[4.0], [5.0], [6.0]])
|
||||
_, vjps = mx.vjp(
|
||||
fun, [mx.array([[1.0], [2.0], [3.0]]), mx.array([[[2.0]]])], [cotan]
|
||||
)
|
||||
updates = mx.array([5.0, 6.0])
|
||||
_, vjps = mx.vjp(fun, [mx.array([1.0, 2.0, 3.0, 4.0]), updates], [cotan])
|
||||
mx.eval(vjps)
|
||||
|
||||
# Update and value are equal
|
||||
self.assertTrue(mx.allclose(vjps[0], mx.array([[4.0], [5.0], [6.0]])))
|
||||
self.assertTrue(mx.allclose(vjps[1], mx.array([[[5.0]]])))
|
||||
self.assertTrue(mx.allclose(vjps[0], mx.array([4.0, 0.0, 6.0, 0.0])))
|
||||
self.assertTrue(mx.allclose(vjps[1], mx.array([5.0, 7.0])))
|
||||
|
||||
def test_scatter_min_vjp(self):
|
||||
def fun(src, updates):
|
||||
x = src.at[1].minimum(updates)
|
||||
x = src.at[mx.array([1, 3])].minimum(updates)
|
||||
return x
|
||||
|
||||
cotan = mx.array([4.0, 5.0, 6.0])
|
||||
_, vjps = mx.vjp(fun, [mx.array([1.0, 2.0, 3.0]), mx.array([[3.0]])], [cotan])
|
||||
cotan = mx.array([4.0, 5.0, 6.0, 7.0])
|
||||
updates = mx.array([5.0, 6.0])
|
||||
_, vjps = mx.vjp(fun, [mx.array([1.0, 2.0, 3.0, 4.0]), updates], [cotan])
|
||||
mx.eval(vjps)
|
||||
|
||||
# Update larger than value
|
||||
self.assertTrue(mx.allclose(vjps[0], mx.array([4.0, 5.0, 6.0])))
|
||||
self.assertTrue(mx.allclose(vjps[1], mx.array([0.0])))
|
||||
self.assertTrue(mx.allclose(vjps[0], mx.array([4.0, 5.0, 6.0, 7.0])))
|
||||
self.assertTrue(mx.allclose(vjps[1], mx.array([0.0, 0.0])))
|
||||
|
||||
cotan = mx.array([[4.0], [5.0], [6.0]])
|
||||
_, vjps = mx.vjp(
|
||||
fun, [mx.array([[1.0], [2.0], [3.0]]), mx.array([[[2.0]]])], [cotan]
|
||||
)
|
||||
updates = mx.array([1.0, 1.0])
|
||||
_, vjps = mx.vjp(fun, [mx.array([1.0, 2.0, 3.0, 4.0]), updates], [cotan])
|
||||
mx.eval(vjps)
|
||||
|
||||
# Update and value are equal
|
||||
self.assertTrue(mx.allclose(vjps[0], mx.array([[4.0], [5.0], [6.0]])))
|
||||
self.assertTrue(mx.allclose(vjps[1], mx.array([[[5.0]]])))
|
||||
self.assertTrue(mx.allclose(vjps[0], mx.array([4.0, 0.0, 6.0, 0.0])))
|
||||
self.assertTrue(mx.allclose(vjps[1], mx.array([5.0, 7.0])))
|
||||
|
||||
def test_slice_update_max_vjp(self):
|
||||
def fun(src, updates):
|
||||
x = src.at[1:3].maximum(updates)
|
||||
return x
|
||||
|
||||
cotan = mx.array([4.0, 5.0, 6.0, 7.0])
|
||||
updates = mx.array([[1.0, 2.0]])
|
||||
_, vjps = mx.vjp(fun, [mx.array([1.0, 2.0, 3.0, 4.0]), updates], [cotan])
|
||||
mx.eval(vjps)
|
||||
|
||||
self.assertTrue(mx.allclose(vjps[0], mx.array([4.0, 5.0, 6.0, 7.0])))
|
||||
self.assertTrue(mx.allclose(vjps[1], mx.array([[0.0, 0.0]])))
|
||||
|
||||
updates = mx.array([[5.0, 6.0]])
|
||||
_, vjps = mx.vjp(fun, [mx.array([1.0, 2.0, 3.0, 4.0]), updates], [cotan])
|
||||
mx.eval(vjps)
|
||||
|
||||
self.assertTrue(mx.allclose(vjps[0], mx.array([4.0, 0.0, 0.0, 7.0])))
|
||||
self.assertTrue(mx.allclose(vjps[1], mx.array([[5.0, 6.0]])))
|
||||
|
||||
def test_slice_update_min_vjp(self):
|
||||
def fun(src, updates):
|
||||
x = src.at[1:3].minimum(updates)
|
||||
return x
|
||||
|
||||
cotan = mx.array([4.0, 5.0, 6.0, 7.0])
|
||||
updates = mx.array([[5.0, 6.0]])
|
||||
_, vjps = mx.vjp(fun, [mx.array([1.0, 2.0, 3.0, 4.0]), updates], [cotan])
|
||||
mx.eval(vjps)
|
||||
|
||||
self.assertTrue(mx.allclose(vjps[0], mx.array([4.0, 5.0, 6.0, 7.0])))
|
||||
self.assertTrue(mx.allclose(vjps[1], mx.array([[0.0, 0.0]])))
|
||||
|
||||
updates = mx.array([[1.0, 1.0]])
|
||||
_, vjps = mx.vjp(fun, [mx.array([1.0, 2.0, 3.0, 4.0]), updates], [cotan])
|
||||
mx.eval(vjps)
|
||||
|
||||
self.assertTrue(mx.allclose(vjps[0], mx.array([4.0, 0.0, 0.0, 7.0])))
|
||||
self.assertTrue(mx.allclose(vjps[1], mx.array([[5.0, 6.0]])))
|
||||
|
||||
def test_slice_update_add_vjp(self):
|
||||
def fun(src, updates):
|
||||
x = src.at[1:3].add(updates)
|
||||
return x
|
||||
|
||||
cotan = mx.array([4.0, 5.0, 6.0, 7.0])
|
||||
updates = mx.array([[1.0, 2.0]])
|
||||
_, vjps = mx.vjp(fun, [mx.array([1.0, 2.0, 3.0, 4.0]), updates], [cotan])
|
||||
mx.eval(vjps)
|
||||
|
||||
self.assertTrue(mx.allclose(vjps[0], mx.array([4.0, 5.0, 6.0, 7.0])))
|
||||
self.assertTrue(mx.allclose(vjps[1], mx.array([[5.0, 6.0]])))
|
||||
|
||||
def test_slice_update_multiply_vjp(self):
|
||||
def fun(src, updates):
|
||||
x = src.at[1:3].multiply(updates)
|
||||
return x
|
||||
|
||||
cotan = mx.array([4.0, 5.0, 6.0, 7.0])
|
||||
updates = mx.array([[2.0, 3.0]])
|
||||
_, vjps = mx.vjp(fun, [mx.array([1.0, 2.0, 3.0, 4.0]), updates], [cotan])
|
||||
mx.eval(vjps)
|
||||
|
||||
self.assertTrue(mx.allclose(vjps[0], mx.array([4.0, 10.0, 18.0, 7.0])))
|
||||
self.assertTrue(mx.allclose(vjps[1], mx.array([[10.0, 18.0]])))
|
||||
|
||||
def test_split_against_slice(self):
|
||||
def f_split(x):
|
||||
|
||||
Reference in New Issue
Block a user