From fad278d5c7ac3729db9c0b26f24f79bf9bcd6fa9 Mon Sep 17 00:00:00 2001 From: Sylvester Kaczmarek <16242628+sylvesterkaczmarek@users.noreply.github.com> Date: Thu, 10 Sep 2026 19:26:15 +0100 Subject: [PATCH] Honor output axes in fake pmap --- chex/_src/fake.py | 14 ++++++++++++-- chex/_src/fake_test.py | 38 ++++++++++++++++++++++++++++++++++++++ 2 files changed, 50 insertions(+), 2 deletions(-) diff --git a/chex/_src/fake.py b/chex/_src/fake.py index a248e26a..97c4f555 100644 --- a/chex/_src/fake.py +++ b/chex/_src/fake.py @@ -125,6 +125,7 @@ def _fake_pmap(fn, axis_name: Optional[Any] = None, *, in_axes=0, + out_axes=0, static_broadcasted_argnums: Union[int, Iterable[int]] = (), jit_result: bool = False, fake_parallel_axis: bool = False, @@ -184,7 +185,10 @@ def fn_without_statics(*args): fn_without_statics = fn vmapped_fn = jax.vmap( - fn_without_statics, in_axes=vmap_in_axes, axis_name=axis_name + fn_without_statics, + in_axes=vmap_in_axes, + out_axes=out_axes, + axis_name=axis_name, ) if jit_result: vmapped_fn = jax.jit(vmapped_fn) @@ -196,7 +200,13 @@ def fn_without_statics(*args): output = vmapped_fn(*call_args) if fake_parallel_axis: - output = jax.tree_util.tree_map(lambda x: jnp.squeeze(x, axis=0), output) + def squeeze_fake_axis(axes, value): + if axes is None: + return value + return jax.tree_util.tree_map(lambda x: jnp.squeeze(x, axis=axes), value) + + output = jax.tree_util.tree_map( + squeeze_fake_axis, out_axes, output, is_leaf=lambda x: x is None) return output diff --git a/chex/_src/fake_test.py b/chex/_src/fake_test.py index 17b42320..ce0d2373 100644 --- a/chex/_src/fake_test.py +++ b/chex/_src/fake_test.py @@ -149,6 +149,44 @@ def foo(x): _assert_pmapped(foo, fn_input, is_pmapped, jit_result) ctx.stop() + @parameterized.product(out_axes=(0, 1, 2), jit_result=(False, True)) + def test_fake_pmap_out_axes(self, out_axes, jit_result): + num_devices = len(jax.devices()) + inputs = jnp.arange(num_devices * 6).reshape(num_devices, 2, 3) + fn = lambda x: x * 2 + expected = jax.pmap(fn, out_axes=out_axes)(inputs) + with fake.fake_pmap(jit_result=jit_result): + actual = jax.pmap(fn, out_axes=out_axes)(inputs) + asserts.assert_trees_all_equal(actual, expected) + + @parameterized.parameters(False, True) + def test_fake_pmap_out_axes_tree(self, jit_result): + num_devices = len(jax.devices()) + inputs = jnp.arange(num_devices * 6).reshape(num_devices, 2, 3) + + def fn(x): + return {'mapped': x * 2, 'constant': jnp.array(7), 'nested': (x, x + 1)} + + out_axes = {'mapped': 1, 'constant': None, 'nested': 2} + expected = jax.pmap(fn, out_axes=out_axes)(inputs) + with fake.fake_pmap(jit_result=jit_result): + actual = jax.pmap(fn, out_axes=out_axes)(inputs) + asserts.assert_trees_all_equal(actual, expected) + + @parameterized.product( + out_axes=(0, 1, 2, {'mapped': 1, 'constant': None, 'nested': 2}), + jit_result=(False, True), + ) + def test_fake_parallel_axis_out_axes(self, out_axes, jit_result): + inputs = jnp.arange(6).reshape(2, 3) + + def fn(x): + return {'mapped': x * 2, 'constant': jnp.ones((2, 3)), 'nested': (x, x + 1)} + + with fake.fake_pmap(fake_parallel_axis=True, jit_result=jit_result): + actual = jax.pmap(fn, out_axes=out_axes)(inputs) + asserts.assert_trees_all_equal(actual, fn(inputs)) + def test_fake_pmap_axis_name(self): with fake.fake_pmap():