Skip to content
Draft
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
45 changes: 45 additions & 0 deletions tests/jax/test_ep_partition.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,45 @@
# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
#
# See LICENSE for license information.

"""Single-process check for the EP dispatch custom_partitioning rule.

The EP FFI needs a NCCL bootstrap to compile or run (see test_multi_process_ep.py
for functional coverage), so this drives the partition rule directly to confirm
it admits a tp-sharded hidden dim and propagates it to recv_tokens.
"""

from collections import namedtuple

import jax
import numpy as np
import pytest
from jax.sharding import Mesh, NamedSharding, PartitionSpec as P

from transformer_engine.jax.cpp_extensions.ep import EpDispatchPrimitive
from transformer_engine.jax.sharding import MeshResource, global_shard_guard

# partition only reads arg_infos[i].sharding.
_ArgInfo = namedtuple("_ArgInfo", ["sharding"])


def test_ep_dispatch_partition_hidden_tp():
"""Dispatch accepts a tp-sharded hidden dim and recv_tokens inherits it."""
if len(jax.devices()) < 4:
pytest.skip("needs >= 4 devices for a (ep=2, tp=2) mesh")
mesh = Mesh(np.asarray(jax.devices()[:4]).reshape(2, 2), axis_names=("ep", "tp"))
# tokens [B, S, H] with H sharded on tp; routing tensors share the ep leading axis.
arg_infos = (
_ArgInfo(NamedSharding(mesh, P("ep"))), # handle_mem
_ArgInfo(NamedSharding(mesh, P("ep", None, None))), # topk_idx
_ArgInfo(NamedSharding(mesh, P("ep", None, "tp"))), # tokens
_ArgInfo(NamedSharding(mesh, P("ep", None, None))), # topk_weights
)
with mesh, global_shard_guard(MeshResource(ep_resource="ep", tp_resource="tp")):
_, _, out_shardings, _ = EpDispatchPrimitive.partition(
8, 0, 128, True, mesh, arg_infos, None
)

recv_tokens_sharding, recv_topk_weights_sharding = out_shardings
assert recv_tokens_sharding.spec == P("ep", None, "tp")
assert recv_topk_weights_sharding.spec == P("ep")
81 changes: 59 additions & 22 deletions transformer_engine/jax/cpp_extensions/ep.py
Original file line number Diff line number Diff line change
Expand Up @@ -111,18 +111,35 @@ def ep_handle_mem_size(cfg: EpLayerConfig) -> int:
)


def _leading_axis_ok(spec):
"""Validate an EP input spec; return ``(ok, ep_axis, outer_axes)``.
def _hidden_axis_ok(ax):
"""Allow the hidden dim of an EP token tensor to be replicated or sharded on
tp_resource. Sequence (tpsp_resource) sharding is not supported.
"""
tp_axis = global_mesh_resource().tp_resource
if ax is None:
return True
if tp_axis is None:
return False
elts = ax if isinstance(ax, tuple) else (ax,)
return all(a == tp_axis for a in elts)


Leading dim is ``ep`` or a tuple ending in ``ep`` (outer dp/fsdp axes
first); all other dims must be replicated.
def _leading_axis_ok(spec, hidden_trailing=False):
"""Validate an EP input spec; return (ok, ep_axis, outer_axes).

Leading dim is ep or a tuple ending in ep (outer dp/fsdp axes first). Middle
dims must be replicated. When hidden_trailing the last dim may also be
sharded on tp_resource; otherwise every trailing dim must be replicated.
"""
gsr = global_mesh_resource()
ep_axis = gsr.ep_resource
outer_axes = tuple(a for a in (gsr.dp_resource, gsr.fsdp_resource) if a is not None)
if len(spec) < 2 or ep_axis is None:
return False, ep_axis, outer_axes
if any(ax is not None for ax in spec[1:]):
middle = spec[1:-1] if hidden_trailing else spec[1:]
if any(ax is not None for ax in middle):
return False, ep_axis, outer_axes
if hidden_trailing and not _hidden_axis_ok(spec[-1]):
return False, ep_axis, outer_axes
leading = spec[0]
elts = leading if isinstance(leading, tuple) else (leading,)
Expand Down Expand Up @@ -168,17 +185,22 @@ def _ep_output_spec(*trailing):
return PartitionSpec((outer, gsr.ep_resource), *trailing)


def _ep_spec_ok(spec, trailing_count):
"""Leading dim shards along ep (and outer dp/fsdp when set); trailing dims
are replicated. JAX may collapse size-1 mesh axes to ``None`` or drop them,
so the leading entry is normalized to a set of named axes before comparing.
def _ep_spec_ok(spec, trailing_count, hidden_trailing=False):
"""Leading dim shards along ep (and outer dp/fsdp when set); middle dims are
replicated. When hidden_trailing the last dim may be sharded on tp_resource,
otherwise all trailing dims are replicated. JAX may collapse size-1 mesh axes
to None or drop them, so the leading entry is normalized to a set of named
axes before comparing.
"""
gsr = global_mesh_resource()
ep_axis = gsr.ep_resource
outer = _ep_outer_axis()
if len(spec) != 1 + trailing_count:
return False
if any(ax is not None for ax in spec[1:]):
middle = spec[1:-1] if hidden_trailing else spec[1:]
if any(ax is not None for ax in middle):
return False
if hidden_trailing and not _hidden_axis_ok(spec[-1]):
return False
leading = spec[0]
elts = leading if isinstance(leading, tuple) else (leading,)
Expand Down Expand Up @@ -404,12 +426,13 @@ def partition(
):
del is_outer, result_infos
tokens_spec = arg_infos[2].sharding.spec
ok, ep_axis, outer_axes = _leading_axis_ok(tokens_spec)
ok, ep_axis, outer_axes = _leading_axis_ok(tokens_spec, hidden_trailing=True)
if not ok:
raise NotImplementedError(
"EpDispatch: tokens leading dim must include ep_resource"
f" ('{ep_axis}'), optionally tupled with {outer_axes},"
f" hidden dim replicated; got spec={tokens_spec}."
f" ('{ep_axis}'), optionally tupled with {outer_axes}; the hidden"
" dim may be replicated or sharded on tp_resource, middle dims"
f" replicated; got spec={tokens_spec}."
)
idx_spec = arg_infos[1].sharding.spec
tw_spec = arg_infos[3].sharding.spec
Expand All @@ -418,11 +441,13 @@ def partition(
"EpDispatch: topk_idx, tokens, topk_weights must share the leading"
f" axis; got topk_idx={idx_spec}, tokens={tokens_spec}, topk_weights={tw_spec}."
)
# Recv outputs share the tokens leading-only spec (trailing dims auto-pad to None).
# recv_tokens inherits the tokens leading and hidden sharding (recv_pr
# replicated); recv_topk_weights carries only the leading axis.
hidden_spec = tokens_spec[-1]
leading_spec = PartitionSpec(tokens_spec[0])
arg_shardings = tuple(a.sharding for a in arg_infos)
out_shardings = (
NamedSharding(mesh, leading_spec),
NamedSharding(mesh, PartitionSpec(tokens_spec[0], None, hidden_spec)),
NamedSharding(mesh, leading_spec),
)

Expand Down Expand Up @@ -578,10 +603,11 @@ def partition(
):
del result_infos
eo_spec = arg_infos[1].sharding.spec
if not _ep_spec_ok(eo_spec, trailing_count=2):
if not _ep_spec_ok(eo_spec, trailing_count=2, hidden_trailing=True):
raise NotImplementedError(
"EpCombine: expert_out must be sharded as PartitionSpec(ep_resource,"
" None, None) (or ((dp, ep), None, None) when dp/fsdp is set)"
" None, None), with the hidden dim optionally sharded on tp_resource"
" (or ((dp, ep), ...) when dp/fsdp is set)"
f" over [num_procs, recv_pr, H]; got spec={eo_spec}."
)
if out_partition_spec is not None:
Expand Down Expand Up @@ -721,10 +747,11 @@ def partition(
):
del result_infos
g_spec = arg_infos[1].sharding.spec
if not _ep_spec_ok(g_spec, trailing_count=2):
if not _ep_spec_ok(g_spec, trailing_count=2, hidden_trailing=True):
raise NotImplementedError(
"EpDispatchBwd: grad must be sharded as PartitionSpec(ep_resource,"
" None, None) (or ((dp, ep), None, None) when dp/fsdp is set)"
" None, None), with the hidden dim optionally sharded on tp_resource"
" (or ((dp, ep), ...) when dp/fsdp is set)"
f" over [num_procs, recv_pr, H]; got spec={g_spec}."
)
gw_spec = arg_infos[2].sharding.spec
Expand All @@ -742,11 +769,14 @@ def partition(
if out_partition_spec is not None:
per_shard_leading = _leading_per_shard(out_leading_shape, out_partition_spec[0], mesh)
out_sharding = NamedSharding(mesh, PartitionSpec(*out_partition_spec))
# grad_topk_weights last dim is top_k, not hidden, so keep it replicated.
gw_sharding = NamedSharding(mesh, PartitionSpec(*out_partition_spec[:-1], None))
else:
per_shard_leading = out_leading_shape
out_sharding = NamedSharding(mesh, PartitionSpec())
gw_sharding = out_sharding
arg_shardings = tuple(a.sharding for a in arg_infos)
out_shardings = [out_sharding, out_sharding]
out_shardings = [out_sharding, gw_sharding]

def sharded_impl(handle_mem, grad, g_recv_topk_weights):
return EpDispatchBwdPrimitive.impl(
Expand Down Expand Up @@ -875,9 +905,16 @@ def partition(
result_infos,
):
del is_outer, result_infos
grad_spec = arg_infos[1].sharding.spec
hidden_spec = grad_spec[-1] if len(grad_spec) else None
if not _hidden_axis_ok(hidden_spec):
raise NotImplementedError(
"EpCombineBwd: grad hidden dim must be replicated or sharded on"
f" tp_resource; got spec={grad_spec}."
)
arg_shardings = tuple(a.sharding for a in arg_infos)
# EP-output leading (trailing dims auto-pad to None).
out_sharding = NamedSharding(mesh, _ep_output_spec())
# EP output leading; hidden dim inherits the grad sharding, recv_pr replicated.
out_sharding = NamedSharding(mesh, _ep_output_spec(None, hidden_spec))

def sharded_impl(handle_mem, grad):
return EpCombineBwdPrimitive.impl(
Expand Down
Loading