Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 46 additions & 7 deletions jax/_src/state/discharge.py
Original file line number Diff line number Diff line change
Expand Up @@ -511,12 +511,14 @@ def transform_array(x, transforms):
case BitcastTransform():
result = bitcast(result, transform.dtype)
case ReshapeTransform():
result = result.reshape(transform.shape)
result = lax.reshape(result, transform.shape)
case _:
raise NotImplementedError(f"Unsupported transform: {transform}")
return result

def transform_swap_array(x, transforms, val):
from jax._src.numpy import lax_numpy # pyrefly: ignore[missing-import]

if transforms is None:
transforms = []

Expand Down Expand Up @@ -555,9 +557,11 @@ def transform_swap_array(x, transforms, val):
# was indexed into.
intermediates.append(new_val)
case BitcastTransform():
intermediates.append(bitcast(new_val, transform.dtype))
new_val = bitcast(new_val, transform.dtype)
intermediates.append(new_val)
case ReshapeTransform():
intermediates.append(new_val.reshape(transform.shape))
new_val = lax.reshape(new_val, transform.shape)
intermediates.append(new_val)
case _:
raise NotImplementedError(f"Unsupported transform: {transform}")

Expand All @@ -583,10 +587,17 @@ def transform_swap_array(x, transforms, val):
intermediate, indexer, transpose_order
)
arrays = _convert_to_gather_arrays(indexer)
new_x = intermediate.at[arrays].set(new_x)
# `asarray` ensures `intermediate` has an `.at` attribute; it may be a
# plain value (e.g. after a reshape/bitcast reverse) rather than a
# jax array.
new_x = lax_numpy.asarray(intermediate).at[arrays].set(new_x)
if transpose_order is not None:
transpose_order_inversed = np.argsort(transpose_order)
new_x = new_x.transpose(transpose_order_inversed)
elif isinstance(transform, ReshapeTransform):
new_x = lax.reshape(new_x, np.shape(intermediate))
elif isinstance(transform, BitcastTransform):
new_x = bitcast(new_x, intermediate.dtype)
else:
raise NotImplementedError(f"Unsupported transform: {transform}")

Expand Down Expand Up @@ -628,12 +639,38 @@ def _optimization_barrier_discharge_rule(
return new_invals, [o for o, r in zip(outs, is_ref) if not r]

def _addupdate_discharge(x, val, idx, tree):
transforms = tree_util.tree_unflatten(tree, idx)
from jax._src.numpy import lax_numpy # pyrefly: ignore[missing-import]

transforms = list(tree_util.tree_unflatten(tree, idx))
if any(isinstance(t, BitcastTransform) for t in transforms):
raise NotImplementedError(
"`addupdate` (`+=`) is not supported on bitcast views. Use explicit"
" read-modify-write (`ref.bitcast(...)[...] = ...`) or `.swap(...)`"
" instead."
)

if transforms and isinstance(transforms[-1], ReshapeTransform):
broadcast_shape = transforms[-1].shape
while transforms and isinstance(transforms[-1], ReshapeTransform):
transforms.pop()
target_shape = (
transforms[-1].get_indexer_shape()
if transforms and isinstance(transforms[-1], indexing.NDIndexer)
else x.shape
)
val = lax_numpy.broadcast_to(val, broadcast_shape).reshape(target_shape)
if not transforms:
return x + val
if len(transforms) > 1:
raise NotImplementedError("Only single indexer is supported.")
raise NotImplementedError(
"`addupdate` does not support combining an indexer with other"
f" transforms (e.g. indexed reshape views); got {transforms}."
)
indexer = transforms[0]
if not isinstance(indexer, indexing.NDIndexer):
raise NotImplementedError(
f"Unsupported transform for `addupdate`: {indexer}"
)

if _is_trivial_indexer(indexer):
return x + val
Expand All @@ -652,7 +689,9 @@ def _addupdate_discharge(x, val, idx, tree):
if transpose_order is not None:
x, indexer = _perform_transpose_before_gather(x, indexer, transpose_order)
arrays = _convert_to_gather_arrays(indexer)
x = x.at[arrays].add(val)
# `asarray` ensures `x` has an `.at` attribute; it may be a plain value
# rather than a jax array.
x = lax_numpy.asarray(x).at[arrays].add(val)
if transpose_order is not None:
transpose_order_inversed = np.argsort(transpose_order)
x = x.transpose(transpose_order_inversed)
Expand Down
29 changes: 29 additions & 0 deletions tests/pallas/pallas_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -300,6 +300,20 @@ def kernel(o_ref):
self.assertEqual(o_ref_shape, (128,))
self.assertAllClose(pids, jnp.zeros(o_ref_shape, dtype=np.int32))

def test_store_reshaped_ref(self):
shape1, shape2 = (4, 3), (2, 6)

@functools.partial(
self.pallas_call,
out_shape=jax.ShapeDtypeStruct(shape1, jnp.float32),
)
def kernel(x_ref, o_ref):
o_ref.reshape(shape2)[...] = x_ref.reshape(shape2)[...]

x = jnp.arange(12, dtype=jnp.float32).reshape(shape1)
y = kernel(x)
np.testing.assert_array_equal(y, x)


class PallasTritonTest(PallasTest):

Expand All @@ -316,6 +330,12 @@ def pallas_call(self, *args, **kwargs):
*args, compiler_params=pltriton.CompilerParams(), **kwargs
)

def test_store_reshaped_ref(self):
with self.assertRaisesRegex(
AttributeError, "object has no attribute 'get_indexer_shape_static'"
):
super().test_store_reshaped_ref()

def test_array_indexing(self):
x = jnp.arange(128, dtype=floatx)

Expand All @@ -342,6 +362,11 @@ def setUp(self):
self.skipTest("Pallas TPU is not available")
super().setUp()

def test_store_reshaped_ref(self):
if not self.INTERPRET:
self.skipTest("Requires TPU sublane tiling (tested in tpu_pallas_test.py)")
super().test_store_reshaped_ref()


class PallasMGPUTest(PallasTest):

Expand All @@ -368,6 +393,10 @@ def skip_if_x64(self):
if floatx == jnp.float64:
self.skipTest("Mosaic GPU does not support float64.")

def test_store_reshaped_ref(self):
self.skip_if_x64()
super().test_store_reshaped_ref()

def test_add_one(self):
self.skip_if_x64()
super().test_add_one()
Expand Down
137 changes: 137 additions & 0 deletions tests/state_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -2135,5 +2135,142 @@ def f(x):
np.testing.assert_array_equal(grad, 2.0 * x)


class RefTransformTest(jtu.JaxTestCase):

def test_swap_reshape(self):
@jax.jit
def f():
x = jnp.arange(12, dtype=jnp.float32).reshape((4, 3))
ref = jax.new_ref(x)
old_val = ref.reshape(6, 2).swap(jnp.zeros((6, 2), dtype=jnp.float32))
return old_val, ref[...]
old_val, new_val = f()
self.assertEqual(old_val.shape, (6, 2))
self.assertAllClose(old_val, jnp.arange(12, dtype=jnp.float32).reshape((6, 2)))
self.assertEqual(new_val.shape, (4, 3))
self.assertAllClose(new_val, jnp.zeros((4, 3), dtype=jnp.float32))

def test_set_reshape(self):
@jax.jit
def f():
x = jnp.zeros((4, 3))
ref = jax.new_ref(x)
ref.reshape(2, 6)[...] = jnp.ones((2, 6))
return ref[...]
self.assertAllClose(f(), jnp.ones((4, 3)))

def test_bitcast_swap_and_set(self):
@jax.jit
def f():
ref = jax.new_ref(jnp.array([[1.0, 2.0], [3.0, 4.0]], dtype=jnp.float32))
old_val = ref.bitcast(jnp.int32).swap(jnp.zeros((2, 2), dtype=jnp.int32))
ref.bitcast(jnp.int32)[...] = jnp.ones((2, 2), dtype=jnp.int32)
return old_val, ref[...]
old_val, new_val = f()
self.assertEqual(old_val.dtype, jnp.int32)
self.assertArraysEqual(
old_val,
lax.bitcast_convert_type(
jnp.array([[1.0, 2.0], [3.0, 4.0]], dtype=jnp.float32), jnp.int32
),
)
self.assertEqual(new_val.dtype, jnp.float32)
self.assertArraysEqual(
new_val,
lax.bitcast_convert_type(jnp.ones((2, 2), dtype=jnp.int32), jnp.float32),
)

def test_set_reshaped_slice(self):
@jax.jit
def f():
x = jnp.zeros((4, 3), dtype=jnp.float32)
ref = jax.new_ref(x)
ref.reshape(2, 6)[0, :3] = jnp.ones(3, dtype=jnp.float32)
return ref[...]
expected = jnp.array([[1.0, 1.0, 1.0], [0.0, 0.0, 0.0], [0.0, 0.0, 0.0], [0.0, 0.0, 0.0]], dtype=jnp.float32)
self.assertAllClose(f(), expected)

def test_swap_reshaped_slice(self):
@jax.jit
def f():
x = jnp.arange(12, dtype=jnp.float32).reshape((4, 3))
ref = jax.new_ref(x)
old_val = ref.reshape(2, 6).at[0, :3].swap(jnp.zeros(3, dtype=jnp.float32))
return old_val, ref[...]
old_val, new_val = f()
self.assertEqual(old_val.shape, (3,))
self.assertAllClose(old_val, jnp.arange(3, dtype=jnp.float32))
expected = jnp.arange(12, dtype=jnp.float32).reshape((4, 3))
expected = expected.reshape(2, 6).at[0, :3].set(0.0).reshape(4, 3)
self.assertAllClose(new_val, expected)

def test_bitcast_slice_set(self):
@jax.jit
def f():
ref = jax.new_ref(jnp.array([[1.0, 2.0], [3.0, 4.0]], dtype=jnp.float32))
ref.bitcast(jnp.int32)[0, :] = jnp.zeros(2, dtype=jnp.int32)
return ref[...]
expected = jnp.array([[0.0, 0.0], [3.0, 4.0]], dtype=jnp.float32)
self.assertArraysEqual(f(), expected)

def test_addupdate_on_reshaped_view(self):
@jax.jit
def f():
ref = jax.new_ref(jnp.zeros((4, 3)))
ref_addupdate(ref.reshape(2, 6), (), jnp.ones((2, 6)))
return ref[...]

expected = jnp.ones((4, 3))
self.assertAllClose(f(), expected)

def test_addupdate_on_reshaped_slice_raises_error(self):
@jax.jit
def f():
ref = jax.new_ref(jnp.zeros((4, 3), dtype=jnp.float32))
ref_addupdate(ref.reshape(2, 6), (0, slice(0, 3)), jnp.ones(3, dtype=jnp.float32))
return ref[...]

with self.assertRaisesRegex(
NotImplementedError, "does not support combining an indexer"
):
f()

def test_addupdate_on_bitcast_raises_error(self):
@jax.jit
def f():
ref = jax.new_ref(jnp.zeros((2, 2), dtype=jnp.float32))
ref_addupdate(ref.bitcast(jnp.int32), (), jnp.ones((2, 2), dtype=jnp.int32))
return ref[...]

with self.assertRaisesRegex(
NotImplementedError, "not supported on bitcast views"
):
f()

def test_addupdate_on_unsupported_transform_raises_error(self):
@jax.jit
def f():
ref = jax.new_ref(jnp.zeros((2, 3), dtype=jnp.float32))
view = state_types.TransformedRef(
ref, (state_types.TransposeTransform((1, 0)),)
)
ref_addupdate(view, (), jnp.ones((3, 2), dtype=jnp.float32))
return ref[...]

with self.assertRaisesRegex(
NotImplementedError, "Unsupported transform for `addupdate`"
):
f()

def test_scalar_ref_reshape(self):
@jax.jit
def f():
ref = jax.new_ref(jnp.array(5.0))
return ref.reshape(1)[...]

self.assertEqual(f().shape, (1,))
self.assertAllClose(f(), jnp.array([5.0]))


if __name__ == '__main__':
absltest.main(testLoader=jtu.JaxTestLoader())
Loading