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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 8 additions & 15 deletions recml/core/ops/hstu_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
)

Expand Down
122 changes: 122 additions & 0 deletions recml/core/utils/distribution_utils.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,122 @@
# Copyright 2024 RecML authors <recommendations-ml@google.com>.
#
# 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)
137 changes: 137 additions & 0 deletions recml/core/utils/distribution_utils_test.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,137 @@
# Copyright 2024 RecML authors <recommendations-ml@google.com>.
#
# 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()
Loading