From 2210ff90323ceb72fa04f0b1addedee182f9b0d8 Mon Sep 17 00:00:00 2001 From: RecML authors Date: Tue, 8 Sep 2026 15:10:59 -0700 Subject: [PATCH] Internal change PiperOrigin-RevId: 978138333 --- recml/core/ops/hstu_ops.py | 23 ++-- recml/core/utils/distribution_utils.py | 122 +++++++++++++++++ recml/core/utils/distribution_utils_test.py | 137 ++++++++++++++++++++ 3 files changed, 267 insertions(+), 15 deletions(-) create mode 100644 recml/core/utils/distribution_utils.py create mode 100644 recml/core/utils/distribution_utils_test.py diff --git a/recml/core/ops/hstu_ops.py b/recml/core/ops/hstu_ops.py index bc8c901..7f17a73 100644 --- a/recml/core/ops/hstu_ops.py +++ b/recml/core/ops/hstu_ops.py @@ -26,8 +26,8 @@ from jax.experimental.pallas.ops.tpu.splash_attention import splash_attention_mask_info import jax.numpy as jnp import jaxtyping as jt -import keras import numpy as np +from recml.core.utils import distribution_utils NUM_LANES = 128 @@ -835,8 +835,8 @@ def pointwise_splash_attention( q_segment_ids: jax.Array | None = None, kv_segment_ids: jax.Array | None = None, sliding_window_size: int | None = None, - qkv_axis_names: tuple[str | None, ...] = (), - segment_id_axis_names: tuple[str | None, ...] = (), + qkv_axis_names: tuple[distribution_utils.AxisSpec, ...] = (), + segment_id_axis_names: tuple[distribution_utils.AxisSpec, ...] = (), scale: float | None = None, block_q: int | None = None, block_kv: int | None = None, @@ -978,16 +978,7 @@ def pointwise_splash_attention( " provided." ) - if (global_abstract_mesh := jax.sharding.get_abstract_mesh()).shape_tuple: - abstract_mesh = global_abstract_mesh - elif (distribution := keras.distribution.distribution()) is not None: - device_mesh: keras.distribution.DeviceMesh = distribution.device_mesh - abstract_mesh = jax.sharding.AbstractMesh( - axis_sizes=tuple(device_mesh.shape), - axis_names=tuple(device_mesh.axis_names), - ) - else: - abstract_mesh = None + abstract_mesh = distribution_utils.get_abstract_mesh() def _kernel_wrapper( query: jax.Array, @@ -1016,16 +1007,18 @@ def _kernel_wrapper( )(query, key, value, segment_ids=segment_ids) if abstract_mesh is not None and abstract_mesh.shape_tuple: + batch_axis = distribution_utils.resolve_batch_axis(abstract_mesh) + if not qkv_axis_names: qkv_axis_names = ( - abstract_mesh.axis_names[0], # batch dimension + batch_axis, # batch dimension None, # heads dimension None, # length dimension None, # hidden dimension ) if not segment_id_axis_names: segment_id_axis_names = ( - abstract_mesh.axis_names[0], # batch dimension + batch_axis, # batch dimension None, # length dimension ) diff --git a/recml/core/utils/distribution_utils.py b/recml/core/utils/distribution_utils.py new file mode 100644 index 0000000..e079588 --- /dev/null +++ b/recml/core/utils/distribution_utils.py @@ -0,0 +1,122 @@ +# Copyright 2024 RecML authors . +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Utilities for inspecting the active JAX/Keras distribution mesh.""" + +import jax +import keras + +# Conventional name of the mesh axis used for model/tensor parallelism. Any +# other axis of the mesh is assumed to shard the batch. +MODEL_AXIS_NAME = 'model' + +AxisSpec = str | tuple[str, ...] | None + + +def get_abstract_mesh() -> jax.sharding.AbstractMesh | None: + """Returns the active JAX abstract mesh, if one can be determined. + + Prefers the mesh installed by an enclosing `jax.sharding.use_mesh` context. + Falls back to deriving an abstract mesh from the active Keras distribution's + device mesh. + + Returns: + The active abstract mesh, or None if there is no mesh in scope. + """ + if (global_abstract_mesh := jax.sharding.get_abstract_mesh()).shape_tuple: + return global_abstract_mesh + distribution = keras.distribution.distribution() + if distribution is None: + return None + device_mesh = getattr(distribution, 'device_mesh', None) + if device_mesh is None: + return None + return jax.sharding.AbstractMesh( + axis_sizes=tuple(device_mesh.shape), + axis_names=tuple(device_mesh.axis_names), + # `axis_types` must have one entry per axis. Passing a bare `AxisType` + # raises `ValueError` on any mesh with more than one axis. + axis_types=tuple( + jax.sharding.AxisType.Auto for _ in device_mesh.axis_names + ), + ) + + +def get_batch_dim_name() -> AxisSpec: + """Returns the `batch_dim_name` of the active Keras distribution, if any. + + Distributions that shard the batch over several mesh axes (e.g. hybrid FSDP + over `('replica', 'fsdp')`) report a tuple here rather than a single name. + + Returns: + The batch dimension name(s), or None if no distribution is active or the + distribution does not declare one. + """ + distribution = keras.distribution.distribution() + if distribution is None: + return None + return getattr(distribution, 'batch_dim_name', None) + + +def normalize_axis(axis: AxisSpec | list[str]) -> AxisSpec: + """Normalizes an axis specification to a bare name, a tuple, or None. + + A single-element sequence is collapsed to the bare axis name so that it can + be used interchangeably with a scalar axis in `jax.sharding.PartitionSpec` + and `jax.lax.all_gather`. + + Args: + axis: An axis name, a sequence of axis names, or None. + + Returns: + None for an empty or missing spec, a bare name for a single axis, or a + tuple of names otherwise. + """ + if not isinstance(axis, (tuple, list)): + return axis + if not axis: + return None + return axis[0] if len(axis) == 1 else tuple(axis) + + +def resolve_batch_axis( + abstract_mesh: jax.sharding.AbstractMesh | None, +) -> AxisSpec: + """Resolves the mesh axis or axes over which the batch is sharded. + + The active Keras distribution is authoritative: if it declares a + `batch_dim_name`, that is used verbatim. Otherwise every mesh axis other + than `MODEL_AXIS_NAME` is assumed to shard the batch, which covers both + plain data parallelism and hybrid FSDP meshes. + + Args: + abstract_mesh: The mesh to resolve against, typically from + `get_abstract_mesh`. + + Returns: + The batch axis name, a tuple of names if the batch is sharded over several + axes, or None if the batch axis cannot be determined. + """ + batch_dim_name = get_batch_dim_name() + if batch_dim_name is not None: + return normalize_axis(batch_dim_name) + if abstract_mesh is None or not abstract_mesh.axis_names: + return None + batch_axes = tuple( + axis for axis in abstract_mesh.axis_names if axis != MODEL_AXIS_NAME + ) + if not batch_axes: + # A mesh consisting only of the model axis still has to place the batch + # somewhere; fall back to the full set of axis names. + batch_axes = tuple(abstract_mesh.axis_names) + return normalize_axis(batch_axes) diff --git a/recml/core/utils/distribution_utils_test.py b/recml/core/utils/distribution_utils_test.py new file mode 100644 index 0000000..ce622a5 --- /dev/null +++ b/recml/core/utils/distribution_utils_test.py @@ -0,0 +1,137 @@ +# Copyright 2024 RecML authors . +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Tests mesh and batch-axis resolution in distribution_utils.""" + +from absl.testing import absltest +from absl.testing import parameterized +import jax +import keras +from keras.src.distribution import distribution_lib +from recml.core.utils import distribution_utils + +_DEVICES = [f'cpu:{i}' for i in range(8)] + + +class _StubDistribution: + """Minimal stand-in for a Keras distribution. + + Real hybrid-FSDP distributions live outside this package, so the two + attributes `distribution_utils` reads are stubbed here instead. + """ + + def __init__(self, device_mesh, batch_dim_name): + self.device_mesh = device_mesh + self.batch_dim_name = batch_dim_name + + +def _abstract_mesh( + axis_names: tuple[str, ...], axis_sizes: tuple[int, ...] +) -> jax.sharding.AbstractMesh: + return jax.sharding.AbstractMesh( + axis_sizes=axis_sizes, + axis_names=axis_names, + axis_types=tuple(jax.sharding.AxisType.Auto for _ in axis_names), + ) + + +class DistributionUtilsTest(parameterized.TestCase): + + def setUp(self): + super().setUp() + # The active distribution is process-global state; make sure no test + # leaks into another. + keras.distribution.set_distribution(None) + self.addCleanup(keras.distribution.set_distribution, None) + + @parameterized.named_parameters( + ('none', None, None), + ('bare_name', 'data', 'data'), + ('empty_tuple', (), None), + ('single_element_tuple', ('data',), 'data'), + ('single_element_list', ['data'], 'data'), + ('multi_element_tuple', ('replica', 'fsdp'), ('replica', 'fsdp')), + ('multi_element_list', ['replica', 'fsdp'], ('replica', 'fsdp')), + ) + def test_normalize_axis(self, axis, expected): + self.assertEqual(distribution_utils.normalize_axis(axis), expected) + + @parameterized.named_parameters( + ('single_axis', ('data',), (8,), 'data'), + ('batch_and_model', ('batch', 'model'), (2, 4), 'batch'), + ('hybrid_fsdp', ('replica', 'fsdp'), (2, 4), ('replica', 'fsdp')), + # Regression: a mesh whose only axis is the model axis must still + # resolve to something rather than an empty tuple. + ('model_axis_only', ('model',), (8,), 'model'), + ) + def test_resolve_batch_axis_from_mesh(self, axis_names, axis_sizes, expected): + mesh = _abstract_mesh(axis_names, axis_sizes) + self.assertEqual(distribution_utils.resolve_batch_axis(mesh), expected) + + def test_resolve_batch_axis_without_mesh_or_distribution(self): + self.assertIsNone(distribution_utils.resolve_batch_axis(None)) + + def test_resolve_batch_axis_prefers_distribution_over_mesh(self): + device_mesh = distribution_lib.DeviceMesh( + (2, 4), ('replica', 'fsdp'), _DEVICES + ) + keras.distribution.set_distribution( + _StubDistribution(device_mesh, ('replica', 'fsdp')) + ) + # The mesh alone would resolve to 'batch'; the distribution wins. + mesh = _abstract_mesh(('batch', 'model'), (2, 4)) + self.assertEqual( + distribution_utils.resolve_batch_axis(mesh), ('replica', 'fsdp') + ) + + def test_resolve_batch_axis_from_data_parallel_distribution(self): + device_mesh = distribution_lib.DeviceMesh((8,), ('batch',), _DEVICES) + keras.distribution.set_distribution( + keras.distribution.DataParallel(device_mesh=device_mesh) + ) + self.assertEqual(distribution_utils.resolve_batch_axis(None), 'batch') + + def test_get_batch_dim_name_without_distribution(self): + self.assertIsNone(distribution_utils.get_batch_dim_name()) + + def test_get_abstract_mesh_without_distribution_or_mesh(self): + self.assertIsNone(distribution_utils.get_abstract_mesh()) + + def test_get_abstract_mesh_from_global_jax_mesh(self): + mesh = jax.sharding.Mesh(jax.devices(), axis_names=('data',)) + with jax.set_mesh(mesh): + abstract_mesh = distribution_utils.get_abstract_mesh() + self.assertIsNotNone(abstract_mesh) + self.assertEqual(abstract_mesh.axis_names, ('data',)) + + def test_get_abstract_mesh_from_distribution(self): + device_mesh = distribution_lib.DeviceMesh( + (2, 4), ('replica', 'fsdp'), _DEVICES + ) + keras.distribution.set_distribution(_StubDistribution(device_mesh, None)) + + mesh = distribution_utils.get_abstract_mesh() + + self.assertIsNotNone(mesh) + self.assertEqual(mesh.axis_names, ('replica', 'fsdp')) + self.assertEqual(mesh.axis_sizes, (2, 4)) + # Regression: `axis_types` must have one entry per axis. Passing a bare + # `AxisType` raises `ValueError` on a mesh with more than one axis. + self.assertEqual( + mesh.axis_types, + (jax.sharding.AxisType.Auto, jax.sharding.AxisType.Auto), + ) + + +if __name__ == '__main__': + absltest.main()