From f6ee56643e4b1d163ceda7bcbc429598d0389608 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Wed, 8 Jul 2026 16:05:30 -0700 Subject: [PATCH 01/38] Add quantized shard-map MoE support --- transformer_engine/jax/moe.py | 238 +++++++++++++++++++++++++++++----- 1 file changed, 203 insertions(+), 35 deletions(-) diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index 887e005de6b..1a4c1ec6ccf 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -25,14 +25,10 @@ a small ``shard_map`` whose ``in_specs`` and ``out_specs`` mirror the same ``((dp, ep), ...)`` layout. -Out-of-scope (for now) ----------------------- -FP8 / MXFP8 quantizer sets are not yet wired on this path; turning -them on requires recipe-aware residual specs and ``ScaledTensor`` -leaves across the ``shard_map`` boundary. ``aux_loss_coeff`` and -``expert_bias`` are supported (the former forces a per-step -all-gather over the routing-side logits, which lives off the critical -path and overlaps with the dispatch collective). +FC1 and FC2 use independent quantizer sets. The sets are differentiable +``custom_vjp`` arguments, are threaded through the per-shard FFN, and are +returned by the backward rule so stateful recipes follow the same update +semantics as :mod:`transformer_engine.jax.dense`. """ from functools import partial @@ -46,6 +42,10 @@ from . import cpp_extensions as tex from .quantize import ( + GroupedNoScaleTensor, + QuantizerSet, + ScaledTensor, + ScaledTensorFactory, TensorUsage, noop_quantizer_set, with_sharding_constraint_by_logical_axes, @@ -206,6 +206,8 @@ class _Ctx: casted_wo_rhs_trans: Any expert_outputs: jnp.ndarray local_group_sizes: jnp.ndarray + fc1_quantizer_set: QuantizerSet + fc2_quantizer_set: QuantizerSet aux_const_buf: Any = None aux_tokens_per_expert: Any = None aux_saved_scores: Any = None @@ -216,6 +218,83 @@ class _Ctx: # ============================================================================= +def _pack_grouped_tensor(tensor): + """Flatten a grouped tensor into arrays safe to cross shard_map. + + TE's grouped tensor PyTree constructors validate array shapes during + ``tree_unflatten``. JAX's shard_map prefix checker temporarily unflattens + sentinel objects, so carrying the tensor object itself across shard_map + fails before lowering. A plain tuple keeps the same residual data without + invoking those constructors. + """ + if isinstance(tensor, GroupedNoScaleTensor): + scale_inv = jnp.empty((0,), dtype=jnp.float32) + else: + scale_inv = tensor.scale_inv + first_dims = ( + tensor.first_dims + if tensor.first_dims is not None + else jnp.empty((0,), dtype=jnp.int32) + ) + last_dims = ( + tensor.last_dims + if tensor.last_dims is not None + else jnp.empty((0,), dtype=jnp.int32) + ) + return tensor.data, scale_inv, tensor.amax, first_dims, last_dims + + +def _unpack_grouped_tensor( + packed, + quantizer, + *, + original_shape, + dq_dtype, + is_colwise, + flatten_axis, +): + """Reconstruct a grouped tensor inside the backward shard_map body.""" + data, scale_inv, amax, first_dims, last_dims = packed + first_dims = first_dims if first_dims.size else None + last_dims = last_dims if last_dims.size else None + if quantizer is None: + return GroupedNoScaleTensor( + data=data, + amax=amax, + first_dims=first_dims, + last_dims=last_dims, + original_shape=original_shape, + ) + return ScaledTensorFactory.create_1x( + data=data, + scale_inv=scale_inv, + amax=amax, + scaling_mode=quantizer.scaling_mode, + dq_dtype=dq_dtype, + is_colwise=is_colwise, + data_layout="N", + flatten_axis=flatten_axis, + first_dims=first_dims, + last_dims=last_dims, + original_shape=original_shape, + pre_swizzled=quantizer.scaling_mode.is_mxfp8_scaling, + ) + + +def _token_residual_spec(batch_pspec_axis, quantized): + """Specs for a packed token-major grouped tensor.""" + data_spec = P(batch_pspec_axis) if quantized else P(batch_pspec_axis, None) + scale_spec = P(batch_pspec_axis) if quantized else P() + return data_spec, scale_spec, P(), P(batch_pspec_axis), P() + + +def _kernel_residual_spec(ep_axis, quantized): + """Specs for a packed expert-kernel grouped tensor.""" + data_spec = P(ep_axis) if quantized else P(ep_axis, None, None) + scale_spec = P(ep_axis) if quantized else P() + return data_spec, scale_spec, P(), P(), P() + + def _ffn_fwd_per_shard( recv_tokens_local: jnp.ndarray, recv_topk_weights_local: jnp.ndarray, @@ -226,6 +305,8 @@ def _ffn_fwd_per_shard( wi_0_bias: Optional[jnp.ndarray], wi_1_bias: Optional[jnp.ndarray], wo_bias: Optional[jnp.ndarray], + fc1_quantizer_set: QuantizerSet, + fc2_quantizer_set: QuantizerSet, *, num_local_experts: int, activation_type: str, @@ -262,9 +343,12 @@ def _ffn_fwd_per_shard( jnp.concatenate([wi_0_bias, wi_1_bias], axis=-1) if wi_0_bias is not None else None ) - q_set = noop_quantizer_set - casted_sorted_x = tex.grouped_quantize(sorted_x, q_set.x, local_group_sizes, flatten_axis=-1) - casted_wi = tex.grouped_quantize(wi_combined, q_set.kernel, flatten_axis=-1) + casted_sorted_x = tex.grouped_quantize( + sorted_x, fc1_quantizer_set.x, local_group_sizes, flatten_axis=-1 + ) + casted_wi = tex.grouped_quantize( + wi_combined, fc1_quantizer_set.kernel, flatten_axis=-1 + ) combined_out = tex.grouped_gemm( casted_sorted_x.get_tensor(usage=TensorUsage.LHS), casted_wi.get_tensor(usage=TensorUsage.RHS), @@ -273,7 +357,17 @@ def _ffn_fwd_per_shard( ) gate_proj_out, up_proj_out = jnp.split(combined_out, 2, axis=-1) casted_sorted_x_lhs_trans = casted_sorted_x.get_tensor(usage=TensorUsage.LHS_TRANS) + casted_sorted_x_lhs_trans = ( + casted_sorted_x_lhs_trans.checkpoint(fc1_quantizer_set.x) + if isinstance(casted_sorted_x_lhs_trans, ScaledTensor) + else casted_sorted_x_lhs_trans + ) casted_wi_rhs_trans = casted_wi.get_tensor(usage=TensorUsage.RHS_TRANS) + casted_wi_rhs_trans = ( + casted_wi_rhs_trans.checkpoint(fc1_quantizer_set.kernel) + if isinstance(casted_wi_rhs_trans, ScaledTensor) + else casted_wi_rhs_trans + ) # Activation inputs (gate_proj_out, up_proj_out) stay in the wi GEMM # output dtype; the activation output (`intermediate`) stays in the @@ -296,9 +390,9 @@ def _ffn_fwd_per_shard( intermediate = jnp.where(active, intermediate * w_b, jnp.zeros_like(intermediate)) casted_intermediate = tex.grouped_quantize( - intermediate, q_set.x, local_group_sizes, flatten_axis=-1 + intermediate, fc2_quantizer_set.x, local_group_sizes, flatten_axis=-1 ) - casted_wo = tex.grouped_quantize(wo, q_set.kernel, flatten_axis=-1) + casted_wo = tex.grouped_quantize(wo, fc2_quantizer_set.kernel, flatten_axis=-1) expert_outputs = tex.grouped_gemm( casted_intermediate.get_tensor(usage=TensorUsage.LHS), casted_wo.get_tensor(usage=TensorUsage.RHS), @@ -306,7 +400,17 @@ def _ffn_fwd_per_shard( bias=wo_bias, ) casted_intermediate_lhs_trans = casted_intermediate.get_tensor(usage=TensorUsage.LHS_TRANS) + casted_intermediate_lhs_trans = ( + casted_intermediate_lhs_trans.checkpoint(fc2_quantizer_set.x) + if isinstance(casted_intermediate_lhs_trans, ScaledTensor) + else casted_intermediate_lhs_trans + ) casted_wo_rhs_trans = casted_wo.get_tensor(usage=TensorUsage.RHS_TRANS) + casted_wo_rhs_trans = ( + casted_wo_rhs_trans.checkpoint(fc2_quantizer_set.kernel) + if isinstance(casted_wo_rhs_trans, ScaledTensor) + else casted_wo_rhs_trans + ) expert_outputs_3d = expert_outputs.reshape(1, expert_outputs.shape[0], expert_outputs.shape[1]) # Reshape local_group_sizes to (1, num_local_experts) so the @@ -314,12 +418,12 @@ def _ffn_fwd_per_shard( # global (num_procs, num_local_experts) layout matching token_counts. local_group_sizes_3d = local_group_sizes.reshape(1, num_local_experts) residuals = ( - casted_sorted_x_lhs_trans, - casted_wi_rhs_trans, + _pack_grouped_tensor(casted_sorted_x_lhs_trans), + _pack_grouped_tensor(casted_wi_rhs_trans), gate_proj_out, up_proj_out, - casted_intermediate_lhs_trans, - casted_wo_rhs_trans, + _pack_grouped_tensor(casted_intermediate_lhs_trans), + _pack_grouped_tensor(casted_wo_rhs_trans), local_group_sizes_3d, ) return expert_outputs_3d, residuals @@ -335,6 +439,8 @@ def _ffn_bwd_per_shard( casted_wo_rhs_trans, local_group_sizes: jnp.ndarray, recv_topk_weights_local: jnp.ndarray, + fc1_quantizer_set: QuantizerSet, + fc2_quantizer_set: QuantizerSet, *, activation_type: str, apply_topk_weights_early: bool, @@ -349,14 +455,51 @@ def _ffn_bwd_per_shard( local_group_sizes = local_group_sizes.reshape(-1).astype(jnp.int32) d_eo_2d = d_expert_outputs_local.reshape(-1, d_expert_outputs_local.shape[-1]) recv_w_flat = recv_topk_weights_local.reshape(-1) - q_set = noop_quantizer_set + num_local_experts = local_group_sizes.size + recv_rows = d_eo_2d.shape[0] + hidden = d_eo_2d.shape[-1] + intermediate_size = gate_proj_out.shape[-1] + casted_sorted_x_lhs_trans = _unpack_grouped_tensor( + casted_sorted_x_lhs_trans, + fc1_quantizer_set.x, + original_shape=(recv_rows, hidden), + dq_dtype=gate_proj_out.dtype, + is_colwise=True, + flatten_axis=1, + ) + casted_wi_rhs_trans = _unpack_grouped_tensor( + casted_wi_rhs_trans, + fc1_quantizer_set.kernel, + original_shape=(num_local_experts, hidden, 2 * intermediate_size), + dq_dtype=gate_proj_out.dtype, + is_colwise=False, + flatten_axis=2, + ) + casted_intermediate_lhs_trans = _unpack_grouped_tensor( + casted_intermediate_lhs_trans, + fc2_quantizer_set.x, + original_shape=(recv_rows, intermediate_size), + dq_dtype=gate_proj_out.dtype, + is_colwise=True, + flatten_axis=1, + ) + casted_wo_rhs_trans = _unpack_grouped_tensor( + casted_wo_rhs_trans, + fc2_quantizer_set.kernel, + original_shape=(num_local_experts, intermediate_size, hidden), + dq_dtype=gate_proj_out.dtype, + is_colwise=False, + flatten_axis=2, + ) # cuBLAS grouped_gemm skips size_g == 0 groups without zero-filling # the output slice; mask 0-token-expert wgrads to zero so the # optimizer never sees uninit memory. wgrad_group_active = (local_group_sizes > 0)[:, None, None] # wo bwd - casted_d_eo = tex.grouped_quantize(d_eo_2d, q_set.dgrad, local_group_sizes, flatten_axis=-1) + casted_d_eo = tex.grouped_quantize( + d_eo_2d, fc2_quantizer_set.dgrad, local_group_sizes, flatten_axis=-1 + ) _casted_d_eo_lhs = casted_d_eo.get_tensor(usage=TensorUsage.LHS) _casted_d_eo_rhs = casted_d_eo.get_tensor(usage=TensorUsage.RHS) d_intermediate = tex.grouped_gemm( @@ -408,7 +551,7 @@ def _ffn_bwd_per_shard( # wgrad result back into d_wi_0 / d_wi_1 halves with jnp.split. d_combined = jnp.concatenate([d_gate_proj_out, d_up_proj_out], axis=-1) casted_d_combined = tex.grouped_quantize( - d_combined, q_set.dgrad, local_group_sizes, flatten_axis=-1 + d_combined, fc1_quantizer_set.dgrad, local_group_sizes, flatten_axis=-1 ) d_sorted_x = tex.grouped_gemm( casted_d_combined.get_tensor(usage=TensorUsage.LHS), @@ -458,6 +601,8 @@ def _moe_fwd_rule( wi_1_bias, wo_bias, expert_bias, + fc1_quantizer_set, + fc2_quantizer_set, num_experts, num_experts_per_tok, activation_type, @@ -659,6 +804,11 @@ def _moe_fwd_rule( if has_bias: ffn_in_specs = ffn_in_specs + (bias_spec, bias_spec, bias_spec) ffn_in_args.extend([wi_0_bias, wi_1_bias, wo_bias]) + # QuantizerSet is a JAX pytree. P() is a tree-prefix specification + # that replicates any recipe state into each FFN shard; stateless + # recipes such as MXFP8 have no array leaves here. + ffn_in_specs = ffn_in_specs + (P(), P()) + ffn_in_args.extend([fc1_quantizer_set, fc2_quantizer_set]) # FFN residuals live entirely on the local ep rank, so the leading # "experts" / "rows" dims map to P() (already shard-local). wi is @@ -669,21 +819,21 @@ def _moe_fwd_rule( # now per-shard dynamic (= per-shard token_counts), so its # residual spec mirrors ep2_spec (one row per ep rank). residuals_spec = ( - P(), # casted_sorted_x_lhs_trans - P(ep_axis, None, None), # casted_wi_rhs_trans - P(), # gate_proj_out - P(), # up_proj_out - P(), # casted_intermediate_lhs_trans - P(ep_axis, None, None), # casted_wo_rhs_trans + _token_residual_spec(batch_pspec_axis, fc1_quantizer_set.x is not None), + _kernel_residual_spec(ep_axis, fc1_quantizer_set.kernel is not None), + P(batch_pspec_axis, None), # gate_proj_out + P(batch_pspec_axis, None), # up_proj_out + _token_residual_spec(batch_pspec_axis, fc2_quantizer_set.x is not None), + _kernel_residual_spec(ep_axis, fc2_quantizer_set.kernel is not None), ep2_spec, # local_group_sizes (1, num_local_experts) per shard ) out_specs = (ep3_spec, residuals_spec) def _body(*args): if has_bias: - (r_tok, r_w, tc, w0, w1, w_o, w0b, w1b, wob) = args + (r_tok, r_w, tc, w0, w1, w_o, w0b, w1b, wob, fc1_qset, fc2_qset) = args else: - (r_tok, r_w, tc, w0, w1, w_o) = args + (r_tok, r_w, tc, w0, w1, w_o, fc1_qset, fc2_qset) = args w0b = w1b = wob = None # NOTE: tex.ep_dispatch_fwd's NCCL EP HT path leaves the recv # buffer uninitialised on fully-empty-receiver ranks (and at @@ -714,6 +864,8 @@ def _body(*args): w0b, w1b, wob, + fc1_qset, + fc2_qset, num_local_experts=num_local_experts, activation_type=activation_type, apply_topk_weights_early=apply_topk_weights_early, @@ -781,6 +933,8 @@ def _body(*args): casted_wo_rhs_trans=casted_wo_rhs_trans, expert_outputs=expert_outputs, local_group_sizes=local_group_sizes, + fc1_quantizer_set=fc1_quantizer_set, + fc2_quantizer_set=fc2_quantizer_set, aux_const_buf=aux_const_buf, aux_tokens_per_expert=aux_tokens_per_expert, aux_saved_scores=aux_saved_scores, @@ -872,14 +1026,16 @@ def _moe_bwd_rule( bwd_in_specs = ( ep3_spec, # d_expert_outputs - P(), # casted_sorted_x_lhs_trans - P(ep_axis, None, None), # casted_wi_rhs_trans - P(), # gate_proj_out - P(), # up_proj_out - P(), # casted_intermediate_lhs_trans - P(ep_axis, None, None), # casted_wo_rhs_trans + _token_residual_spec(batch_pspec_axis, ctx.fc1_quantizer_set.x is not None), + _kernel_residual_spec(ep_axis, ctx.fc1_quantizer_set.kernel is not None), + P(batch_pspec_axis, None), # gate_proj_out + P(batch_pspec_axis, None), # up_proj_out + _token_residual_spec(batch_pspec_axis, ctx.fc2_quantizer_set.x is not None), + _kernel_residual_spec(ep_axis, ctx.fc2_quantizer_set.kernel is not None), ep2_spec, # local_group_sizes (1, num_local_experts) per shard ep2_spec, # recv_topk_weights + P(), # FC1 quantizer-set state, replicated into each shard + P(), # FC2 quantizer-set state, replicated into each shard ) bwd_in_args = [ d_expert_outputs, @@ -891,6 +1047,8 @@ def _moe_bwd_rule( ctx.casted_wo_rhs_trans, ctx.local_group_sizes, ctx.recv_topk_weights, + ctx.fc1_quantizer_set, + ctx.fc2_quantizer_set, ] bwd_out_specs = ( ep3_spec, # d_sorted_x @@ -1056,6 +1214,8 @@ def _bwd_body(*args): d_wi_1_bias if has_bias else None, d_wo_bias if has_bias else None, d_expert_bias, + ctx.fc1_quantizer_set, + ctx.fc2_quantizer_set, ) @@ -1064,7 +1224,7 @@ def _bwd_body(*args): # ============================================================================= -@partial(jax.custom_vjp, nondiff_argnums=tuple(range(9, 26))) +@partial(jax.custom_vjp, nondiff_argnums=tuple(range(11, 28))) def _moe( x, gate_kernel, @@ -1075,6 +1235,8 @@ def _moe( wi_1_bias, wo_bias, expert_bias, + fc1_quantizer_set, + fc2_quantizer_set, num_experts, num_experts_per_tok, activation_type, @@ -1103,6 +1265,8 @@ def _moe( wi_1_bias, wo_bias, expert_bias, + fc1_quantizer_set, + fc2_quantizer_set, num_experts, num_experts_per_tok, activation_type, @@ -1137,6 +1301,8 @@ def moe( wi_1_bias: Optional[jnp.ndarray] = None, wo_bias: Optional[jnp.ndarray] = None, expert_bias: Optional[jnp.ndarray] = None, + fc1_quantizer_set: QuantizerSet = noop_quantizer_set, + fc2_quantizer_set: QuantizerSet = noop_quantizer_set, *, num_experts: int, num_experts_per_tok: int, @@ -1248,6 +1414,8 @@ def moe( wi_1_bias, wo_bias, expert_bias_arg, + fc1_quantizer_set, + fc2_quantizer_set, num_experts, num_experts_per_tok, activation_type, From 1ae5c3d21ecd8b3864ead45c901cf2a08c1119d7 Mon Sep 17 00:00:00 2001 From: tdophung Date: Thu, 23 Jul 2026 14:38:00 -0700 Subject: [PATCH 02/38] JAX: integrate cuDNN CuTeDSL fused MoE Add the opt-in grouped GEMM SwiGLU and dSwiGLU forward/backward path, MXFP8 scale handling, eligibility fallbacks, EP padding and routing-gradient fixes, and isolated multiprocess validation flow. --- docs/envvars.rst | 10 + docs/examples/jax/cutedsl_moe_pipeclean.rst | 85 +++ docs/examples/te_jax_integration.rst | 1 + tests/jax/cutedsl_smoke.py | 86 +++ tests/jax/repro_cutedsl_moe_shardy.py | 124 ++++ tests/jax/requirements_cutedsl.txt | 4 + tests/jax/run_te_ep_moe.sh | 130 ++-- tests/jax/test_cutedsl_moe.py | 248 ++++++++ tests/jax/test_te_ep_moe.py | 228 ++++++- .../jax/cutedsl_extensions/__init__.py | 4 + .../jax/cutedsl_extensions/moe.py | 576 ++++++++++++++++++ transformer_engine/jax/flax/moe.py | 18 +- transformer_engine/jax/moe.py | 517 ++++++++++++---- 13 files changed, 1856 insertions(+), 175 deletions(-) create mode 100644 docs/examples/jax/cutedsl_moe_pipeclean.rst create mode 100644 tests/jax/cutedsl_smoke.py create mode 100644 tests/jax/repro_cutedsl_moe_shardy.py create mode 100644 tests/jax/requirements_cutedsl.txt create mode 100644 tests/jax/test_cutedsl_moe.py create mode 100644 transformer_engine/jax/cutedsl_extensions/__init__.py create mode 100644 transformer_engine/jax/cutedsl_extensions/moe.py diff --git a/docs/envvars.rst b/docs/envvars.rst index b3765a06bde..a14c9290b51 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -498,6 +498,16 @@ JAX-Specific Variables :Default: None :Description: Test level for JAX unit tests (``"L0"``, ``"L1"``, ``"L2"``). Used internally by the test suite. +.. envvar:: NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION + + :Type: ``int`` (0 or 1) + :Default: ``0`` + :Description: **(JAX only)** Enable the experimental cuDNN frontend CuTeDSL fusion + for MXFP8 MoE FC1 grouped GEMM, SwiGLU, and grouped quantization. Explicit + opt-in requires an eligible SM100 SwiGLU MXFP8 MoE call and the optional + CUTLASS/cuDNN frontend JAX runtime packages; unsupported calls fail with the + full validation reason list. + JAX Triton Extensions ^^^^^^^^^^^^^^^^^^^^^ diff --git a/docs/examples/jax/cutedsl_moe_pipeclean.rst b/docs/examples/jax/cutedsl_moe_pipeclean.rst new file mode 100644 index 00000000000..f725eb45bc1 --- /dev/null +++ b/docs/examples/jax/cutedsl_moe_pipeclean.rst @@ -0,0 +1,85 @@ +CuTeDSL MXFP8 MoE pipeclean +=========================== + +The JAX MoE path has an opt-in Blackwell fusion for FC1 grouped GEMM, +SwiGLU, and rowwise/colwise MXFP8 quantization. All CuTe-specific code lives +in ``transformer_engine/jax/cutedsl_extensions``; the surrounding EP path +continues to use ``shard_map`` and sees only shard-local tensors. + +Environment baseline +-------------------- + +Use the versions in ``tests/jax/requirements_cutedsl.txt`` on an SM100 CUDA +host. Build Transformer Engine with JAX and NCCL EP support, then run:: + + python3 tests/jax/cutedsl_smoke.py + python3 -m pytest -c tests/jax/pytest.ini tests/jax/test_cutedsl_moe.py -v + bash tests/jax/run_te_ep_moe.sh + NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION=1 \ + bash tests/jax/run_te_ep_moe.sh -k TestTeEpMoeCudnnCutedslFusion + +The first command is intentionally independent of Transformer Engine. It +must report ``cutlass.jax.is_available()`` and execute its vector-add kernel +inside ``jax.jit`` before failures in the fused MoE tests are investigated. +The EP launcher currently requires four or more ranks even though the EP +mesh itself uses groups of two ranks. The CuTeDSL opt-in is strict, so run +only the CuTeDSL fusion class with that environment variable enabled. + +Support and fallback contract +----------------------------- + +The default value of ``NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION`` is ``0`` and +never imports CUTLASS DSL. Existing unfused runs therefore remain supported +when the optional packages or SM100 hardware are absent. Explicit opt-in +validates the GPU architecture, SwiGLU shapes and biases, MXFP8 quantizers, +JAX/CUTLASS FFI availability, and the cuDNN frontend source layout before +kernel lowering. An unsupported explicit opt-in fails with the complete validation +list instead of failing later during kernel compilation. + +The pinned cuDNN frontend layout is:: + + cudnn/grouped_gemm/utils.py + cudnn/grouped_gemm/grouped_gemm_swiglu/grouped_gemm_swiglu_quant.py + +and must export ``BlockScaledContiguousGroupedGemmKernel``. The fused path +uses 256-token dispatch alignment; the unfused path remains at 128. + +Custom partitioning migration design +------------------------------------ + +Do not add a ``custom_partitioning`` wrapper until nvbug 6432162 is fixed. +Today the five kernel outputs are: + +* combined projection: ``[tokens, 2 * intermediate, 1]`` +* rowwise MXFP8 payload: ``[tokens, intermediate, 1]`` +* colwise MXFP8 payload: ``[tokens, intermediate, 1]`` +* rowwise inverse scales: one flat buffer +* colwise inverse scales: one flat buffer + +The draft Shardy factors are ``tokens`` (sharded by the caller's DP/EP token +axes), ``experts`` (sharded by EP), ``hidden``, ``intermediate``, distinct +singleton factors for each singleton dimension, and a constant factor +``swiglu_pair=2``. Inputs map as follows: + +* A ``[M,K,1]``: ``(tokens, hidden, a_l)`` +* B ``[E,2I,K]``: ``(experts, intermediate*swiglu_pair, hidden)`` +* padded offsets ``[E]``: ``(experts,)`` +* probability ``[M,1,1]``: ``(tokens, prob_n, prob_l)`` +* pre-swizzled A/B scales: opaque until their JAX boundary is structured + +The first three outputs inherit the ``tokens`` factor and otherwise remain +local. The two scale outputs must eventually be exposed as structured block +layouts carrying ``tokens``/``intermediate`` (and the per-expert padding +factor where applicable), then flattened only inside ``cutlass_call``. Do +not assign a Shardy factor to the current giant flat dimension: that is the +failure mode tracked by nvbug 6432162. Until structured scales and the bug +fix are both available, the leaf stays under ``shard_map`` and needs no +partitioning rule. + +``tests/jax/repro_cutedsl_moe_shardy.py`` lowers a shape-only custom +partitioning leaf with the production EP2/FSDP2 output shapes. Run it before +starting the migration; it should lower successfully after the Shardy fix. +After that gate passes, replace the shape-only leaf with +``grouped_gemm_swiglu_mxfp8``, structure both scale outputs, add the rule +above, and compare its shardings and numerics with the existing ``shard_map`` +tests before removing ``shard_map``. diff --git a/docs/examples/te_jax_integration.rst b/docs/examples/te_jax_integration.rst index a15a10e0b3a..c25f2a276b3 100644 --- a/docs/examples/te_jax_integration.rst +++ b/docs/examples/te_jax_integration.rst @@ -93,3 +93,4 @@ Conventions used across these documents jax/collective_gemm jax/attention jax/expert_parallelism + jax/cutedsl_moe_pipeclean diff --git a/tests/jax/cutedsl_smoke.py b/tests/jax/cutedsl_smoke.py new file mode 100644 index 00000000000..e36806ede98 --- /dev/null +++ b/tests/jax/cutedsl_smoke.py @@ -0,0 +1,86 @@ +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. +"""Standalone CuTeDSL/JAX toolchain smoke test. + +Run this before the TE tests. It verifies CUTLASS FFI runtime discovery and +executes a trivial CuTeDSL kernel from inside ``jax.jit``. +""" + +from importlib.metadata import version + +import cuda.bindings.driver as cuda +import cutlass.cute as cute +import cutlass.jax as cjax +import jax +import jax.numpy as jnp +import numpy as np + + +BLOCK = 256 + + +@cute.kernel +def _vector_add_kernel(a: cute.Tensor, b: cute.Tensor, c: cute.Tensor): + tidx, _, _ = cute.arch.thread_idx() + bidx, _, _ = cute.arch.block_idx() + + frag_a = cute.make_rmem_tensor(cute.size(a, mode=[0]), a.element_type) + frag_b = cute.make_rmem_tensor(cute.size(b, mode=[0]), b.element_type) + frag_c = cute.make_rmem_tensor(cute.size(c, mode=[0]), c.element_type) + cute.autovec_copy(a[None, tidx, bidx], frag_a) + cute.autovec_copy(b[None, tidx, bidx], frag_b) + frag_c.store(frag_a.load() + frag_b.load()) + cute.autovec_copy(frag_c, c[None, tidx, bidx]) + + +@cute.jit +def _launch_vector_add( + stream: cuda.CUstream, + a: cute.Tensor, + b: cute.Tensor, + c: cute.Tensor, +): + _vector_add_kernel(a, b, c).launch( + grid=[a.shape[-1], 1, 1], + block=[a.shape[-2], 1, 1], + stream=stream, + ) + + +@jax.jit +def _cutlass_add(a, b): + size = a.shape[0] + padded_size = ((size + BLOCK - 1) // BLOCK) * BLOCK + a_3d = jnp.pad(a, (0, padded_size - size)).reshape(1, BLOCK, -1) + b_3d = jnp.pad(b, (0, padded_size - size)).reshape(1, BLOCK, -1) + call = cjax.cutlass_call( + _launch_vector_add, + output_shape_dtype=jax.ShapeDtypeStruct.like(a_3d), + use_static_tensors=True, + ) + return call(a_3d, b_3d).reshape(-1)[:size] + + +def main() -> None: + if not cjax.is_available(): + raise RuntimeError( + "cutlass.jax.is_available() is false; verify cute_dsl_runtime.so discovery" + ) + devices = jax.devices("gpu") + if not devices: + raise RuntimeError("No JAX GPU device is available") + + a = jnp.arange(1024, dtype=jnp.float32) + b = jnp.arange(1024, dtype=jnp.float32) * 2 + actual = _cutlass_add(a, b) + actual.block_until_ready() + np.testing.assert_array_equal(np.asarray(actual), np.asarray(a + b)) + print( + "CuTeDSL JAX smoke passed: " + f"jax={jax.__version__}, cutlass={version('nvidia-cutlass-dsl')}, device={devices[0]}" + ) + + +if __name__ == "__main__": + main() diff --git a/tests/jax/repro_cutedsl_moe_shardy.py b/tests/jax/repro_cutedsl_moe_shardy.py new file mode 100644 index 00000000000..0821116adc4 --- /dev/null +++ b/tests/jax/repro_cutedsl_moe_shardy.py @@ -0,0 +1,124 @@ +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. +"""Shape-only reproducer for nvbug 6432162 and CuTeDSL MoE outputs. + +This intentionally does not call the kernel. It isolates Shardy propagation +over the exact five output shapes used by the production EP2/FSDP2 slice. +""" + +import math + +import jax +import jax.numpy as jnp +import numpy as np +from jax.experimental.custom_partitioning import ( + CompoundFactor, + SdyShardingRule, + custom_partitioning, +) +from jax.sharding import Mesh, NamedSharding, PartitionSpec as P + + +EXPERTS = 16 +M = 266240 +HIDDEN = 1792 +INTERMEDIATE = 2048 +PAIR = 2 + + +def _ceil_div(x, y): + return (x + y - 1) // y + + +ROW_SCALE_SIZE = 32 * 4 * _ceil_div(M, 128) * 4 * _ceil_div(_ceil_div(INTERMEDIATE, 32), 4) +COL_SCALE_SIZE = 32 * 4 * _ceil_div(INTERMEDIATE, 128) * 4 * _ceil_div(_ceil_div(M, 32), 4) + + +def _shape_only_impl(a, b, sfa, sfb, padded_offsets, prob): + del a, b, sfa, sfb, padded_offsets, prob + return ( + jnp.zeros((M, PAIR * INTERMEDIATE, 1), dtype=jnp.bfloat16), + jnp.zeros((M, INTERMEDIATE, 1), dtype=jnp.float8_e4m3fn), + jnp.zeros((M, INTERMEDIATE, 1), dtype=jnp.float8_e4m3fn), + jnp.zeros((ROW_SCALE_SIZE,), dtype=jnp.float8_e8m0fnu), + jnp.zeros((COL_SCALE_SIZE,), dtype=jnp.float8_e8m0fnu), + ) + + +_shape_only = custom_partitioning(_shape_only_impl) + + +def _partition(mesh, arg_infos, result_infos): + arg_shardings = tuple(info.sharding for info in arg_infos) + result_shardings = tuple(info.sharding for info in result_infos) + return mesh, _shape_only_impl, result_shardings, arg_shardings + + +def _shardy_rule(mesh, value_types, result_types): + del mesh, value_types, result_types + combined = CompoundFactor("intermediate", "swiglu_pair") + return SdyShardingRule( + ( + ("tokens", "hidden", "a_l"), + ("experts", combined, "hidden"), + ("sfa_flat",), + ("sfb_flat",), + ("experts",), + ("tokens", "prob_n", "prob_l"), + ), + ( + ("tokens", combined, "combined_l"), + ("tokens", "intermediate", "row_l"), + ("tokens", "intermediate", "col_l"), + ("row_scale_flat",), + ("col_scale_flat",), + ), + swiglu_pair=PAIR, + ) + + +_shape_only.def_partition(partition=_partition, sharding_rule=_shardy_rule) + + +def main() -> None: + jax.config.update("jax_use_shardy_partitioner", True) + devices = np.asarray(jax.devices()) + if devices.size < 2: + raise RuntimeError("This Shardy reproducer requires at least two JAX devices") + ep = 2 + dp = devices.size // ep + devices = devices[: dp * ep].reshape(dp, ep) + mesh = Mesh(devices, ("data", "expert")) + token_sharding = NamedSharding(mesh, P(("data", "expert"), None, None)) + expert_sharding = NamedSharding(mesh, P("expert", None, None)) + replicated = NamedSharding(mesh, P()) + + args = ( + jax.ShapeDtypeStruct((M, HIDDEN, 1), jnp.float8_e4m3fn, sharding=token_sharding), + jax.ShapeDtypeStruct( + (EXPERTS, PAIR * INTERMEDIATE, HIDDEN), + jnp.float8_e4m3fn, + sharding=expert_sharding, + ), + jax.ShapeDtypeStruct((1,), jnp.float8_e8m0fnu, sharding=replicated), + jax.ShapeDtypeStruct((1,), jnp.float8_e8m0fnu, sharding=replicated), + jax.ShapeDtypeStruct( + (EXPERTS,), + jnp.int32, + sharding=NamedSharding(mesh, P("expert")), + ), + jax.ShapeDtypeStruct((M, 1, 1), jnp.float32, sharding=token_sharding), + ) + with jax.set_mesh(mesh): + lowered = jax.jit(_shape_only).lower(*args) + print( + "Shardy lowering passed for CuTeDSL MoE outputs: " + f"M={M}, I={INTERMEDIATE}, row_scale={ROW_SCALE_SIZE}, " + f"col_scale={COL_SCALE_SIZE}, devices={math.prod(mesh.devices.shape)}" + ) + print(lowered.compiler_ir(dialect="stablehlo")) + + +if __name__ == "__main__": + main() diff --git a/tests/jax/requirements_cutedsl.txt b/tests/jax/requirements_cutedsl.txt new file mode 100644 index 00000000000..9cf81761a31 --- /dev/null +++ b/tests/jax/requirements_cutedsl.txt @@ -0,0 +1,4 @@ +# Known-good baseline for the JAX CuTeDSL MoE pipeclean. +jax[cuda13]>=0.9.1 +nvidia-cudnn-frontend==1.25.0 +nvidia-cutlass-dsl[cu13]==4.5.2 diff --git a/tests/jax/run_te_ep_moe.sh b/tests/jax/run_te_ep_moe.sh index 32d5f21956a..1f9c4ee8f5b 100755 --- a/tests/jax/run_te_ep_moe.sh +++ b/tests/jax/run_te_ep_moe.sh @@ -33,6 +33,9 @@ echo " test file : $TEST_FILE" echo " coordinator : $TE_EP_MOE_COORDINATOR_ADDRESS" echo " XLA_PYTHON_CLIENT_PREALLOCATE: $XLA_PYTHON_CLIENT_PREALLOCATE" echo " XLA_PYTHON_CLIENT_MEM_FRACTION: $XLA_PYTHON_CLIENT_MEM_FRACTION" +if [ "$#" -gt 0 ]; then + echo " extra pytest args : $*" +fi echo "============================================================" if [ -n "${TE_EP_MOE_MP_LOG_DIR:-}" ]; then @@ -44,6 +47,8 @@ fi echo "Per-process logs: $LOG_DIR" PIDS=() +EXITS=() +PHASE_FAILED=0 cleanup() { for pid in "${PIDS[@]:-}"; do @@ -60,63 +65,96 @@ cleanup() { } trap cleanup EXIT INT TERM -for i in $(seq 0 $((NUM_GPUS - 1))); do - LOG_FILE="$LOG_DIR/proc_${i}.log" - PYTEST_CMD=( - python3 -m pytest -c "$PYTEST_INI" - "$TEST_FILE" - -p no:typeguard - -v -s - --num-process="$NUM_GPUS" - --process-id="$i" - ) - if [ "$i" -eq 0 ]; then - echo "=== Live output from process 0 ===" - "${PYTEST_CMD[@]}" 2>&1 | tee "$LOG_FILE" & - else - "${PYTEST_CMD[@]}" > "$LOG_FILE" 2>&1 & - fi - PIDS+=("$!") -done +run_phase() { + local phase_name="$1" + local fusion_env="$2" + shift 2 + local -a phase_args=("$@") + local phase_log_dir="$LOG_DIR/$phase_name" + mkdir -p "$phase_log_dir" + PIDS=() + EXITS=() -EXITS=() -for pid in "${PIDS[@]}"; do - if wait "$pid"; then - EXITS+=("0") + echo + echo "============================================================" + echo "Phase: $phase_name" + echo " NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION=$fusion_env" + echo " phase pytest args : ${phase_args[*]:-}" + echo " logs : $phase_log_dir" + echo "============================================================" + + for i in $(seq 0 $((NUM_GPUS - 1))); do + local log_file="$phase_log_dir/proc_${i}.log" + local -a pytest_cmd=( + python3 -m pytest -c "$PYTEST_INI" + "$TEST_FILE" + -p no:typeguard + -v -s + --num-process="$NUM_GPUS" + --process-id="$i" + "${phase_args[@]}" + ) + if [ "$i" -eq 0 ]; then + echo "=== Live output from process 0 ($phase_name) ===" + env NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION="$fusion_env" \ + "${pytest_cmd[@]}" 2>&1 | tee "$log_file" & + else + env NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION="$fusion_env" \ + "${pytest_cmd[@]}" > "$log_file" 2>&1 & + fi + PIDS+=("$!") + done + + for pid in "${PIDS[@]}"; do + if wait "$pid"; then + EXITS+=("0") + else + EXITS+=("$?") + fi + done + + echo + echo "Per-process exit codes for $phase_name:" + for i in "${!EXITS[@]}"; do + echo " proc $i -> ${EXITS[$i]}" + done + + local failed=0 + for e in "${EXITS[@]}"; do + if [ "$e" != "0" ] && [ "$e" != "5" ]; then + failed=1 + break + fi + done + if [ "$failed" -ne 0 ]; then + PHASE_FAILED=1 + echo "[run_te_ep_moe.sh] phase $phase_name FAILED" + echo " process 0 tail:" + tail -20 "$phase_log_dir/proc_0.log" 2>/dev/null || true else - EXITS+=("$?") + echo "[run_te_ep_moe.sh] phase $phase_name PASSED" fi -done +} -echo -echo "============================================================" -echo "Per-process exit codes:" -for i in "${!EXITS[@]}"; do - echo " proc $i -> ${EXITS[$i]}" -done - -# Treat exit 0 (pass) and exit 5 (pytest "no tests collected", which the -# file emits via pytest.skip(allow_module_level=True) on pre-Blackwell -# GPUs) as success. -FAILED=0 -for e in "${EXITS[@]}"; do - if [ "$e" != "0" ] && [ "$e" != "5" ]; then - FAILED=1 - break - fi -done +if [ "${NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION:-0}" = "1" ]; then + # Keep ordinary CUDA C++ and CuTeDSL coverage in separate Python process + # groups. TE EP/NCCL caches layer alignment process-wide, so 128-token and + # 256-token dispatch-alignment tests cannot safely share one interpreter. + run_phase "ordinary" "0" -k "not TestTeEpMoeCudnnCutedslFusion" "$@" + run_phase "cutedsl" "1" -k "TestTeEpMoeCudnnCutedslFusion" "$@" +else + run_phase "ordinary" "0" "$@" +fi echo -if [ "$FAILED" -eq 0 ]; then - echo "[run_te_ep_moe.sh] all processes PASSED" +if [ "$PHASE_FAILED" -eq 0 ]; then + echo "[run_te_ep_moe.sh] all phases PASSED" if [ -z "${TE_EP_MOE_MP_LOG_DIR:-}" ]; then rm -rf "$LOG_DIR" fi exit 0 fi -echo "[run_te_ep_moe.sh] at least one process FAILED" +echo "[run_te_ep_moe.sh] at least one phase FAILED" echo " retaining logs at $LOG_DIR for diagnosis" -echo " process 0 tail:" -tail -20 "$LOG_DIR/proc_0.log" 2>/dev/null || true exit 1 diff --git a/tests/jax/test_cutedsl_moe.py b/tests/jax/test_cutedsl_moe.py new file mode 100644 index 00000000000..b46aac78c7c --- /dev/null +++ b/tests/jax/test_cutedsl_moe.py @@ -0,0 +1,248 @@ +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. +"""Tests for the JAX binding to cuDNN frontend's CuTeDSL MoE kernel.""" + +import sys + +import jax +import jax.numpy as jnp +import numpy as np +import pytest + +from transformer_engine.jax import cpp_extensions as tex +from transformer_engine.jax.cutedsl_extensions.moe import ( + grouped_gemm_dswiglu_mxfp8, + grouped_gemm_swiglu_mxfp8, + load_grouped_gemm_swiglu_kernel, + pack_swiglu_pair, + unpack_swiglu_pair, +) +from transformer_engine.jax.quantize import ( + QuantizerFactory, + ScaledTensorFactory, + ScalingMode, + TensorUsage, +) + + +def test_swiglu_block_pack_round_trip(): + """Gate/up packing alternates 32-column blocks and is reversible.""" + gate = jnp.arange(2 * 3 * 64, dtype=jnp.float32).reshape(2, 3, 64) + up = gate + 1000 + interleaved = pack_swiglu_pair(gate, up) + + np.testing.assert_array_equal(interleaved[..., :32], gate[..., :32]) + np.testing.assert_array_equal(interleaved[..., 32:64], up[..., :32]) + unpacked_gate, unpacked_up = unpack_swiglu_pair(interleaved) + np.testing.assert_array_equal(unpacked_gate, gate) + np.testing.assert_array_equal(unpacked_up, up) + + +def test_direct_kernel_loader_bypasses_torch_api(): + """The direct source loader must not execute cuDNN's Torch API package.""" + kernel_cls = load_grouped_gemm_swiglu_kernel() + assert kernel_cls.__name__ == "BlockScaledContiguousGroupedGemmKernel" + assert "cudnn.grouped_gemm.grouped_gemm_swiglu.api" not in sys.modules + # cuDNN frontend 1.25's shared utility source imports torch only for an + # unused annotation. The loader supplies a non-executable sentinel rather + # than importing the Torch package. + torch_module = sys.modules.get("torch") + assert torch_module is None or getattr( + torch_module, "__transformer_engine_cutedsl_stub__", False + ) + + +def test_swiglu_forward_fused_output_parity(): + """The forward fused call matches TE projection plus JAX SwiGLU reference.""" + try: + from transformer_engine_jax import get_device_compute_capability + + if get_device_compute_capability(0) != 100: + pytest.skip("cuDNN frontend grouped GEMM SwiGLU requires SM100") + load_grouped_gemm_swiglu_kernel() + import cutlass.jax # noqa: F401 # pylint: disable=unused-import,import-outside-toplevel + except (ImportError, RuntimeError) as exc: + pytest.skip(f"CuTeDSL JAX dependencies are unavailable: {exc}") + + experts, rows, hidden, intermediate = 1, 256, 128, 128 + key = jax.random.PRNGKey(123) + x = jax.random.normal(key, (rows, hidden), dtype=jnp.bfloat16) + wi_0 = jax.random.normal( + jax.random.fold_in(key, 1), (experts, hidden, intermediate), dtype=jnp.bfloat16 + ) + wi_1 = jax.random.normal( + jax.random.fold_in(key, 2), (experts, hidden, intermediate), dtype=jnp.bfloat16 + ) + wi = pack_swiglu_pair(wi_0, wi_1) + group_sizes = jnp.asarray([rows], dtype=jnp.int32) + quantizers = QuantizerFactory.create_set( + scaling_mode=ScalingMode.MXFP8_1D_SCALING, + fwd_dtype=jnp.float8_e4m3fn, + bwd_dtype=jnp.float8_e5m2, + is_2x2x=True, + n_groups=experts, + ) + + @jax.jit + def run(x_arg, wi_arg): + casted_x = tex.grouped_quantize( + x_arg, quantizers.x, group_sizes, flatten_axis=-1 + ).get_tensor(TensorUsage.LHS) + casted_wi = tex.grouped_quantize(wi_arg, quantizers.kernel, flatten_axis=-1).get_tensor( + TensorUsage.RHS + ) + reference = tex.grouped_gemm( + casted_x, + casted_wi, + contracting_dims=((1,), (1,)), + ) + physical_wi = casted_wi.data.reshape(experts, hidden, 2 * intermediate).transpose(0, 2, 1) + combined, swiglu_row, swiglu_col, scale_row, scale_col = grouped_gemm_swiglu_mxfp8( + casted_x.data.reshape(rows, hidden, 1), + physical_wi, + casted_x.scale_inv, + casted_wi.scale_inv, + jnp.cumsum(group_sizes), + jnp.ones((rows, 1, 1), dtype=jnp.float32), + compute_dtype=jnp.bfloat16, + output_dtype=jnp.float8_e4m3fn, + ) + return ( + reference, + combined.reshape(rows, 2 * intermediate), + swiglu_row, + swiglu_col, + scale_row, + scale_col, + ) + + reference, combined, swiglu_row, swiglu_col, scale_row, scale_col = run(x, wi) + jax.block_until_ready((reference, combined, swiglu_row, swiglu_col)) + np.testing.assert_array_equal(combined, reference) + assert swiglu_row.shape == swiglu_col.shape == (rows, intermediate, 1) + assert swiglu_row.dtype == swiglu_col.dtype == jnp.float8_e4m3fn + assert scale_row.dtype == scale_col.dtype == jnp.float8_e8m0fnu + + gate, up = unpack_swiglu_pair(combined) + swiglu_reference = np.asarray(jax.nn.silu(gate) * up, dtype=np.float32) + for is_colwise, payload, scale in ( + (False, swiglu_row, scale_row), + (True, swiglu_col, scale_col), + ): + scaled = ScaledTensorFactory.create_1x( + payload.reshape(-1), + scale, + scaling_mode=ScalingMode.MXFP8_1D_SCALING, + dq_dtype=jnp.bfloat16, + is_colwise=is_colwise, + data_layout="N", + flatten_axis=1, + first_dims=group_sizes, + original_shape=(rows, intermediate), + pre_swizzled=True, + ) + dequantized = np.asarray(scaled.dequantize()[0], dtype=np.float32) + assert np.all(np.isfinite(dequantized)) + relative_error = np.linalg.norm(dequantized - swiglu_reference) / np.linalg.norm( + swiglu_reference + ) + assert relative_error < 0.05 + + +def test_dswiglu_backward_quantized_output_parity(): + """The backward fused call emits the expected packed MXFP8 VJP.""" + try: + from transformer_engine_jax import get_device_compute_capability + + if get_device_compute_capability(0) != 100: + pytest.skip("cuDNN frontend grouped GEMM dSwiGLU requires SM100") + load_grouped_gemm_swiglu_kernel() + from transformer_engine.jax.cutedsl_extensions.moe import load_grouped_gemm_dswiglu_kernel + + load_grouped_gemm_dswiglu_kernel() + import cutlass.jax # noqa: F401 # pylint: disable=unused-import,import-outside-toplevel + except (ImportError, RuntimeError) as exc: + pytest.skip(f"CuTeDSL JAX dependencies are unavailable: {exc}") + + experts, rows, hidden, intermediate = 1, 256, 128, 128 + key = jax.random.PRNGKey(321) + d_eo = jax.random.normal(key, (rows, hidden), dtype=jnp.bfloat16) + wo = jax.random.normal( + jax.random.fold_in(key, 1), (experts, intermediate, hidden), dtype=jnp.bfloat16 + ) + gate = jax.random.normal(jax.random.fold_in(key, 2), (rows, intermediate), dtype=jnp.bfloat16) + up = jax.random.normal(jax.random.fold_in(key, 3), (rows, intermediate), dtype=jnp.bfloat16) + packed_forward = pack_swiglu_pair(gate, up) + group_sizes = jnp.asarray([rows], dtype=jnp.int32) + quantizers = QuantizerFactory.create_set( + scaling_mode=ScalingMode.MXFP8_1D_SCALING, + fwd_dtype=jnp.float8_e4m3fn, + bwd_dtype=jnp.float8_e4m3fn, + is_2x2x=True, + n_groups=experts, + ) + + @jax.jit + def run(d_eo_arg, wo_arg, packed_arg): + casted_d_eo = tex.grouped_quantize( + d_eo_arg, quantizers.dgrad, group_sizes, flatten_axis=-1 + ).get_tensor(TensorUsage.LHS) + casted_wo = tex.grouped_quantize(wo_arg, quantizers.kernel, flatten_axis=-1).get_tensor( + TensorUsage.RHS_TRANS + ) + reference_intermediate = tex.grouped_gemm( + casted_d_eo, + casted_wo, + contracting_dims=((1,), (2,)), + ) + d_row, d_col, scale_row, scale_col, dprob = grouped_gemm_dswiglu_mxfp8( + casted_d_eo.data.reshape(rows, hidden, 1), + casted_wo.data.reshape(experts, intermediate, hidden).transpose(1, 2, 0), + packed_arg.reshape(rows, 2 * intermediate, 1), + casted_d_eo.scale_inv, + casted_wo.scale_inv, + jnp.cumsum(group_sizes), + jnp.ones((rows, 1, 1), dtype=jnp.float32), + output_dtype=quantizers.dgrad.q_dtype, + ) + return reference_intermediate, d_row, d_col, scale_row, scale_col, dprob + + reference_intermediate, d_row, d_col, scale_row, scale_col, dprob = run( + d_eo, wo, packed_forward + ) + jax.block_until_ready((reference_intermediate, d_row, d_col, dprob)) + assert d_row.shape == d_col.shape == (rows, 2 * intermediate, 1) + assert d_row.dtype == d_col.dtype == jnp.float8_e4m3fn + assert scale_row.dtype == scale_col.dtype == jnp.float8_e8m0fnu + assert dprob.shape == (rows, 1, 1) + + sigmoid = jax.nn.sigmoid(gate.astype(jnp.float32)) + swish = gate.astype(jnp.float32) * sigmoid + ref = reference_intermediate.astype(jnp.float32) + d_up = ref * swish + d_gate = ref * up.astype(jnp.float32) * sigmoid * (1 + gate.astype(jnp.float32) * (1 - sigmoid)) + packed_reference = np.asarray(pack_swiglu_pair(d_gate, d_up), dtype=np.float32) + + for is_colwise, payload, scale in ( + (False, d_row, scale_row), + (True, d_col, scale_col), + ): + scaled = ScaledTensorFactory.create_1x( + payload.reshape(-1), + scale, + scaling_mode=ScalingMode.MXFP8_1D_SCALING, + dq_dtype=jnp.bfloat16, + is_colwise=is_colwise, + data_layout="N", + flatten_axis=1, + first_dims=group_sizes, + original_shape=(rows, 2 * intermediate), + pre_swizzled=True, + ) + dequantized = np.asarray(scaled.dequantize()[0], dtype=np.float32) + assert np.all(np.isfinite(dequantized)) + relative_error = np.linalg.norm(dequantized - packed_reference) / np.linalg.norm( + packed_reference + ) + assert relative_error < 0.08 diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index d08765e1842..cfabb91aef5 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -118,7 +118,15 @@ def _read_mp_options(): ) from transformer_engine.jax.flax import _MoEBlock as MoEBlock -from transformer_engine.jax.moe import _ALIGN_SIZE, moe, record_ep_bootstrap_signature_for_moe +from transformer_engine.common.recipe import MXFP8BlockScaling +from transformer_engine.jax import autocast +from transformer_engine.jax.moe import ( + _ALIGN_SIZE, + _CUDNN_CUTEDSL_ALIGN_SIZE, + _use_cudnn_cutedsl_fusion_from_env, + moe, + record_ep_bootstrap_signature_for_moe, +) from transformer_engine.jax.ep import ep_bootstrap from transformer_engine.jax.sharding import MeshResource, global_shard_guard @@ -144,14 +152,28 @@ def _read_mp_options(): ) # Small shapes so the parity tests stay tight on bf16. The block still -# has all four ranks participating in dispatch/combine. +# has all four ranks participating in dispatch/combine. The explicit +# MaxText-shape regression mode is selected by its dedicated test invocation; +# keeping it opt-in avoids making every distributed parity test allocate the +# production-slice buffers. +_MAXTEXT_CUTEDSL_REGRESSION = os.getenv("TE_EP_MOE_MAXTEXT_CUTEDSL_REGRESSION", "0") == "1" DTYPE = jnp.bfloat16 -BATCH = EP_SIZE * FSDP_SIZE * 2 # 8 on 4-GPU, 16 on 8-GPU -SEQ = 32 -HIDDEN = 64 -INTER = 128 -NUM_EXPERTS = 8 -TOPK = 2 +if _MAXTEXT_CUTEDSL_REGRESSION: + # Matches run-dsv3-prod-slice-ep2-fsdp2.sh: global batch 16, sequence + # length 4096, 32 experts, top-k 8, H=1792, expert MLP=2048. + BATCH = 16 + SEQ = 4096 + HIDDEN = 1792 + INTER = 2048 + NUM_EXPERTS = 32 + TOPK = 8 +else: + BATCH = EP_SIZE * FSDP_SIZE * 2 # 8 on 4-GPU, 16 on 8-GPU + SEQ = 32 + HIDDEN = 64 + INTER = 128 + NUM_EXPERTS = 8 + TOPK = 2 # bf16 grouped_gemm + softmax-topk + ep all-to-all stack drifts ~1e-1 vs a # fp32 numpy reference. Keep these tight enough to catch real bugs but @@ -180,7 +202,7 @@ def _read_mp_options(): # ----------------------------------------------------------------------------- -def _compute_worst_case_recv_pr(): +def _compute_worst_case_recv_pr(alignment=_ALIGN_SIZE): """Per-rank recv buffer the bootstrap must reserve. NCCL EP HT expert-major uses one flat recv buffer with variable @@ -194,10 +216,10 @@ def _compute_worst_case_recv_pr(): tokens_per_ep_group = EP_SIZE * max_tokens_per_rank max_local_assignments = tokens_per_ep_group * min(TOPK, num_local_experts) max_nonempty_experts = min(num_local_experts, max_local_assignments) - padded_total_bound = max_local_assignments + (_ALIGN_SIZE - 1) * max_nonempty_experts - aligned_total_bound = ((padded_total_bound + _ALIGN_SIZE - 1) // _ALIGN_SIZE) * _ALIGN_SIZE + padded_total_bound = max_local_assignments + (alignment - 1) * max_nonempty_experts + aligned_total_bound = ((padded_total_bound + alignment - 1) // alignment) * alignment per_expert_bound = ( - num_local_experts * ((tokens_per_ep_group + _ALIGN_SIZE - 1) // _ALIGN_SIZE) * _ALIGN_SIZE + num_local_experts * ((tokens_per_ep_group + alignment - 1) // alignment) * alignment ) return min(per_expert_bound, aligned_total_bound) @@ -217,7 +239,9 @@ def mesh(): num_procs = jax.process_count() max_tokens_per_rank = (BATCH // num_procs) * SEQ - recv_capacity_per_rank = _compute_worst_case_recv_pr() + fusion_enabled = _use_cudnn_cutedsl_fusion_from_env() + alignment = _CUDNN_CUTEDSL_ALIGN_SIZE if fusion_enabled else _ALIGN_SIZE + recv_capacity_per_rank = _compute_worst_case_recv_pr(alignment) # Eager bootstrap: ep_bootstrap does a host-side NCCL UID allgather # and cannot run from inside jax.jit. Sized to the worst-case recv_pr @@ -363,6 +387,7 @@ def _make_block( aux_loss_coeff=0.0, use_expert_routing_bias=False, score_function="softmax", + scaling_factor=1.0, expert_bias_init=None, ): kwargs = dict( @@ -374,6 +399,7 @@ def _make_block( aux_loss_coeff=aux_loss_coeff, use_expert_routing_bias=use_expert_routing_bias, score_function=score_function, + scaling_factor=scaling_factor, dtype=DTYPE, ) # Custom expert_bias_init lets tests inject a non-zero expert_bias without @@ -679,6 +705,182 @@ def loss_fn(params, x): ) +class TestTeEpMoeCudnnCutedslFusion: + """End-to-end MXFP8 coverage for the opt-in FC1+SwiGLU+quant fusion.""" + + @pytest.mark.parametrize("apply_topk_weights_early", [False, True]) + def test_mxfp8_forward_and_backward(self, mesh, apply_topk_weights_early): + if not _use_cudnn_cutedsl_fusion_from_env(): + pytest.skip("run separately with NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION=1") + block = _make_block(apply_topk_weights_early=apply_topk_weights_early) + x = _make_inputs(jax.random.PRNGKey(30)) + mesh_resource = MeshResource(ep_resource=EP_AXIS, fsdp_resource=FSDP_AXIS) + + with _ctx(mesh), autocast( + enabled=True, + recipe=MXFP8BlockScaling(), + mesh_resource=mesh_resource, + ): + x_sh = _shard_inputs(x, mesh) + variables = jax.jit(block.init)(jax.random.PRNGKey(31), x_sh) + fused, _ = jax.jit(block.apply)(variables, x_sh) + + def loss_fn(vars_arg, x_arg): + output, _ = block.apply(vars_arg, x_arg) + return jnp.mean(output.astype(jnp.float32) ** 2) + + grads, grad_x = jax.jit(jax.grad(loss_fn, argnums=(0, 1)))(variables, x_sh) + jax.block_until_ready((fused, grads, grad_x)) + + # Do not compile the unfused EP path in this CuTeDSL process. The C++ + # backend caches the first layer alignment process-wide, so a + # 256-aligned fused run and a 128-aligned unfused run must live in + # separate Python process groups. run_te_ep_moe.sh covers the unfused + # CUDA C++ path in its ordinary phase. + fused_np = _to_global_numpy(fused, mesh).astype(np.float32) + assert np.all(np.isfinite(fused_np)) + + params_np = _params_global_numpy(variables, mesh) + reference, _ = _pure_jax_moe_reference( + jnp.asarray(jax.device_get(x)), + jnp.asarray(params_np["gate_kernel"]), + jnp.asarray(params_np["wi_0"]), + jnp.asarray(params_np["wi_1"]), + jnp.asarray(params_np["wo"]), + num_experts=NUM_EXPERTS, + num_experts_per_tok=TOPK, + ) + reference_np = np.asarray(jax.device_get(reference), dtype=np.float32) + relative_error = np.linalg.norm(fused_np - reference_np) / np.linalg.norm(reference_np) + assert relative_error < 0.2 + + for name in ("gate_kernel", "wi_0", "wi_1", "wo"): + grad = _to_global_numpy(_unwrap(grads["params"][name]), mesh).astype(np.float32) + assert np.all(np.isfinite(grad)), f"{name} fused MXFP8 grad has NaN/Inf" + assert np.any(grad != 0), f"{name} fused MXFP8 grad is identically zero" + grad_x_np = _to_global_numpy(grad_x, mesh).astype(np.float32) + assert np.all(np.isfinite(grad_x_np)) + assert np.any(grad_x_np != 0) + + # A finite gradient is not sufficient: a layout error can produce a + # finite but numerically invalid VJP that corrupts parameters on the + # first optimizer update. Compare all learnable gradients to the same + # pure-JAX reference used by the non-quantized VJP tests. + def reference_loss(params, x_arg): + output, _ = _pure_jax_moe_reference( + x_arg, + params["gate_kernel"], + params["wi_0"], + params["wi_1"], + params["wo"], + num_experts=NUM_EXPERTS, + num_experts_per_tok=TOPK, + ) + return jnp.mean(output.astype(jnp.float32) ** 2) + + reference_params = { + name: jnp.asarray(params_np[name]) for name in ("gate_kernel", "wi_0", "wi_1", "wo") + } + reference_grads = jax.jit(jax.grad(reference_loss))( + reference_params, jnp.asarray(jax.device_get(x)) + ) + for name, reference_grad in reference_grads.items(): + fused_grad = _to_global_numpy(_unwrap(grads["params"][name]), mesh).astype(np.float32) + reference_grad = np.asarray(jax.device_get(reference_grad), dtype=np.float32) + relative_error = np.linalg.norm(fused_grad - reference_grad) / max( + np.linalg.norm(reference_grad), 1e-12 + ) + assert ( + relative_error < 0.35 + ), f"{name} fused MXFP8 VJP relative error {relative_error:.4f} exceeds 0.35" + + # Exercise the failure mode seen in MaxText: apply one optimizer-like + # update and require the next forward pass to remain finite. + updated_variables = jax.tree_util.tree_map( + lambda param, grad: param - jnp.asarray(1e-3, param.dtype) * grad.astype(param.dtype), + variables, + grads, + ) + with _ctx(mesh), autocast( + enabled=True, + recipe=MXFP8BlockScaling(), + mesh_resource=mesh_resource, + ): + updated_output, _ = jax.jit(block.apply)(updated_variables, _shard_inputs(x, mesh)) + updated_output.block_until_ready() + updated_output_np = _to_global_numpy(updated_output, mesh).astype(np.float32) + assert np.all(np.isfinite(updated_output_np)), "post-update output has NaN/Inf" + + def test_maxtext_shape_vjp_update_stays_finite(self, mesh): + """Regression for the NaN observed after MaxText's first update.""" + if not _use_cudnn_cutedsl_fusion_from_env(): + pytest.skip("run separately with NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION=1") + if not _MAXTEXT_CUTEDSL_REGRESSION: + pytest.skip("set TE_EP_MOE_MAXTEXT_CUTEDSL_REGRESSION=1 for the large regression") + + block = _make_block( + score_function="sigmoid", + use_expert_routing_bias=True, + scaling_factor=2.5, + ) + x = _make_inputs(jax.random.PRNGKey(40)) + mesh_resource = MeshResource(ep_resource=EP_AXIS, fsdp_resource=FSDP_AXIS) + with _ctx(mesh), autocast( + enabled=True, + recipe=MXFP8BlockScaling(), + mesh_resource=mesh_resource, + ): + x_sh = _shard_inputs(x, mesh) + variables = jax.jit(block.init)(jax.random.PRNGKey(41), x_sh) + + compiled_forward = jax.jit(block.apply) + forward_0, _ = compiled_forward(variables, x_sh) + forward_1, _ = compiled_forward(variables, x_sh) + jax.block_until_ready((forward_0, forward_1)) + forward_0_local = np.asarray(jax.device_get(forward_0.addressable_data(0))) + forward_1_local = np.asarray(jax.device_get(forward_1.addressable_data(0))) + assert np.all(np.isfinite(forward_0_local)), "first repeated forward has NaN/Inf" + assert np.all(np.isfinite(forward_1_local)), "second repeated forward has NaN/Inf" + + def loss_fn(vars_arg, x_arg): + output, _ = block.apply(vars_arg, x_arg) + return jnp.mean(output.astype(jnp.float32) ** 2) + + def train_step(vars_arg, x_arg, learning_rate): + loss, grads = jax.value_and_grad(loss_fn)(vars_arg, x_arg) + updated_vars = jax.tree_util.tree_map( + lambda param, grad: param + - learning_rate.astype(param.dtype) * grad.astype(param.dtype), + vars_arg, + grads, + ) + return updated_vars, loss, grads + + compiled_train_step = jax.jit(train_step) + # MaxText's two-step warmup uses LR=0 at step 0. The second + # invocation therefore checks that replaying the same compiled + # kernel/VJP is safe even before any parameter value changes. + variables, loss_0, grads = compiled_train_step( + variables, x_sh, jnp.asarray(0.0, jnp.float32) + ) + jax.block_until_ready((loss_0, grads)) + assert np.isfinite(float(loss_0.addressable_data(0))), "step-0 loss has NaN/Inf" + for path, grad in jax.tree_util.tree_leaves_with_path(grads): + grad_local = np.asarray(jax.device_get(_unwrap(grad).addressable_data(0))) + assert np.all(np.isfinite(grad_local)), f"gradient {path} has NaN/Inf" + variables, loss_1, _ = compiled_train_step( + variables, x_sh, jnp.asarray(1.5e-5, jnp.float32) + ) + jax.block_until_ready(loss_1) + assert np.isfinite(float(loss_1.addressable_data(0))), "step-1 loss has NaN/Inf" + + updated_output, _ = jax.jit(block.apply)(variables, x_sh) + updated_output.block_until_ready() + + updated_local = np.asarray(jax.device_get(updated_output.addressable_data(0))) + assert np.all(np.isfinite(updated_local)), "MaxText-shape post-update output has NaN/Inf" + + class TestTeEpMoeAuxLoss: """Aux-loss path. Consolidated into: * ``test_aux_loss``: one run that checks the returned scalar's diff --git a/transformer_engine/jax/cutedsl_extensions/__init__.py b/transformer_engine/jax/cutedsl_extensions/__init__.py new file mode 100644 index 00000000000..27730a4e555 --- /dev/null +++ b/transformer_engine/jax/cutedsl_extensions/__init__.py @@ -0,0 +1,4 @@ +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. +"""JAX wrappers for kernels authored with the CUTLASS CuTe DSL.""" diff --git a/transformer_engine/jax/cutedsl_extensions/moe.py b/transformer_engine/jax/cutedsl_extensions/moe.py new file mode 100644 index 00000000000..b51b2892e6c --- /dev/null +++ b/transformer_engine/jax/cutedsl_extensions/moe.py @@ -0,0 +1,576 @@ +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. +"""JAX binding for cuDNN frontend's MXFP8 grouped-GEMM SwiGLU kernel. + +The kernel is intentionally imported from the installed +``nvidia-cudnn-frontend`` distribution. The public cuDNN frontend package +initializers import the Torch API wrappers eagerly, so this module loads the +kernel source directly under a private namespace. No kernel implementation is +vendored into Transformer Engine. +""" + +from __future__ import annotations + +from functools import lru_cache +import importlib.metadata +import importlib.util +import sys +import threading +import types +from typing import Any + +import jax +import jax.numpy as jnp + + +_CUDNN_FRONTEND_DISTRIBUTION = "nvidia-cudnn-frontend" +_PRIVATE_PACKAGE = "_transformer_engine_cudnn_grouped_gemm" +_LOAD_LOCK = threading.Lock() + + +def _namespace_package(name: str, path) -> types.ModuleType: + module = types.ModuleType(name) + module.__package__ = name + module.__path__ = [str(path)] + sys.modules[name] = module + return module + + +def _load_source_module(name: str, path): + spec = importlib.util.spec_from_file_location(name, path) + if spec is None or spec.loader is None: + raise ImportError(f"Could not create an import spec for {path}") + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + try: + spec.loader.exec_module(module) + except Exception: + sys.modules.pop(name, None) + raise + return module + + +def _patch_atomic_add_float32_for_cutlass_dsl(utils_module) -> None: + """Patch cuDNN frontend utility source for newer CUTLASS DSL bindings.""" + try: + import cutlass + from cutlass._mlir.dialects import nvvm + from cutlass.cute.typing import Float32 + except ImportError: + return + + def atomic_add_float32(ptr, value: Float32, *, loc=None, ip=None) -> Float32: + a = value.ir_value(loc=loc, ip=ip) + try: + old_value = nvvm.atomicrmw( + op=cutlass._mlir.dialects.nvvm.AtomicOpKind.FADD, + ptr=ptr, + a=a, + res=a.type, + loc=loc, + ip=ip, + ) + except TypeError as exc: + if "res" not in str(exc): + raise + old_value = nvvm.atomicrmw( + op=cutlass._mlir.dialects.nvvm.AtomicOpKind.FADD, + ptr=ptr, + a=a, + loc=loc, + ip=ip, + ) + return Float32(old_value) + + utils_module.atomic_add_float32 = atomic_add_float32 + + +def _ensure_torch_annotation_stub() -> bool: + # Do not shadow a real Torch installation. If it is installed, the + # unmodified cuDNN frontend source can import it normally; the sentinel is + # only for JAX-only environments where no Torch module exists. + try: + importlib.metadata.distribution("torch") + torch_installed = True + except importlib.metadata.PackageNotFoundError: + torch_installed = False + torch_stubbed = "torch" not in sys.modules and not torch_installed + if torch_stubbed: + torch_stub = types.ModuleType("torch") + torch_stub.Tensor = object + torch_stub.float4_e2m1fn_x2 = object() + torch_stub.__transformer_engine_cutedsl_stub__ = True + sys.modules["torch"] = torch_stub + return torch_stubbed + + +def _load_grouped_gemm_kernel_module(subpackage: str, filename: str): + distribution = importlib.metadata.distribution(_CUDNN_FRONTEND_DISTRIBUTION) + grouped_root = distribution.locate_file("cudnn/grouped_gemm") + utils_path = grouped_root / "utils.py" + kernel_root = grouped_root / subpackage + kernel_path = kernel_root / filename + if not utils_path.is_file() or not kernel_path.is_file(): + raise ImportError( + f"{_CUDNN_FRONTEND_DISTRIBUTION} {distribution.version} does not contain " + f"the {subpackage} CuTeDSL kernel sources" + ) + + _namespace_package(_PRIVATE_PACKAGE, grouped_root) + private_subpackage = f"{_PRIVATE_PACKAGE}.{subpackage}" + _namespace_package(private_subpackage, kernel_root) + + torch_stubbed = _ensure_torch_annotation_stub() + try: + utils_module = _load_source_module(f"{_PRIVATE_PACKAGE}.utils", utils_path) + _patch_atomic_add_float32_for_cutlass_dsl(utils_module) + module_name = f"{private_subpackage}.{filename.removesuffix('.py')}" + return _load_source_module(module_name, kernel_path) + except Exception: + if torch_stubbed: + sys.modules.pop("torch", None) + raise + + +@lru_cache(maxsize=1) +def load_grouped_gemm_swiglu_kernel(): + """Load the cuDNN frontend forward kernel class without its Torch API.""" + with _LOAD_LOCK: + kernel_module = _load_grouped_gemm_kernel_module( + "grouped_gemm_swiglu", "grouped_gemm_swiglu_quant.py" + ) + try: + return kernel_module.BlockScaledContiguousGroupedGemmKernel + except AttributeError as exc: + distribution = importlib.metadata.distribution(_CUDNN_FRONTEND_DISTRIBUTION) + raise ImportError( + f"{_CUDNN_FRONTEND_DISTRIBUTION} {distribution.version} has an " + "incompatible grouped-GEMM SwiGLU kernel API" + ) from exc + + +@lru_cache(maxsize=1) +def load_grouped_gemm_dswiglu_kernel(): + """Load the cuDNN frontend dSwiGLU backward kernel class.""" + with _LOAD_LOCK: + kernel_module = _load_grouped_gemm_kernel_module( + "grouped_gemm_dswiglu", "grouped_gemm_dswiglu_quant.py" + ) + try: + return kernel_module.BlockScaledContiguousGroupedGemmKernel + except AttributeError as exc: + distribution = importlib.metadata.distribution(_CUDNN_FRONTEND_DISTRIBUTION) + raise ImportError( + f"{_CUDNN_FRONTEND_DISTRIBUTION} {distribution.version} has an " + "incompatible grouped-GEMM dSwiGLU kernel API" + ) from exc + + +def pack_swiglu_pair(gate: jax.Array, up: jax.Array) -> jax.Array: + """Interleave 32-column gate/up blocks as required by the kernel.""" + if gate.shape != up.shape: + raise ValueError(f"gate shape {gate.shape} must match up shape {up.shape}") + if gate.shape[-1] % 32: + raise ValueError(f"SwiGLU intermediate dimension {gate.shape[-1]} must be divisible by 32") + blocks = gate.shape[-1] // 32 + return jnp.stack( + ( + gate.reshape(*gate.shape[:-1], blocks, 32), + up.reshape(*up.shape[:-1], blocks, 32), + ), + axis=-2, + ).reshape(*gate.shape[:-1], 2 * gate.shape[-1]) + + +def unpack_swiglu_pair(interleaved: jax.Array) -> tuple[jax.Array, jax.Array]: + """Undo :func:`pack_swiglu_pair`.""" + if interleaved.shape[-1] % 64: + raise ValueError( + f"Interleaved SwiGLU dimension {interleaved.shape[-1]} must be divisible by 64" + ) + intermediate = interleaved.shape[-1] // 2 + blocks = intermediate // 32 + paired = interleaved.reshape(*interleaved.shape[:-1], blocks, 2, 32) + return ( + paired[..., 0, :].reshape(*interleaved.shape[:-1], intermediate), + paired[..., 1, :].reshape(*interleaved.shape[:-1], intermediate), + ) + + +def _ceil_div(x: int, y: int) -> int: + return (x + y - 1) // y + + +@lru_cache(maxsize=None) +def _make_launcher( + expert_count: int, + sf_vec_size: int, + mma_tiler_m: int, + mma_tiler_n: int, +): + try: + import cutlass + from cutlass import cute + except ImportError as exc: + raise ImportError( + "NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION=1 requires nvidia-cutlass-dsl" + ) from exc + + kernel_cls = load_grouped_gemm_swiglu_kernel() + use_2cta_instrs = mma_tiler_m == 256 + cluster_shape = (2, 1) if use_2cta_instrs else (1, 1) + kernel = kernel_cls( + sf_vec_size=sf_vec_size, + acc_dtype=cutlass.Float32, + use_2cta_instrs=use_2cta_instrs, + mma_tiler_mn=(mma_tiler_m, mma_tiler_n), + cluster_shape_mn=cluster_shape, + vector_f32=False, + generate_sfd=True, + # TE's grouped colwise tensor stores an independently padded scale + # segment for every expert. The kernel's default colwise SFD layout + # treats the concatenated M dimension as one matrix, which makes the + # FC2 wgrad read another expert's scales (and eventually emit NaNs). + discrete_col_sfd=True, + expert_cnt=expert_count, + use_mono_increase_expert_idx=True, + ) + max_active_clusters = cutlass.utils.HardwareInfo().get_max_active_clusters( + cluster_shape[0] * cluster_shape[1] + ) + + @cute.jit + def launch( + stream, + a, + b, + sfa, + sfb, + padded_offsets, + alpha, + prob, + norm_const, + c, + d, + d_col, + sfd_row, + sfd_col, + ): + kernel( + a, + b, + c, + d, + d_col, + sfa, + sfb, + sfd_row, + sfd_col, + None, + norm_const, + padded_offsets, + alpha, + prob, + max_active_clusters, + stream, + ) + + return launch + + +def grouped_gemm_swiglu_mxfp8( + a: jax.Array, + b: jax.Array, + sfa: jax.Array, + sfb: jax.Array, + padded_offsets: jax.Array, + prob: jax.Array, + *, + compute_dtype: Any, + output_dtype: Any, +) -> tuple[jax.Array, jax.Array, jax.Array, jax.Array, jax.Array]: + """Call cuDNN frontend's grouped MXFP8 GEMM+SwiGLU+quant kernel. + + Args: + a: Quantized activation payload with physical shape ``[M, K, 1]``. + b: Quantized, block-interleaved colwise weights with physical shape + ``[E, N, K]``. + sfa/sfb: Pre-swizzled E8M0 inverse-scale buffers. + padded_offsets: Exclusive, 256-aligned end offset for each expert. + prob: Per-row multiplier, physical shape ``[M, 1, 1]``. + compute_dtype: Unquantized combined-projection dtype. + output_dtype: MXFP8 payload dtype for the quantized SwiGLU output. + + Returns: + Raw combined projection, rowwise payload, colwise payload, rowwise + inverse scales, and colwise inverse scales. + """ + try: + from cutlass.jax import TensorSpec, cutlass_call + except ImportError as exc: + raise ImportError( + "NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION=1 requires CUTLASS JAX bindings" + ) from exc + + if a.ndim != 3 or a.shape[-1] != 1: + raise ValueError(f"Expected A[M,K,1], got {a.shape}") + if b.ndim != 3: + raise ValueError(f"Expected physical B[E,N,K], got {b.shape}") + expert_count, n, k_b = b.shape + m, k_a, _ = a.shape + if k_a != k_b: + raise ValueError(f"A K={k_a} does not match B K={k_b}") + if n % 2: + raise ValueError(f"Combined SwiGLU N={n} must be even") + intermediate = n // 2 + + # The public kernel requires 256-aligned expert ranges. M is static at + # trace time, while the individual offsets remain runtime values. + if m % 256: + raise ValueError(f"Padded activation rows M={m} must be divisible by 256") + + sf_vec_size = 32 + row_scale_size = ( + 32 * 4 * _ceil_div(m, 128) * 4 * _ceil_div(_ceil_div(intermediate, sf_vec_size), 4) + ) + col_scale_size = ( + 32 * 4 * _ceil_div(intermediate, 128) * 4 * _ceil_div(_ceil_div(m, sf_vec_size), 4) + ) + outputs = ( + jax.ShapeDtypeStruct((m, n, 1), compute_dtype), + jax.ShapeDtypeStruct((m, intermediate, 1), output_dtype), + jax.ShapeDtypeStruct((m, intermediate, 1), output_dtype), + jax.ShapeDtypeStruct((row_scale_size,), jnp.float8_e8m0fnu), + jax.ShapeDtypeStruct((col_scale_size,), jnp.float8_e8m0fnu), + ) + launcher = _make_launcher(expert_count, sf_vec_size, 256, 256) + call = cutlass_call( + launcher, + output_shape_dtype=outputs, + input_spec=( + # Singleton L must be outermost so K, rather than L, is the + # leading (stride-1) dimension seen by the kernel. + TensorSpec(layout=(1, 0, 2)), + # physical colwise [E,N,K] -> logical K-major [N,K,E] + TensorSpec(mode=(1, 2, 0)), + TensorSpec(), + TensorSpec(), + TensorSpec(), + TensorSpec(), + TensorSpec(), + TensorSpec(), + ), + output_spec=( + TensorSpec(layout=(1, 0, 2)), + TensorSpec(layout=(1, 0, 2)), + TensorSpec(layout=(1, 0, 2)), + TensorSpec(), + TensorSpec(), + ), + allow_cuda_graph=True, + ) + alpha = jnp.ones((expert_count,), dtype=jnp.float32) + # With norm_const=1 the generated E8M0 factors are directly the + # scale-inverse values expected by TE's MXFP8 tensor representation. + norm_const = jnp.ones((1,), dtype=jnp.float32) + return call( + a, + b, + sfa.reshape(-1), + sfb.reshape(-1), + padded_offsets.astype(jnp.int32), + alpha, + prob.astype(jnp.float32), + norm_const, + ) + + +@lru_cache(maxsize=None) +def _make_dswiglu_launcher( + expert_count: int, + sf_vec_size: int, + mma_tiler_m: int, + mma_tiler_n: int, +): + try: + import cutlass + from cutlass import cute + except ImportError as exc: + raise ImportError( + "NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION=1 requires nvidia-cutlass-dsl" + ) from exc + + kernel_cls = load_grouped_gemm_dswiglu_kernel() + use_2cta_instrs = mma_tiler_m == 256 + cluster_shape = (2, 1) if use_2cta_instrs else (1, 1) + kernel = kernel_cls( + sf_vec_size=sf_vec_size, + acc_dtype=cutlass.Float32, + use_2cta_instrs=use_2cta_instrs, + mma_tiler_mn=(mma_tiler_m, mma_tiler_n), + cluster_shape_mn=cluster_shape, + vectorized_f32=False, + discrete_col_sfd=True, + expert_cnt=expert_count, + use_mono_increase_expert_idx=True, + ) + max_active_clusters = cutlass.utils.HardwareInfo().get_max_active_clusters( + cluster_shape[0] * cluster_shape[1] + ) + + @cute.jit + def launch( + stream, + a, + b, + c, + sfa, + sfb, + padded_offsets, + alpha, + beta, + prob, + norm_const, + d, + d_col, + sfd_row, + sfd_col, + dprob, + ): + kernel( + a, + b, + c, + d, + d_col, + sfa, + sfb, + sfd_row, + sfd_col, + None, + norm_const, + padded_offsets, + alpha, + beta, + prob, + dprob, + max_active_clusters, + stream, + ) + + return launch + + +def grouped_gemm_dswiglu_mxfp8( + a: jax.Array, + b: jax.Array, + c: jax.Array, + sfa: jax.Array, + sfb: jax.Array, + padded_offsets: jax.Array, + prob: jax.Array, + *, + output_dtype: Any, +) -> tuple[jax.Array, jax.Array, jax.Array, jax.Array, jax.Array]: + """Call cuDNN frontend's grouped MXFP8 GEMM+dSwiGLU+quant kernel. + + Args: + a: Quantized upstream gradient payload with physical shape ``[M, K, 1]``. + b: Quantized FC2 weights with physical shape ``[N, K, E]``. + c: Forward packed gate/up activations with shape ``[M, 2N, 1]``. + sfa/sfb: Pre-swizzled E8M0 inverse-scale buffers. + padded_offsets: Exclusive, 256-aligned end offset for each expert. + prob: Per-row multiplier from early top-k weighting, or ones. + output_dtype: MXFP8 payload dtype for the quantized packed output. + + Returns: + Rowwise and colwise packed dSwiGLU payloads, rowwise and colwise + inverse scales, and the per-row probability gradient. + """ + try: + from cutlass.jax import TensorSpec, cutlass_call + except ImportError as exc: + raise ImportError( + "NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION=1 requires CUTLASS JAX bindings" + ) from exc + + if a.ndim != 3 or a.shape[-1] != 1: + raise ValueError(f"Expected A[M,K,1], got {a.shape}") + if b.ndim != 3: + raise ValueError(f"Expected B[N,K,E], got {b.shape}") + if c.ndim != 3 or c.shape[-1] != 1: + raise ValueError(f"Expected C[M,2N,1], got {c.shape}") + m, k_a, _ = a.shape + n, k_b, expert_count = b.shape + if k_a != k_b: + raise ValueError(f"A K={k_a} does not match B K={k_b}") + if c.shape != (m, 2 * n, 1): + raise ValueError(f"Expected C shape {(m, 2 * n, 1)}, got {c.shape}") + if m % 256: + raise ValueError(f"Padded activation rows M={m} must be divisible by 256") + + sf_vec_size = 32 + row_scale_size = 32 * 4 * _ceil_div(m, 128) * 4 * _ceil_div( + _ceil_div(2 * n, sf_vec_size), 4 + ) + col_scale_size = 32 * 4 * _ceil_div(2 * n, 128) * 4 * _ceil_div( + _ceil_div(m, sf_vec_size), 4 + ) + outputs = ( + jax.ShapeDtypeStruct((m, 2 * n, 1), output_dtype), + jax.ShapeDtypeStruct((m, 2 * n, 1), output_dtype), + jax.ShapeDtypeStruct((row_scale_size,), jnp.float8_e8m0fnu), + jax.ShapeDtypeStruct((col_scale_size,), jnp.float8_e8m0fnu), + jax.ShapeDtypeStruct((m, 1, 1), jnp.float32), + ) + launcher = _make_dswiglu_launcher(expert_count, sf_vec_size, 256, 256) + call = cutlass_call( + launcher, + output_shape_dtype=outputs, + input_spec=( + TensorSpec(layout=(1, 0, 2)), + TensorSpec(layout=(1, 0, 2)), + TensorSpec(layout=(1, 0, 2)), + TensorSpec(), + TensorSpec(), + TensorSpec(), + TensorSpec(), + TensorSpec(), + TensorSpec(), + TensorSpec(), + ), + output_spec=( + TensorSpec(layout=(1, 0, 2)), + TensorSpec(layout=(1, 0, 2)), + TensorSpec(), + TensorSpec(), + TensorSpec(layout=(1, 0, 2)), + ), + allow_cuda_graph=True, + ) + alpha = jnp.ones((expert_count,), dtype=jnp.float32) + beta = jnp.ones((expert_count,), dtype=jnp.float32) + norm_const = jnp.ones((1,), dtype=jnp.float32) + return call( + a, + b, + c, + sfa.reshape(-1), + sfb.reshape(-1), + padded_offsets.astype(jnp.int32), + alpha, + beta, + prob.astype(jnp.float32), + norm_const, + ) + + +__all__ = [ + "grouped_gemm_dswiglu_mxfp8", + "grouped_gemm_swiglu_mxfp8", + "load_grouped_gemm_dswiglu_kernel", + "load_grouped_gemm_swiglu_kernel", + "pack_swiglu_pair", + "unpack_swiglu_pair", +] diff --git a/transformer_engine/jax/flax/moe.py b/transformer_engine/jax/flax/moe.py index 3629346e33c..c345ab23519 100644 --- a/transformer_engine/jax/flax/moe.py +++ b/transformer_engine/jax/flax/moe.py @@ -39,7 +39,7 @@ from ..moe import moe from ..router import ScoreFunction -from ..sharding import get_active_resource_axis +from ..sharding import get_active_resource_axis, get_mesh_axis_size from .module import TransformerEngineBase PRNGKey = Any @@ -105,10 +105,9 @@ class _MoEBlock(TransformerEngineBase): *inside* each shard before ``ep_combine`` (saves one global reduction at the cost of an extra broadcast). Default ``False``. - The per-expert dispatch-slot alignment is fixed internally at 128 - tokens (see ``moe._ALIGN_SIZE``) -- the value required by NCCL EP - HT and satisfied by every current TE grouped-GEMM recipe -- and is - therefore not exposed as a per-instance knob. + The default per-expert dispatch-slot alignment is 128 tokens. The + opt-in cuDNN CuTeDSL MXFP8 fusion uses its required 256-token + alignment; neither value is exposed as a per-instance knob. dtype : jnp.dtype Compute / parameter dtype. @@ -190,6 +189,11 @@ def __call__(self, inputs: Array) -> Tuple[Array, Optional[Array]]: ), f"_MoEBlock expects [batch, sequence, hidden] input, got shape {inputs.shape}" _, _, hidden_size = inputs.shape + ep_axis = get_active_resource_axis("ep_resource") + num_local_experts = self.num_experts // get_mesh_axis_size(ep_axis) + fc1_quantizer_set = self.generate_quantizer_set("_moe_fc1", n_groups=num_local_experts) + fc2_quantizer_set = self.generate_quantizer_set("_moe_fc2", n_groups=num_local_experts) + # Param registrations -- must run OUTSIDE any JAX transform that # alters the variable scope (e.g. shard_map). The functional # ``moe(...)`` opens its own shard_map internally for the EP @@ -249,8 +253,6 @@ def __call__(self, inputs: Array) -> Tuple[Array, Optional[Array]]: jnp.float32, ) - ep_axis = get_active_resource_axis("ep_resource") - return moe( inputs, gate_kernel, @@ -261,6 +263,8 @@ def __call__(self, inputs: Array) -> Tuple[Array, Optional[Array]]: wi_1_bias, wo_bias, expert_bias, + fc1_quantizer_set, + fc2_quantizer_set, num_experts=self.num_experts, num_experts_per_tok=self.num_experts_per_tok, activation_type=self.activation_type, diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index 1a4c1ec6ccf..113472176a0 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -33,6 +33,7 @@ from functools import partial from typing import Any, Optional, Tuple, Union +import os import warnings import flax.struct @@ -63,6 +64,116 @@ # TE grouped-GEMM recipes (bf16/fp16/fp8/mxfp8) are satisfied by the # same 128-token tile, so a single constant covers every supported path. _ALIGN_SIZE = 128 +_CUDNN_CUTEDSL_ALIGN_SIZE = 256 +_CUDNN_CUTEDSL_ENV = "NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION" + + +def _use_cudnn_cutedsl_fusion_from_env() -> bool: + value = os.getenv(_CUDNN_CUTEDSL_ENV, "0") + if value not in ("0", "1"): + raise ValueError(f"{_CUDNN_CUTEDSL_ENV} must be '0' or '1', got {value!r}") + return value == "1" + + +def _cudnn_cutedsl_fusion_rejection_reasons( + x, + wi_0, + wi_1, + wi_0_bias, + wi_1_bias, + fc1_quantizer_set, + fc2_quantizer_set, + *, + num_experts, + activation_type, + ep_axis, +) -> list[str]: + """Return reasons this call cannot use the CuTeDSL fusion, or an empty list.""" + from transformer_engine_jax import get_device_compute_capability + + from .cutedsl_extensions.moe import ( + load_grouped_gemm_dswiglu_kernel, + load_grouped_gemm_swiglu_kernel, + ) + from .quantize import GroupedQuantizer, ScalingMode + + errors = [] + try: + compute_capability = get_device_compute_capability(0) + except RuntimeError as exc: + errors.append(f"could not query GPU compute capability: {exc}") + else: + if compute_capability != 100: + errors.append(f"requires an SM100 GPU, got SM{compute_capability}") + if str(activation_type).lower() != "silu": + errors.append("requires activation_type='silu'") + if wi_0_bias is not None or wi_1_bias is not None: + errors.append("does not support FC1 gate/up bias") + if wi_0.shape != wi_1.shape or wi_0.ndim != 3: + errors.append(f"requires matching rank-3 wi_0/wi_1, got {wi_0.shape}/{wi_1.shape}") + elif wi_0.shape[-1] % 32: + errors.append(f"requires intermediate size divisible by 32, got {wi_0.shape[-1]}") + if x.dtype not in (jnp.bfloat16, jnp.float16): + errors.append(f"requires BF16 or FP16 activations, got {x.dtype}") + + required_quantizers = { + "fc1.x": fc1_quantizer_set.x, + "fc1.kernel": fc1_quantizer_set.kernel, + "fc1.dgrad": fc1_quantizer_set.dgrad, + "fc2.x": fc2_quantizer_set.x, + "fc2.kernel": fc2_quantizer_set.kernel, + "fc2.dgrad": fc2_quantizer_set.dgrad, + } + for name, quantizer in required_quantizers.items(): + if not isinstance(quantizer, GroupedQuantizer): + errors.append(f"requires a grouped MXFP8 quantizer for {name}") + elif quantizer.scaling_mode != ScalingMode.MXFP8_1D_SCALING: + errors.append(f"requires MXFP8_1D_SCALING for {name}") + elif not quantizer.q_layout.is_rowwise_colwise: + errors.append(f"requires rowwise+colwise quantization for {name}") + + if all(isinstance(q, GroupedQuantizer) for q in required_quantizers.values()): + if fc1_quantizer_set.x.q_dtype != fc1_quantizer_set.kernel.q_dtype: + errors.append("requires identical FC1 activation and weight MXFP8 payload dtypes") + if fc2_quantizer_set.dgrad.q_dtype != fc2_quantizer_set.kernel.q_dtype: + errors.append( + "requires identical FC2 dgrad and weight MXFP8 payload dtypes for " + "dSwiGLU backward fusion" + ) + supported = (jnp.float8_e4m3fn, jnp.float8_e5m2) + for name, quantizer in required_quantizers.items(): + if quantizer.q_dtype not in supported: + errors.append(f"unsupported MXFP8 payload dtype {quantizer.q_dtype} for {name}") + + mesh = _get_mesh() + if mesh is not None and not mesh.empty and ep_axis in mesh.shape: + num_local_experts = num_experts // mesh.shape[ep_axis] + if num_local_experts > 1024: + errors.append(f"requires at most 1024 local experts, got {num_local_experts}") + + try: + import cutlass.jax as cutlass_jax + except ImportError as exc: + errors.append(f"nvidia-cutlass-dsl with JAX bindings is required: {exc}") + else: + try: + if not cutlass_jax.is_available(): + errors.append( + "cutlass.jax.is_available() is false; check cute_dsl_runtime.so discovery" + ) + except (AttributeError, RuntimeError) as exc: + errors.append(f"CUTLASS JAX runtime check failed: {exc}") + + for kernel_name, loader in ( + ("SwiGLU forward", load_grouped_gemm_swiglu_kernel), + ("dSwiGLU backward", load_grouped_gemm_dswiglu_kernel), + ): + try: + loader() + except (ImportError, ModuleNotFoundError, RuntimeError) as exc: + errors.append(f"could not load the cuDNN frontend {kernel_name} kernel: {exc}") + + return errors def _with_sharding_constraint_cast_bwd(x: jnp.ndarray, sharding) -> jnp.ndarray: @@ -232,14 +343,10 @@ def _pack_grouped_tensor(tensor): else: scale_inv = tensor.scale_inv first_dims = ( - tensor.first_dims - if tensor.first_dims is not None - else jnp.empty((0,), dtype=jnp.int32) + tensor.first_dims if tensor.first_dims is not None else jnp.empty((0,), dtype=jnp.int32) ) last_dims = ( - tensor.last_dims - if tensor.last_dims is not None - else jnp.empty((0,), dtype=jnp.int32) + tensor.last_dims if tensor.last_dims is not None else jnp.empty((0,), dtype=jnp.int32) ) return tensor.data, scale_inv, tensor.amax, first_dims, last_dims @@ -311,6 +418,7 @@ def _ffn_fwd_per_shard( num_local_experts: int, activation_type: str, apply_topk_weights_early: bool, + use_cudnn_cutedsl_fusion: bool, ): """Per-shard FFN forward. @@ -334,28 +442,113 @@ def _ffn_fwd_per_shard( wi_1 = wi_1.astype(sorted_x.dtype) wo = wo.astype(sorted_x.dtype) - # Concat wi_0/wi_1 along the trailing axis (NOT stack on a new - # axis). grouped_gemm requires the 3D (G, K, N) weight layout with - # contracting_dims=((1,), (1,)); a 4D stack variant walks off the - # end of the RHS and returns NaN. - wi_combined = jnp.concatenate([wi_0, wi_1], axis=-1) + # The cuDNN frontend kernel consumes alternating 32-column gate/up + # blocks. The existing TE grouped-GEMM path consumes the two full + # projections concatenated along N. + if use_cudnn_cutedsl_fusion: + from .cutedsl_extensions.moe import pack_swiglu_pair + + wi_combined = pack_swiglu_pair(wi_0, wi_1) + else: + # Concat along the trailing axis (NOT stack on a new axis). + # grouped_gemm requires the 3D (G, K, N) weight layout with + # contracting_dims=((1,), (1,)). + wi_combined = jnp.concatenate([wi_0, wi_1], axis=-1) wi_combined_bias = ( jnp.concatenate([wi_0_bias, wi_1_bias], axis=-1) if wi_0_bias is not None else None ) + # EP dispatch zero-fills aligned padding rows in recv_tokens, so + # inactive slots do not contaminate grouped quantization scales. casted_sorted_x = tex.grouped_quantize( sorted_x, fc1_quantizer_set.x, local_group_sizes, flatten_axis=-1 ) - casted_wi = tex.grouped_quantize( - wi_combined, fc1_quantizer_set.kernel, flatten_axis=-1 - ) - combined_out = tex.grouped_gemm( - casted_sorted_x.get_tensor(usage=TensorUsage.LHS), - casted_wi.get_tensor(usage=TensorUsage.RHS), - contracting_dims=((1,), (1,)), - bias=wi_combined_bias, - ) - gate_proj_out, up_proj_out = jnp.split(combined_out, 2, axis=-1) + casted_wi = tex.grouped_quantize(wi_combined, fc1_quantizer_set.kernel, flatten_axis=-1) + casted_intermediate = None + if use_cudnn_cutedsl_fusion: + from .cutedsl_extensions.moe import ( + grouped_gemm_swiglu_mxfp8, + unpack_swiglu_pair, + ) + + casted_sorted_x_lhs = casted_sorted_x.get_tensor(usage=TensorUsage.LHS) + casted_wi_rhs = casted_wi.get_tensor(usage=TensorUsage.RHS) + padded_offsets = jnp.cumsum(local_group_sizes, dtype=jnp.int32) + prob = ( + recv_w_flat[:, None, None] + if apply_topk_weights_early + else jnp.ones((sorted_x.shape[0], 1, 1), dtype=jnp.float32) + ) + ( + combined_out_3d, + intermediate_row, + intermediate_col, + intermediate_scale_row, + intermediate_scale_col, + ) = grouped_gemm_swiglu_mxfp8( + casted_sorted_x_lhs.data.reshape(sorted_x.shape[0], hidden, 1), + # TE stores the colwise payload in logical [E,K,N] order even + # though its values were quantized along K. Materialize the + # K-major [E,N,K] layout consumed by the CuTeDSL kernel. + casted_wi_rhs.data.reshape(num_local_experts, hidden, wi_combined.shape[-1]).transpose( + 0, 2, 1 + ), + casted_sorted_x_lhs.scale_inv, + casted_wi_rhs.scale_inv, + padded_offsets, + prob, + compute_dtype=sorted_x.dtype, + output_dtype=fc2_quantizer_set.x.q_dtype, + ) + combined_out = combined_out_3d.reshape(sorted_x.shape[0], wi_combined.shape[-1]) + gate_proj_out, up_proj_out = unpack_swiglu_pair(combined_out) + + intermediate_shape = (sorted_x.shape[0], gate_proj_out.shape[-1]) + scaling_mode = fc2_quantizer_set.x.scaling_mode + row_scale_size = scaling_mode.get_grouped_scale_shape( + intermediate_shape, + num_local_experts, + False, + is_padded=True, + flatten_axis=1, + )[0] + col_scale_size = scaling_mode.get_grouped_scale_shape( + intermediate_shape, + num_local_experts, + True, + is_padded=True, + flatten_axis=1, + )[0] + intermediate_scale_row = jnp.pad( + intermediate_scale_row, + (0, row_scale_size - intermediate_scale_row.size), + ) + intermediate_scale_col = jnp.pad( + intermediate_scale_col, + (0, col_scale_size - intermediate_scale_col.size), + ) + casted_intermediate = ScaledTensorFactory.create( + data=intermediate_row.reshape(-1), + scale_inv=intermediate_scale_row, + colwise_data=intermediate_col.reshape(-1), + colwise_scale_inv=intermediate_scale_col, + scaling_mode=scaling_mode, + dq_dtype=sorted_x.dtype, + data_layout=fc2_quantizer_set.x.data_layout, + q_layout=fc2_quantizer_set.x.q_layout, + flatten_axis=1, + first_dims=local_group_sizes, + original_shape=intermediate_shape, + pre_swizzled=True, + ) + else: + combined_out = tex.grouped_gemm( + casted_sorted_x.get_tensor(usage=TensorUsage.LHS), + casted_wi.get_tensor(usage=TensorUsage.RHS), + contracting_dims=((1,), (1,)), + bias=wi_combined_bias, + ) + gate_proj_out, up_proj_out = jnp.split(combined_out, 2, axis=-1) casted_sorted_x_lhs_trans = casted_sorted_x.get_tensor(usage=TensorUsage.LHS_TRANS) casted_sorted_x_lhs_trans = ( casted_sorted_x_lhs_trans.checkpoint(fc1_quantizer_set.x) @@ -375,23 +568,24 @@ def _ffn_fwd_per_shard( # that's all bf16; for FP8/FP4 the downstream grouped_quantize is what # transitions to the target precision. act_fn = _convert_to_activation_function(activation_type) - intermediate = act_fn(gate_proj_out) * up_proj_out + if not use_cudnn_cutedsl_fusion: + intermediate = act_fn(gate_proj_out) * up_proj_out - if apply_topk_weights_early: + if apply_topk_weights_early and not use_cudnn_cutedsl_fusion: # Fold the per-token combine weights into the FFN intermediate; # the downstream wo GEMM is linear so this is equivalent to the - # late-weighting path. Padded recv slots can contain uninitialized - # data, so overwrite inactive rows with literal zeros instead of - # relying on multiplication by a zero mask (IEEE NaN * 0 = NaN). + # late-weighting path. EP dispatch zero-fills aligned padding tokens + # and routing weights; grouped ops exclude any trailing over-allocation, + # so a separate validity select is unnecessary. # ``w_b`` is cast to ``intermediate.dtype`` so the multiply doesn't # promote expert_outputs above the EP buffer's element width. - w_b = recv_w_flat[:, None].astype(intermediate.dtype) - active = (recv_w_flat != 0)[:, None] - intermediate = jnp.where(active, intermediate * w_b, jnp.zeros_like(intermediate)) + intermediate = intermediate * recv_w_flat[:, None].astype(intermediate.dtype) - casted_intermediate = tex.grouped_quantize( - intermediate, fc2_quantizer_set.x, local_group_sizes, flatten_axis=-1 - ) + if not use_cudnn_cutedsl_fusion: + casted_intermediate = tex.grouped_quantize( + intermediate, fc2_quantizer_set.x, local_group_sizes, flatten_axis=-1 + ) + assert casted_intermediate is not None casted_wo = tex.grouped_quantize(wo, fc2_quantizer_set.kernel, flatten_axis=-1) expert_outputs = tex.grouped_gemm( casted_intermediate.get_tensor(usage=TensorUsage.LHS), @@ -445,6 +639,7 @@ def _ffn_bwd_per_shard( activation_type: str, apply_topk_weights_early: bool, has_bias: bool, + use_cudnn_cutedsl_fusion: bool, ): """Per-shard FFN backward. @@ -491,9 +686,9 @@ def _ffn_bwd_per_shard( is_colwise=False, flatten_axis=2, ) - # cuBLAS grouped_gemm skips size_g == 0 groups without zero-filling - # the output slice; mask 0-token-expert wgrads to zero so the - # optimizer never sees uninit memory. + # Grouped GEMM skips size_g == 0 experts without writing their wgrad + # slice. Mask those slices to zero so the optimizer never sees garbage; + # this is unrelated to EP alignment padding (handled by dispatch zero-fill). wgrad_group_active = (local_group_sizes > 0)[:, None, None] # wo bwd @@ -502,11 +697,6 @@ def _ffn_bwd_per_shard( ) _casted_d_eo_lhs = casted_d_eo.get_tensor(usage=TensorUsage.LHS) _casted_d_eo_rhs = casted_d_eo.get_tensor(usage=TensorUsage.RHS) - d_intermediate = tex.grouped_gemm( - _casted_d_eo_lhs, - casted_wo_rhs_trans, - contracting_dims=((1,), (2,)), - ) d_wo = tex.grouped_gemm( casted_intermediate_lhs_trans, _casted_d_eo_rhs, @@ -515,44 +705,128 @@ def _ffn_bwd_per_shard( d_wo = jnp.where(wgrad_group_active, d_wo, jnp.zeros_like(d_wo)) d_wo_bias = tex.grouped_dbias(d_eo_2d, local_group_sizes) if has_bias else None - act_fn = _convert_to_activation_function(activation_type) - if apply_topk_weights_early: - # intermediate' = intermediate * w * mask. Split the cotangent - # across both factors before the activation bwd consumes it. Padded - # recv slots may still be NaN in the saved activation residuals, so - # use zero-filled residuals on inactive rows before the activation VJP. - w_b = recv_w_flat[:, None].astype(d_intermediate.dtype) - active = (recv_w_flat != 0)[:, None] - gate_proj_for_bwd = jnp.where(active, gate_proj_out, jnp.zeros_like(gate_proj_out)) - up_proj_for_bwd = jnp.where(active, up_proj_out, jnp.zeros_like(up_proj_out)) - intermediate_unweighted = act_fn(gate_proj_for_bwd) * up_proj_for_bwd - d_recv_w_from_intermediate = jnp.sum( - d_intermediate * intermediate_unweighted, - axis=-1, - ).astype(recv_w_flat.dtype) - d_intermediate = jnp.where(active, d_intermediate * w_b, jnp.zeros_like(d_intermediate)) + if use_cudnn_cutedsl_fusion: + from .cutedsl_extensions.moe import ( + grouped_gemm_dswiglu_mxfp8, + pack_swiglu_pair, + ) + + packed_forward = pack_swiglu_pair(gate_proj_out, up_proj_out) + padded_offsets = jnp.cumsum(local_group_sizes, dtype=jnp.int32) + prob = ( + recv_w_flat[:, None, None] + if apply_topk_weights_early + else jnp.ones((recv_rows, 1, 1), dtype=jnp.float32) + ) + ( + d_combined_row, + d_combined_col, + d_combined_scale_row, + d_combined_scale_col, + _dprob, + ) = grouped_gemm_dswiglu_mxfp8( + _casted_d_eo_lhs.data.reshape(recv_rows, hidden, 1), + casted_wo_rhs_trans.data.reshape( + num_local_experts, intermediate_size, hidden + ).transpose(1, 2, 0), + packed_forward.reshape(recv_rows, 2 * intermediate_size, 1), + _casted_d_eo_lhs.scale_inv, + casted_wo_rhs_trans.scale_inv, + padded_offsets, + prob, + output_dtype=fc1_quantizer_set.dgrad.q_dtype, + ) + d_combined_shape = (recv_rows, 2 * intermediate_size) + scaling_mode = fc1_quantizer_set.dgrad.scaling_mode + row_scale_size = scaling_mode.get_grouped_scale_shape( + d_combined_shape, + num_local_experts, + False, + is_padded=True, + flatten_axis=1, + )[0] + col_scale_size = scaling_mode.get_grouped_scale_shape( + d_combined_shape, + num_local_experts, + True, + is_padded=True, + flatten_axis=1, + )[0] + d_combined_scale_row = jnp.pad( + d_combined_scale_row, + (0, row_scale_size - d_combined_scale_row.size), + ) + d_combined_scale_col = jnp.pad( + d_combined_scale_col, + (0, col_scale_size - d_combined_scale_col.size), + ) + casted_d_combined = ScaledTensorFactory.create( + data=d_combined_row.reshape(-1), + scale_inv=d_combined_scale_row, + colwise_data=d_combined_col.reshape(-1), + colwise_scale_inv=d_combined_scale_col, + scaling_mode=scaling_mode, + dq_dtype=gate_proj_out.dtype, + data_layout=fc1_quantizer_set.dgrad.data_layout, + q_layout=fc1_quantizer_set.dgrad.q_layout, + flatten_axis=1, + first_dims=local_group_sizes, + original_shape=d_combined_shape, + pre_swizzled=True, + ) + if apply_topk_weights_early: + # The cuDNN frontend dSwiGLU kernel accumulates dprob with + # atomics and expects a zero-initialized output buffer. cutlass_call + # allocates custom-call outputs uninitialized, so compute this + # routing-weight cotangent explicitly until we can pass dprob as an + # initialized input/output buffer. + d_intermediate_for_prob = tex.grouped_gemm( + _casted_d_eo_lhs, + casted_wo_rhs_trans, + contracting_dims=((1,), (2,)), + ) + act_fn = _convert_to_activation_function(activation_type) + intermediate_unweighted = act_fn(gate_proj_out) * up_proj_out + d_recv_w_from_intermediate = jnp.sum( + d_intermediate_for_prob * intermediate_unweighted, + axis=-1, + ) + else: + d_recv_w_from_intermediate = jnp.zeros_like(recv_w_flat) + d_recv_w_from_intermediate = d_recv_w_from_intermediate.astype(recv_w_flat.dtype) + d_combined_for_bias = None else: - gate_proj_for_bwd = gate_proj_out - up_proj_for_bwd = up_proj_out - d_recv_w_from_intermediate = jnp.zeros_like(recv_w_flat) - - # Activation bwd, symmetric with the fwd: silu' and the two - # elementwise products run in the GEMM dtype (no fp32 island), so - # the chain rule composes through at the same precision the wi/wo - # GEMMs consume. - act_gp, dact_pullback = jax.vjp(act_fn, gate_proj_for_bwd) - d_up_proj_out = d_intermediate * act_gp - (d_gate_proj_out,) = dact_pullback(d_intermediate * up_proj_for_bwd) - - # wi bwd (fused gate/up via concat). Mirror the fused fwd: pack the - # gate/up cotangents along the trailing axis, run a single - # grouped_quantize + two grouped_gemm pair (one dgrad, one wgrad) - # against the fused casted_wi_rhs_trans residual, then split the - # wgrad result back into d_wi_0 / d_wi_1 halves with jnp.split. - d_combined = jnp.concatenate([d_gate_proj_out, d_up_proj_out], axis=-1) - casted_d_combined = tex.grouped_quantize( - d_combined, fc1_quantizer_set.dgrad, local_group_sizes, flatten_axis=-1 - ) + d_intermediate = tex.grouped_gemm( + _casted_d_eo_lhs, + casted_wo_rhs_trans, + contracting_dims=((1,), (2,)), + ) + act_fn = _convert_to_activation_function(activation_type) + if apply_topk_weights_early: + # intermediate' = intermediate * w. Split the cotangent across + # both factors before the activation bwd consumes it. EP zero-fills + # aligned padding, and grouped ops exclude trailing over-allocation. + w_b = recv_w_flat[:, None].astype(d_intermediate.dtype) + intermediate_unweighted = act_fn(gate_proj_out) * up_proj_out + d_recv_w_from_intermediate = jnp.sum( + d_intermediate * intermediate_unweighted, + axis=-1, + ).astype(recv_w_flat.dtype) + d_intermediate = d_intermediate * w_b + else: + d_recv_w_from_intermediate = jnp.zeros_like(recv_w_flat) + + # Activation bwd, symmetric with the fwd: silu' and the two + # elementwise products run in the GEMM dtype (no fp32 island), so + # the chain rule composes through at the same precision the wi/wo + # GEMMs consume. + act_gp, dact_pullback = jax.vjp(act_fn, gate_proj_out) + d_up_proj_out = d_intermediate * act_gp + (d_gate_proj_out,) = dact_pullback(d_intermediate * up_proj_out) + d_combined_for_bias = jnp.concatenate([d_gate_proj_out, d_up_proj_out], axis=-1) + casted_d_combined = tex.grouped_quantize( + d_combined_for_bias, fc1_quantizer_set.dgrad, local_group_sizes, flatten_axis=-1 + ) d_sorted_x = tex.grouped_gemm( casted_d_combined.get_tensor(usage=TensorUsage.LHS), casted_wi_rhs_trans, @@ -564,9 +838,14 @@ def _ffn_bwd_per_shard( contracting_dims=((0,), (0,)), ) d_wi_combined = jnp.where(wgrad_group_active, d_wi_combined, jnp.zeros_like(d_wi_combined)) - d_wi_0, d_wi_1 = jnp.split(d_wi_combined, 2, axis=-1) + if use_cudnn_cutedsl_fusion: + from .cutedsl_extensions.moe import unpack_swiglu_pair + + d_wi_0, d_wi_1 = unpack_swiglu_pair(d_wi_combined) + else: + d_wi_0, d_wi_1 = jnp.split(d_wi_combined, 2, axis=-1) if has_bias: - d_wi_combined_bias = tex.grouped_dbias(d_combined, local_group_sizes) + d_wi_combined_bias = tex.grouped_dbias(d_combined_for_bias, local_group_sizes) d_wi_0_bias, d_wi_1_bias = jnp.split(d_wi_combined_bias, 2, axis=-1) else: d_wi_0_bias = None @@ -620,6 +899,7 @@ def _moe_fwd_rule( wo_kernel_axes, dtype, apply_topk_weights_early, + use_cudnn_cutedsl_fusion, ): """Forward: gate -> topk -> ep_dispatch -> shard_map(FFN) -> ep_combine. @@ -659,10 +939,15 @@ def _moe_fwd_rule( tokens_per_ep_group = num_ep * max_tokens_per_rank max_local_assignments = tokens_per_ep_group * min(K, num_local_experts) max_nonempty_experts = min(num_local_experts, max_local_assignments) - padded_total_bound = max_local_assignments + (_ALIGN_SIZE - 1) * max_nonempty_experts - aligned_total_bound = ((padded_total_bound + _ALIGN_SIZE - 1) // _ALIGN_SIZE) * _ALIGN_SIZE + dispatch_alignment = _CUDNN_CUTEDSL_ALIGN_SIZE if use_cudnn_cutedsl_fusion else _ALIGN_SIZE + padded_total_bound = max_local_assignments + (dispatch_alignment - 1) * max_nonempty_experts + aligned_total_bound = ( + (padded_total_bound + dispatch_alignment - 1) // dispatch_alignment + ) * dispatch_alignment per_expert_bound = ( - num_local_experts * ((tokens_per_ep_group + _ALIGN_SIZE - 1) // _ALIGN_SIZE) * _ALIGN_SIZE + num_local_experts + * ((tokens_per_ep_group + dispatch_alignment - 1) // dispatch_alignment) + * dispatch_alignment ) recv_pr = min(per_expert_bound, aligned_total_bound) @@ -780,7 +1065,7 @@ def _moe_fwd_rule( # ---------------- TE EP dispatch (global view) ---------------- cfg = tex.EpLayerConfig( top_k=K, - dispatch_output_per_expert_alignment=_ALIGN_SIZE, + dispatch_output_per_expert_alignment=dispatch_alignment, ) token_counts, handle_mem = tex.ep_prepare(cfg, topk_idx_3d) recv_tokens, recv_topk_weights = tex.ep_dispatch_fwd( @@ -835,21 +1120,11 @@ def _body(*args): else: (r_tok, r_w, tc, w0, w1, w_o, fc1_qset, fc2_qset) = args w0b = w1b = wob = None - # NOTE: tex.ep_dispatch_fwd's NCCL EP HT path leaves the recv - # buffer uninitialised on fully-empty-receiver ranks (and at - # padded slots on partially-loaded ranks). We don't need a - # zero-init guard here anymore because: - # 1. ``tc`` (per-expert padded counts) is plumbed into - # grouped_gemm as group_sizes, so cuBLAS skips both - # 0-token experts and the trailing overalloc tail. - # 2. The per-group wgrad masks in _ffn_bwd_per_shard zero - # ``d_wo`` / ``d_wi_combined`` slices for 0-token-globally - # experts (cuBLAS skips size_g==0 groups without - # zero-filling, which would otherwise leak NaN into the - # user's optimizer). - # 3. All other downstream consumers (ep_combine, - # ep_dispatch_bwd) are handle_mem-aware and read only - # valid positions. + # NCCL EP zero-fills aligned expert padding. Any trailing + # over-allocation remains safe because ``tc`` is plumbed into + # grouped operations and EP consumers use handle metadata to read + # only valid positions. Zero-token-expert wgrads are masked in + # ``_ffn_bwd_per_shard`` separately from padding. # If a future caller adds a non-group-aware reader of r_tok # (e.g. an inspect probe over the full recv tile), re-add the # ``jax.lax.cond(jnp.any(r_w != 0), identity, zeros_like)`` @@ -869,6 +1144,7 @@ def _body(*args): num_local_experts=num_local_experts, activation_type=activation_type, apply_topk_weights_early=apply_topk_weights_early, + use_cudnn_cutedsl_fusion=use_cudnn_cutedsl_fusion, ) expert_outputs, ffn_residuals = shard_map( @@ -965,6 +1241,7 @@ def _moe_bwd_rule( wo_kernel_axes, dtype, apply_topk_weights_early, + use_cudnn_cutedsl_fusion, residuals, cotangents, ): @@ -1009,14 +1286,10 @@ def _moe_bwd_rule( d_expert_outputs = grad_pre_combine d_recv_w_from_combine = jnp.zeros_like(ctx.recv_topk_weights) else: - # Reverse the late-weighting multiply. Padded expert-major rows are - # part of the physical grouped-GEMM ranges, so write literal zero - # cotangents for inactive rows instead of relying on NaN * 0. + # Reverse the late-weighting multiply. NCCL EP zero-fills aligned + # padding, while grouped/EP consumers ignore trailing over-allocation. w = ctx.recv_topk_weights[..., None].astype(grad_pre_combine.dtype) - mask_bool = (ctx.recv_topk_weights != 0)[..., None] - d_expert_outputs = jnp.where( - mask_bool, grad_pre_combine * w, jnp.zeros_like(grad_pre_combine) - ) + d_expert_outputs = grad_pre_combine * w d_recv_w_from_combine = (grad_pre_combine * ctx.expert_outputs).sum(axis=-1) d_recv_w_from_combine = d_recv_w_from_combine.astype(ctx.recv_topk_weights.dtype) @@ -1076,6 +1349,7 @@ def _bwd_body(*args): activation_type=activation_type, apply_topk_weights_early=apply_topk_weights_early, has_bias=has_bias, + use_cudnn_cutedsl_fusion=use_cudnn_cutedsl_fusion, ) # Weight grads accumulate per-DP-shard inside the body; psum across # DP axes so each replica sees the full sum (matches out_specs @@ -1224,7 +1498,7 @@ def _bwd_body(*args): # ============================================================================= -@partial(jax.custom_vjp, nondiff_argnums=tuple(range(11, 28))) +@partial(jax.custom_vjp, nondiff_argnums=tuple(range(11, 29))) def _moe( x, gate_kernel, @@ -1254,6 +1528,7 @@ def _moe( wo_kernel_axes, dtype, apply_topk_weights_early, + use_cudnn_cutedsl_fusion, ): primal, _ = _moe_fwd_rule( x, @@ -1284,6 +1559,7 @@ def _moe( wo_kernel_axes, dtype, apply_topk_weights_early, + use_cudnn_cutedsl_fusion, ) return primal @@ -1342,9 +1618,11 @@ def moe( ``fused_moe_aux_loss`` kernel sees a global ``[T_global, E]`` view; this lives off the dispatch critical path. - Note that the per-expert dispatch-slot alignment is fixed internally - at 128 tokens (``_ALIGN_SIZE``); see that constant's docstring for - rationale and how to extend if a future recipe needs >128. + The default per-expert dispatch-slot alignment is 128 tokens. Setting + ``NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION=1`` requires the call to match + the eligible MXFP8 SwiGLU surface, uses the cuDNN frontend CuTeDSL FC1 + fusion, and raises dispatch alignment to the kernel-required 256 tokens. + Ineligible opt-in calls fail with the reasons they cannot use CuTeDSL. Axis-name parameters: @@ -1375,6 +1653,26 @@ def moe( surrounding design rationale. """ score_function = _validate_score_function(score_function) + use_cudnn_cutedsl_fusion = False + if _use_cudnn_cutedsl_fusion_from_env(): + rejection_reasons = _cudnn_cutedsl_fusion_rejection_reasons( + x, + wi_0, + wi_1, + wi_0_bias, + wi_1_bias, + fc1_quantizer_set, + fc2_quantizer_set, + num_experts=num_experts, + activation_type=activation_type, + ep_axis=ep_axis, + ) + if rejection_reasons: + raise ValueError( + f"{_CUDNN_CUTEDSL_ENV}=1 is unsupported for this moe() call:\n- " + + "\n- ".join(rejection_reasons) + ) + use_cudnn_cutedsl_fusion = True # Enforce ((outer_dp..., ep), None, None) on inbound activations. The # EP comm groups consecutive global ranks (dp_color = rank // ep_size), @@ -1433,6 +1731,7 @@ def moe( wo_kernel_axes, dtype, apply_topk_weights_early, + use_cudnn_cutedsl_fusion, ) if aux_loss_coeff <= 0.0: aux_loss = None From 8a9f704bd80bfd46dd3cf25002e8ec2ca3d431b6 Mon Sep 17 00:00:00 2001 From: tdophung Date: Thu, 23 Jul 2026 14:38:25 -0700 Subject: [PATCH 03/38] JAX: add TVM-FFI grouped GEMM SwiGLU bridge Add the framework-neutral cuDNN-FE compiler patch and JAX TVM-FFI adapter, fix the compact multi-expert weight ABI, and cover descriptor-only compilation, eight-expert ragged parity, multiprocess partitioning, and integrated training updates. --- docs/examples/jax/cutedsl_moe_pipeclean.rst | 93 +- ...work-neutral-grouped-SwiGLU-compiler.patch | 813 ++++++++++++++++++ tests/jax/requirements_cutedsl.txt | 6 +- tests/jax/run_tvm_ffi_fused_moe_e2e.sh | 92 ++ .../run_tvm_ffi_grouped_mlp_multiprocess.sh | 62 ++ tests/jax/test_cutedsl_moe.py | 123 ++- tests/jax/test_te_ep_moe.py | 47 +- tests/jax/tvm_ffi_grouped_mlp_multiprocess.py | 359 ++++++++ .../jax/cpp_extensions/__init__.py | 1 + .../jax/cpp_extensions/grouped_gemm_swiglu.py | 350 ++++++++ transformer_engine/jax/moe.py | 38 +- 11 files changed, 1879 insertions(+), 105 deletions(-) create mode 100644 tests/jax/patches/0001-Add-framework-neutral-grouped-SwiGLU-compiler.patch create mode 100755 tests/jax/run_tvm_ffi_fused_moe_e2e.sh create mode 100755 tests/jax/run_tvm_ffi_grouped_mlp_multiprocess.sh create mode 100644 tests/jax/tvm_ffi_grouped_mlp_multiprocess.py create mode 100644 transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py diff --git a/docs/examples/jax/cutedsl_moe_pipeclean.rst b/docs/examples/jax/cutedsl_moe_pipeclean.rst index f725eb45bc1..8df5ddec841 100644 --- a/docs/examples/jax/cutedsl_moe_pipeclean.rst +++ b/docs/examples/jax/cutedsl_moe_pipeclean.rst @@ -2,53 +2,61 @@ CuTeDSL MXFP8 MoE pipeclean =========================== The JAX MoE path has an opt-in Blackwell fusion for FC1 grouped GEMM, -SwiGLU, and rowwise/colwise MXFP8 quantization. All CuTe-specific code lives -in ``transformer_engine/jax/cutedsl_extensions``; the surrounding EP path -continues to use ``shard_map`` and sees only shard-local tensors. +SwiGLU, and rowwise/colwise MXFP8 quantization. The forward leaf is compiled +from abstract tensor descriptors by cuDNN-FE and invoked as a native +``tvm_ffi.Function``. The surrounding EP path continues to use ``shard_map`` +and sees only shard-local tensors. The fused dSwiGLU backward leaf continues +to use CUTLASS DSL's JAX integration. Environment baseline -------------------- Use the versions in ``tests/jax/requirements_cutedsl.txt`` on an SM100 CUDA -host. Build Transformer Engine with JAX and NCCL EP support, then run:: - - python3 tests/jax/cutedsl_smoke.py - python3 -m pytest -c tests/jax/pytest.ini tests/jax/test_cutedsl_moe.py -v - bash tests/jax/run_te_ep_moe.sh - NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION=1 \ - bash tests/jax/run_te_ep_moe.sh -k TestTeEpMoeCudnnCutedslFusion - -The first command is intentionally independent of Transformer Engine. It -must report ``cutlass.jax.is_available()`` and execute its vector-add kernel -inside ``jax.jit`` before failures in the fused MoE tests are investigated. -The EP launcher currently requires four or more ranks even though the EP -mesh itself uses groups of two ranks. The CuTeDSL opt-in is strict, so run -only the CuTeDSL fusion class with that environment variable enabled. +host and install the sibling cuDNN-FE checkout with its compile-only extra. +The end-to-end driver installs both editable trees, runs standalone forward +and backward parity, exercises custom partitioning and ``shard_map`` across +two processes, and runs the four-GPU MoEBlock integration test:: + + CUDNN_FE_ROOT=/path/to/cudnn-frontend NUM_GPUS=4 \ + bash tests/jax/run_tvm_ffi_fused_moe_e2e.sh + +If ``CUDNN_FE_ROOT`` does not exist, the driver clones the pinned NVIDIA +upstream revision and applies the bundled compile-only patch. An existing +checkout that already provides the API is left unchanged. Set +``INSTALL_DEPS=0`` to reuse an existing environment. The EP launcher currently +requires four or more ranks even though the EP mesh itself uses groups of two +ranks. Logs from every stage are retained in the timestamped artifact directory +printed by the driver. Support and fallback contract ----------------------------- The default value of ``NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION`` is ``0`` and -never imports CUTLASS DSL. Existing unfused runs therefore remain supported -when the optional packages or SM100 hardware are absent. Explicit opt-in -validates the GPU architecture, SwiGLU shapes and biases, MXFP8 quantizers, -JAX/CUTLASS FFI availability, and the cuDNN frontend source layout before -kernel lowering. An unsupported explicit opt-in fails with the complete validation -list instead of failing later during kernel compilation. - -The pinned cuDNN frontend layout is:: - - cudnn/grouped_gemm/utils.py - cudnn/grouped_gemm/grouped_gemm_swiglu/grouped_gemm_swiglu_quant.py - -and must export ``BlockScaledContiguousGroupedGemmKernel``. The fused path +does not import the optional TVM-FFI or CUTLASS DSL packages. Existing +unfused runs therefore remain supported when the optional packages or SM100 +hardware are absent. Explicit opt-in validates the GPU architecture, SwiGLU +shapes and biases, MXFP8 quantizers, JAX FFI availability, and both compiler +paths before lowering. Unsupported configurations emit one warning with the +complete validation list and use the unfused implementation. + +The cuDNN-FE compile-only API is exported from:: + + cudnn/grouped_gemm/grouped_gemm_swiglu/compile.py + +It accepts self-describing operand metadata, compiles without Torch or live +buffers, and returns a native function plus an exact ABI descriptor. The +native launch wrapper reorders the eight arguments, five results, and stream +into the volatile raw kernel ABI. This is necessary because ``jax-tvm-ffi`` +currently describes only complete argument and result groups. The fused path uses 256-token dispatch alignment; the unfused path remains at 128. Custom partitioning migration design ------------------------------------ -Do not add a ``custom_partitioning`` wrapper until nvbug 6432162 is fixed. -Today the five kernel outputs are: +The standalone single-expert primitive now has a ``custom_partitioning`` +wrapper and is tested against an explicit ``shard_map`` path. Integrated +multi-expert MoE remains inside the established ``shard_map`` boundary. +The five kernel outputs are: * combined projection: ``[tokens, 2 * intermediate, 1]`` * rowwise MXFP8 payload: ``[tokens, intermediate, 1]`` @@ -68,18 +76,9 @@ singleton factors for each singleton dimension, and a constant factor * pre-swizzled A/B scales: opaque until their JAX boundary is structured The first three outputs inherit the ``tokens`` factor and otherwise remain -local. The two scale outputs must eventually be exposed as structured block -layouts carrying ``tokens``/``intermediate`` (and the per-expert padding -factor where applicable), then flattened only inside ``cutlass_call``. Do -not assign a Shardy factor to the current giant flat dimension: that is the -failure mode tracked by nvbug 6432162. Until structured scales and the bug -fix are both available, the leaf stays under ``shard_map`` and needs no -partitioning rule. - -``tests/jax/repro_cutedsl_moe_shardy.py`` lowers a shape-only custom -partitioning leaf with the production EP2/FSDP2 output shapes. Run it before -starting the migration; it should lower successfully after the Shardy fix. -After that gate passes, replace the shape-only leaf with -``grouped_gemm_swiglu_mxfp8``, structure both scale outputs, add the rule -above, and compare its shardings and numerics with the existing ``shard_map`` -tests before removing ``shard_map``. +local. Scale buffers remain opaque flat native-ABI values at lowering time. +The standalone partitioner can shard them along the data axis because every +shard compiles a complete local scale buffer. Multi-expert production +partitioning stays below ``shard_map`` so expert-local padded offsets and +pre-swizzled scale layouts are preserved without exposing an invalid global +flat-buffer factor to Shardy. diff --git a/tests/jax/patches/0001-Add-framework-neutral-grouped-SwiGLU-compiler.patch b/tests/jax/patches/0001-Add-framework-neutral-grouped-SwiGLU-compiler.patch new file mode 100644 index 00000000000..d26ca1e4854 --- /dev/null +++ b/tests/jax/patches/0001-Add-framework-neutral-grouped-SwiGLU-compiler.patch @@ -0,0 +1,813 @@ +From fe78f4c3f3050b537362ebca83b734adfbff1093 Mon Sep 17 00:00:00 2001 +From: tdophung +Date: Wed, 22 Jul 2026 12:15:31 -0700 +Subject: [PATCH] Add framework-neutral grouped SwiGLU compiler + +--- + .../gemm_fusions/grouped_gemm_swiglu.md | 41 ++ + pyproject.toml | 5 + + python/cudnn/__init__.py | 6 + + python/cudnn/grouped_gemm/__init__.py | 114 +++--- + .../grouped_gemm_swiglu/__init__.py | 32 +- + .../grouped_gemm_swiglu/compile.py | 366 ++++++++++++++++++ + python/cudnn/grouped_gemm/utils.py | 10 +- + .../test_grouped_gemm_swiglu_compile.py | 96 +++++ + 8 files changed, 597 insertions(+), 73 deletions(-) + create mode 100644 python/cudnn/grouped_gemm/grouped_gemm_swiglu/compile.py + create mode 100644 test/python/fe_api/grouped_gemm/test_grouped_gemm_swiglu_compile.py + +diff --git a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_swiglu.md b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_swiglu.md +index 1ac83b7..bf73160 100644 +--- a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_swiglu.md ++++ b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_swiglu.md +@@ -96,6 +96,47 @@ $$ + + ## API Usage + ++### Framework-neutral compile-only API ++ ++The compile-only entrypoint is intended for frameworks that own their output ++allocation and launch through TVM-FFI. It consumes tensor metadata, compiles a ++native `tvm_ffi.Function`, and returns a `GroupedGemmAbiDescriptor` describing ++the exact positional ABI of that function. It does not allocate framework ++buffers and does not launch the kernel. ++ ++Install the compiler path without PyTorch: ++ ++```bash ++pip install -e ".[cutedsl-compile]" ++``` ++ ++```python ++from cudnn import ( ++ GroupedGemmElementType as DType, ++ GroupedGemmOperandDesc as Operand, ++ GroupedGemmSwigluConfig, ++ compile_grouped_gemm_swiglu, ++) ++ ++operands = { ++ "a": Operand("a", DType.FLOAT8_E4M3FN, (256, 128, 1), (1, 0, 2)), ++ "b": Operand("b", DType.FLOAT8_E4M3FN, (256, 128, 1), (1, 0, 2)), ++ # sfa, sfb, padded_offsets, alpha, prob, norm_const and the five ++ # output descriptors are omitted here for brevity. ++} ++native_function, abi = compile_grouped_gemm_swiglu( ++ operands=operands, ++ config=GroupedGemmSwigluConfig(), ++) ++``` ++ ++The supported native ABI groups the eight input buffers, the five output ++buffers, and the CUDA stream. This grouping can be registered directly with ++`jax-tvm-ffi` as `abi.arg_spec == ("args", "rets", "ctx.stream")`. A native ++CuTeDSL reorder launcher hides the volatile raw-kernel operand order, so no ++Python callback runs during execution. The `amax` slot is disabled in this ++five-output prototype. ++ + ### High-level Wrapper + + ```python +diff --git a/pyproject.toml b/pyproject.toml +index 544b28d..c68805d 100644 +--- a/pyproject.toml ++++ b/pyproject.toml +@@ -61,6 +61,11 @@ cutedsl = [ + "apache-tvm-ffi", + "torch-c-dlpack-ext", + ] ++cutedsl-compile = [ ++ "nvidia-cutlass-dsl[cu13]>=4.5.0", ++ "cuda-python", ++ "apache-tvm-ffi>=0.1.12", ++] + + [dependency-groups] + dev = [ +diff --git a/python/cudnn/__init__.py b/python/cudnn/__init__.py +index 6e8b998..e4493a4 100644 +--- a/python/cudnn/__init__.py ++++ b/python/cudnn/__init__.py +@@ -298,6 +298,12 @@ _LAZY_OPTIONAL_IMPORTS = { + "grouped_gemm": (".grouped_gemm", None), + "GroupedGemmSwigluSm100": (".grouped_gemm", "GroupedGemmSwigluSm100"), + "grouped_gemm_swiglu_wrapper_sm100": (".grouped_gemm", "grouped_gemm_swiglu_wrapper_sm100"), ++ "GroupedGemmElementType": (".grouped_gemm", "GroupedGemmElementType"), ++ "GroupedGemmOperandDesc": (".grouped_gemm", "GroupedGemmOperandDesc"), ++ "GroupedGemmSwigluConfig": (".grouped_gemm", "GroupedGemmSwigluConfig"), ++ "GroupedGemmAbiEntry": (".grouped_gemm", "GroupedGemmAbiEntry"), ++ "GroupedGemmAbiDescriptor": (".grouped_gemm", "GroupedGemmAbiDescriptor"), ++ "compile_grouped_gemm_swiglu": (".grouped_gemm", "compile_grouped_gemm_swiglu"), + "GroupedGemmDswigluSm100": (".grouped_gemm", "GroupedGemmDswigluSm100"), + "grouped_gemm_dswiglu_wrapper_sm100": (".grouped_gemm", "grouped_gemm_dswiglu_wrapper_sm100"), + "GroupedGemmSreluSm100": (".grouped_gemm", "GroupedGemmSreluSm100"), +diff --git a/python/cudnn/grouped_gemm/__init__.py b/python/cudnn/grouped_gemm/__init__.py +index 818a1ab..0f69996 100644 +--- a/python/cudnn/grouped_gemm/__init__.py ++++ b/python/cudnn/grouped_gemm/__init__.py +@@ -1,68 +1,50 @@ +-# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. ++# Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + # SPDX-License-Identifier: MIT + +-from .grouped_gemm_swiglu.api import ( +- GroupedGemmSwigluSm100, +- grouped_gemm_swiglu_wrapper_sm100, +-) +- +-from .grouped_gemm_dswiglu.api import ( +- GroupedGemmDswigluSm100, +- grouped_gemm_dswiglu_wrapper_sm100, +-) +- +-from .grouped_gemm_quant.api import ( +- GroupedGemmQuantSm100, +- grouped_gemm_quant_wrapper_sm100, +-) +- +-from .grouped_gemm_srelu.api import ( +- GroupedGemmSreluSm100, +- grouped_gemm_srelu_wrapper_sm100, +-) +- +-from .grouped_gemm_dsrelu.api import ( +- GroupedGemmDsreluSm100, +- grouped_gemm_dsrelu_wrapper_sm100, +-) +- +-from .grouped_gemm_glu.api import ( +- GroupedGemmGluSm100, +- grouped_gemm_glu_wrapper_sm100, +-) +- +-from .grouped_gemm_glu_hadamard.api import ( +- GroupedGemmGluHadamardSm100, +- grouped_gemm_glu_hadamard_wrapper_sm100, +-) +- +-from .grouped_gemm_dglu.api import ( +- GroupedGemmDgluSm100, +- grouped_gemm_dglu_wrapper_sm100, +-) +- +-from .grouped_gemm_wgrad.api import ( +- GroupedGemmWgradSm100, +- grouped_gemm_wgrad_wrapper_sm100, +-) +- +-__all__ = [ +- "GroupedGemmSwigluSm100", +- "grouped_gemm_swiglu_wrapper_sm100", +- "GroupedGemmDswigluSm100", +- "grouped_gemm_dswiglu_wrapper_sm100", +- "GroupedGemmQuantSm100", +- "grouped_gemm_quant_wrapper_sm100", +- "GroupedGemmSreluSm100", +- "grouped_gemm_srelu_wrapper_sm100", +- "GroupedGemmDsreluSm100", +- "grouped_gemm_dsrelu_wrapper_sm100", +- "GroupedGemmGluSm100", +- "grouped_gemm_glu_wrapper_sm100", +- "GroupedGemmGluHadamardSm100", +- "grouped_gemm_glu_hadamard_wrapper_sm100", +- "GroupedGemmDgluSm100", +- "grouped_gemm_dglu_wrapper_sm100", +- "GroupedGemmWgradSm100", +- "grouped_gemm_wgrad_wrapper_sm100", +-] ++"""Lazy public exports for grouped GEMM CuTeDSL APIs. ++ ++Keeping these imports lazy lets framework-neutral compile-only entrypoints be ++used in environments that intentionally do not install PyTorch. ++""" ++ ++from importlib import import_module ++ ++ ++_LAZY_IMPORTS = { ++ "GroupedGemmSwigluSm100": (".grouped_gemm_swiglu", "GroupedGemmSwigluSm100"), ++ "grouped_gemm_swiglu_wrapper_sm100": (".grouped_gemm_swiglu", "grouped_gemm_swiglu_wrapper_sm100"), ++ "GroupedGemmElementType": (".grouped_gemm_swiglu", "GroupedGemmElementType"), ++ "GroupedGemmOperandDesc": (".grouped_gemm_swiglu", "GroupedGemmOperandDesc"), ++ "GroupedGemmSwigluConfig": (".grouped_gemm_swiglu", "GroupedGemmSwigluConfig"), ++ "GroupedGemmAbiEntry": (".grouped_gemm_swiglu", "GroupedGemmAbiEntry"), ++ "GroupedGemmAbiDescriptor": (".grouped_gemm_swiglu", "GroupedGemmAbiDescriptor"), ++ "compile_grouped_gemm_swiglu": (".grouped_gemm_swiglu", "compile_grouped_gemm_swiglu"), ++ "GroupedGemmDswigluSm100": (".grouped_gemm_dswiglu.api", "GroupedGemmDswigluSm100"), ++ "grouped_gemm_dswiglu_wrapper_sm100": (".grouped_gemm_dswiglu.api", "grouped_gemm_dswiglu_wrapper_sm100"), ++ "GroupedGemmQuantSm100": (".grouped_gemm_quant.api", "GroupedGemmQuantSm100"), ++ "grouped_gemm_quant_wrapper_sm100": (".grouped_gemm_quant.api", "grouped_gemm_quant_wrapper_sm100"), ++ "GroupedGemmSreluSm100": (".grouped_gemm_srelu.api", "GroupedGemmSreluSm100"), ++ "grouped_gemm_srelu_wrapper_sm100": (".grouped_gemm_srelu.api", "grouped_gemm_srelu_wrapper_sm100"), ++ "GroupedGemmDsreluSm100": (".grouped_gemm_dsrelu.api", "GroupedGemmDsreluSm100"), ++ "grouped_gemm_dsrelu_wrapper_sm100": (".grouped_gemm_dsrelu.api", "grouped_gemm_dsrelu_wrapper_sm100"), ++ "GroupedGemmGluSm100": (".grouped_gemm_glu.api", "GroupedGemmGluSm100"), ++ "grouped_gemm_glu_wrapper_sm100": (".grouped_gemm_glu.api", "grouped_gemm_glu_wrapper_sm100"), ++ "GroupedGemmGluHadamardSm100": (".grouped_gemm_glu_hadamard.api", "GroupedGemmGluHadamardSm100"), ++ "grouped_gemm_glu_hadamard_wrapper_sm100": (".grouped_gemm_glu_hadamard.api", "grouped_gemm_glu_hadamard_wrapper_sm100"), ++ "GroupedGemmDgluSm100": (".grouped_gemm_dglu.api", "GroupedGemmDgluSm100"), ++ "grouped_gemm_dglu_wrapper_sm100": (".grouped_gemm_dglu.api", "grouped_gemm_dglu_wrapper_sm100"), ++ "GroupedGemmWgradSm100": (".grouped_gemm_wgrad.api", "GroupedGemmWgradSm100"), ++ "grouped_gemm_wgrad_wrapper_sm100": (".grouped_gemm_wgrad.api", "grouped_gemm_wgrad_wrapper_sm100"), ++} ++ ++ ++def __getattr__(name): ++ if name not in _LAZY_IMPORTS: ++ raise AttributeError(name) ++ module_name, attr_name = _LAZY_IMPORTS[name] ++ value = getattr(import_module(module_name, package=__name__), attr_name) ++ globals()[name] = value ++ return value ++ ++ ++__all__ = list(_LAZY_IMPORTS) +diff --git a/python/cudnn/grouped_gemm/grouped_gemm_swiglu/__init__.py b/python/cudnn/grouped_gemm/grouped_gemm_swiglu/__init__.py +index 6bcd935..77483cf 100644 +--- a/python/cudnn/grouped_gemm/grouped_gemm_swiglu/__init__.py ++++ b/python/cudnn/grouped_gemm/grouped_gemm_swiglu/__init__.py +@@ -8,12 +8,36 @@ This module provides the forward grouped GEMM with SwiGLU activation + for MoE (Mixture of Experts) workloads on SM100+ GPUs. + """ + +-from .api import ( +- GroupedGemmSwigluSm100, +- grouped_gemm_swiglu_wrapper_sm100, +-) ++from importlib import import_module ++ ++ ++_LAZY_IMPORTS = { ++ "GroupedGemmSwigluSm100": (".api", "GroupedGemmSwigluSm100"), ++ "grouped_gemm_swiglu_wrapper_sm100": (".api", "grouped_gemm_swiglu_wrapper_sm100"), ++ "GroupedGemmElementType": (".compile", "GroupedGemmElementType"), ++ "GroupedGemmOperandDesc": (".compile", "GroupedGemmOperandDesc"), ++ "GroupedGemmSwigluConfig": (".compile", "GroupedGemmSwigluConfig"), ++ "GroupedGemmAbiEntry": (".compile", "GroupedGemmAbiEntry"), ++ "GroupedGemmAbiDescriptor": (".compile", "GroupedGemmAbiDescriptor"), ++ "compile_grouped_gemm_swiglu": (".compile", "compile_grouped_gemm_swiglu"), ++} ++ ++ ++def __getattr__(name): ++ if name not in _LAZY_IMPORTS: ++ raise AttributeError(name) ++ module_name, attr_name = _LAZY_IMPORTS[name] ++ value = getattr(import_module(module_name, package=__name__), attr_name) ++ globals()[name] = value ++ return value + + __all__ = [ + "GroupedGemmSwigluSm100", + "grouped_gemm_swiglu_wrapper_sm100", ++ "GroupedGemmElementType", ++ "GroupedGemmOperandDesc", ++ "GroupedGemmSwigluConfig", ++ "GroupedGemmAbiEntry", ++ "GroupedGemmAbiDescriptor", ++ "compile_grouped_gemm_swiglu", + ] +diff --git a/python/cudnn/grouped_gemm/grouped_gemm_swiglu/compile.py b/python/cudnn/grouped_gemm/grouped_gemm_swiglu/compile.py +new file mode 100644 +index 0000000..04f5d30 +--- /dev/null ++++ b/python/cudnn/grouped_gemm/grouped_gemm_swiglu/compile.py +@@ -0,0 +1,391 @@ ++# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. ++# SPDX-License-Identifier: MIT ++ ++"""Framework-neutral compile-only API for grouped GEMM + SwiGLU. ++ ++This module deliberately has no PyTorch dependency. It consumes tensor ++metadata, compiles a native TVM-FFI function, and returns a self-describing ++ABI. It allocates no framework buffers and never launches the kernel. ++""" ++ ++from __future__ import annotations ++ ++from dataclasses import dataclass ++from enum import Enum ++import hashlib ++import json ++import threading ++from typing import Any, Literal, Mapping ++ ++from cuda.bindings import driver as cuda ++import cutlass ++import cutlass.cute as cute ++from cutlass.cute.runtime import make_fake_stream ++ ++from .grouped_gemm_swiglu_quant import BlockScaledContiguousGroupedGemmKernel ++ ++ ++class GroupedGemmElementType(str, Enum): ++ """Framework-neutral element types accepted by the compile API.""" ++ ++ FLOAT4_E2M1FN_X2 = "float4_e2m1fn_x2" ++ FLOAT8_E4M3FN = "float8_e4m3fn" ++ FLOAT8_E5M2 = "float8_e5m2" ++ FLOAT8_E8M0FNU = "float8_e8m0fnu" ++ FLOAT16 = "float16" ++ BFLOAT16 = "bfloat16" ++ FLOAT32 = "float32" ++ INT32 = "int32" ++ UINT8 = "uint8" ++ ++ ++_CUTLASS_DTYPES = { ++ GroupedGemmElementType.FLOAT4_E2M1FN_X2: "Float4E2M1FN", ++ GroupedGemmElementType.FLOAT8_E4M3FN: "Float8E4M3FN", ++ GroupedGemmElementType.FLOAT8_E5M2: "Float8E5M2", ++ GroupedGemmElementType.FLOAT8_E8M0FNU: "Float8E8M0FNU", ++ GroupedGemmElementType.FLOAT16: "Float16", ++ GroupedGemmElementType.BFLOAT16: "BFloat16", ++ GroupedGemmElementType.FLOAT32: "Float32", ++ GroupedGemmElementType.INT32: "Int32", ++ GroupedGemmElementType.UINT8: "Uint8", ++} ++ ++ ++def _normalize_dtype(dtype: GroupedGemmElementType | str) -> GroupedGemmElementType: ++ if isinstance(dtype, GroupedGemmElementType): ++ return dtype ++ return GroupedGemmElementType(dtype) ++ ++ ++def _cutlass_dtype(dtype: GroupedGemmElementType | str) -> type[cutlass.Numeric]: ++ dtype = _normalize_dtype(dtype) ++ try: ++ return getattr(cutlass, _CUTLASS_DTYPES[dtype]) ++ except AttributeError as exc: ++ raise RuntimeError(f"Installed nvidia-cutlass-dsl does not provide {dtype.value}") from exc ++ ++ ++@dataclass(frozen=True) ++class GroupedGemmOperandDesc: ++ """Storage-free metadata for one logical kernel operand.""" ++ ++ role: str ++ dtype: GroupedGemmElementType | str ++ shape: tuple[int, ...] ++ stride_order: tuple[int, ...] ++ align: int = 16 ++ ++ def __post_init__(self): ++ dtype = _normalize_dtype(self.dtype) ++ shape = tuple(int(dim) for dim in self.shape) ++ stride_order = tuple(int(dim) for dim in self.stride_order) ++ if not self.role: ++ raise ValueError("Operand role must not be empty") ++ if not shape or any(dim < 0 for dim in shape): ++ raise ValueError(f"{self.role}: shape must contain non-negative dimensions, got {shape}") ++ if tuple(sorted(stride_order)) != tuple(range(len(shape))): ++ raise ValueError(f"{self.role}: invalid stride_order {stride_order} for rank {len(shape)}") ++ if self.align <= 0 or self.align & (self.align - 1): ++ raise ValueError(f"{self.role}: align must be a positive power of two, got {self.align}") ++ object.__setattr__(self, "dtype", dtype) ++ object.__setattr__(self, "shape", shape) ++ object.__setattr__(self, "stride_order", stride_order) ++ ++ ++@dataclass(frozen=True) ++class GroupedGemmSwigluConfig: ++ """Compile-time configuration for the SM100 grouped SwiGLU kernel.""" ++ ++ sf_vec_size: int = 32 ++ acc_dtype: GroupedGemmElementType | str = GroupedGemmElementType.FLOAT32 ++ mma_tiler_mn: tuple[int, int] = (256, 256) ++ cluster_shape_mn: tuple[int, int] = (2, 1) ++ vector_f32: bool = False ++ discrete_col_sfd: bool = True ++ use_mono_increase_expert_idx: bool = True ++ num_cluster_overlap_margin: int = 0 ++ ++ def __post_init__(self): ++ object.__setattr__(self, "acc_dtype", _normalize_dtype(self.acc_dtype)) ++ object.__setattr__(self, "mma_tiler_mn", tuple(self.mma_tiler_mn)) ++ object.__setattr__(self, "cluster_shape_mn", tuple(self.cluster_shape_mn)) ++ if self.sf_vec_size not in (16, 32): ++ raise ValueError(f"sf_vec_size must be 16 or 32, got {self.sf_vec_size}") ++ if self.acc_dtype != GroupedGemmElementType.FLOAT32: ++ raise ValueError("Grouped GEMM SwiGLU requires a float32 accumulator") ++ ++ ++@dataclass(frozen=True) ++class GroupedGemmAbiEntry: ++ """One positional slot in the returned native TVM-FFI function.""" ++ ++ role: str ++ kind: Literal["arg", "ret", "stream"] ++ dtype: GroupedGemmElementType | None ++ shape: tuple[int, ...] = () ++ stride_order: tuple[int, ...] = () ++ optional: bool = False ++ ++ ++@dataclass(frozen=True) ++class GroupedGemmAbiDescriptor: ++ """The exact positional ABI and stable registration key for a compile.""" ++ ++ entries: tuple[GroupedGemmAbiEntry, ...] ++ key: str ++ ++ @property ++ def arg_spec(self) -> tuple[str, ...]: ++ """jax-tvm-ffi routing for this native, inputs-then-results ABI.""" ++ kinds = tuple(entry.kind for entry in self.entries) ++ first_ret = next((idx for idx, kind in enumerate(kinds) if kind == "ret"), len(kinds)) ++ first_stream = next((idx for idx, kind in enumerate(kinds) if kind == "stream"), len(kinds)) ++ if any(kind != "arg" for kind in kinds[:first_ret]): ++ raise ValueError("ABI args must be contiguous") ++ if any(kind != "ret" for kind in kinds[first_ret:first_stream]): ++ raise ValueError("ABI returns must be contiguous") ++ if kinds[first_stream:] != ("stream",): ++ raise ValueError("ABI must contain exactly one trailing stream") ++ return ("args", "rets", "ctx.stream") ++ ++ ++_INPUT_ROLES = ("a", "b", "sfa", "sfb", "padded_offsets", "alpha", "prob", "norm_const") ++_OUTPUT_ROLES = ("c", "d", "d_col", "sfd_row", "sfd_col") ++_ALL_ROLES = _INPUT_ROLES + _OUTPUT_ROLES ++_COMPILE_CACHE: dict[str, tuple[Any, GroupedGemmAbiDescriptor]] = {} ++_COMPILE_LOCK = threading.Lock() ++ ++ ++def _ceil_div(value: int, divisor: int) -> int: ++ return (value + divisor - 1) // divisor ++ ++ ++def _scale_size(rows: int, cols: int, sf_vec_size: int, *, colwise: bool) -> int: ++ if colwise: ++ rows, cols = cols, rows ++ return 32 * 4 * _ceil_div(rows, 128) * 4 * _ceil_div(_ceil_div(cols, sf_vec_size), 4) ++ ++ ++def _expect_shape(operand: GroupedGemmOperandDesc, shape: tuple[int, ...]) -> None: ++ if operand.shape != shape: ++ raise ValueError(f"{operand.role}: expected shape {shape}, got {operand.shape}") ++ ++ ++def _validate_operands(operands: Mapping[str, GroupedGemmOperandDesc], config: GroupedGemmSwigluConfig) -> None: ++ missing = sorted(set(_ALL_ROLES) - set(operands)) ++ extra = sorted(set(operands) - set(_ALL_ROLES)) ++ if missing or extra: ++ raise ValueError(f"Invalid operand roles: missing={missing}, extra={extra}") ++ for role, operand in operands.items(): ++ if operand.role != role: ++ raise ValueError(f"Operand mapping key {role!r} does not match descriptor role {operand.role!r}") ++ ++ a, b = operands["a"], operands["b"] ++ if len(a.shape) != 3 or a.shape[2] != 1: ++ raise ValueError(f"a: expected [M,K,1], got {a.shape}") ++ if len(b.shape) != 3: ++ raise ValueError(f"b: expected [N,K,E], got {b.shape}") ++ m, k, _ = a.shape ++ n, b_k, experts = b.shape ++ if k != b_k: ++ raise ValueError(f"A/B K mismatch: {k} != {b_k}") ++ if m % BlockScaledContiguousGroupedGemmKernel.FIX_PAD_SIZE: ++ raise ValueError(f"M must be {BlockScaledContiguousGroupedGemmKernel.FIX_PAD_SIZE}-aligned, got {m}") ++ if n % 2: ++ raise ValueError(f"Combined SwiGLU N must be even, got {n}") ++ if experts > 1024: ++ raise ValueError(f"Expert count must be <= 1024, got {experts}") ++ ++ _expect_shape(operands["c"], (m, n, 1)) ++ _expect_shape(operands["d"], (m, n // 2, 1)) ++ _expect_shape(operands["d_col"], (m, n // 2, 1)) ++ _expect_shape(operands["padded_offsets"], (experts,)) ++ _expect_shape(operands["alpha"], (experts,)) ++ _expect_shape(operands["prob"], (m, 1, 1)) ++ _expect_shape(operands["norm_const"], (1,)) ++ expected_sfa = _scale_size(m, k, config.sf_vec_size, colwise=False) ++ expected_sfb = _scale_size(n, k, config.sf_vec_size, colwise=False) * experts ++ if len(operands["sfa"].shape) != 1 or operands["sfa"].shape[0] < expected_sfa: ++ raise ValueError(f"sfa: expected a flat buffer with at least {expected_sfa} values") ++ if len(operands["sfb"].shape) != 1 or operands["sfb"].shape[0] < expected_sfb: ++ raise ValueError(f"sfb: expected a flat buffer with at least {expected_sfb} values") ++ _expect_shape(operands["sfd_row"], (_scale_size(m, n // 2, config.sf_vec_size, colwise=False),)) ++ _expect_shape(operands["sfd_col"], (_scale_size(m, n // 2, config.sf_vec_size, colwise=True),)) ++ ++ if a.dtype != b.dtype or a.dtype not in (GroupedGemmElementType.FLOAT8_E4M3FN, GroupedGemmElementType.FLOAT8_E5M2): ++ raise ValueError(f"A/B must have one matching FP8 dtype, got {a.dtype}/{b.dtype}") ++ sf_dtype = GroupedGemmElementType.FLOAT8_E8M0FNU ++ for role in ("sfa", "sfb", "sfd_row", "sfd_col"): ++ if operands[role].dtype != sf_dtype: ++ raise ValueError(f"{role}: expected {sf_dtype.value}, got {operands[role].dtype.value}") ++ if operands["c"].dtype not in (GroupedGemmElementType.BFLOAT16, GroupedGemmElementType.FLOAT16): ++ raise ValueError(f"c: expected bfloat16 or float16, got {operands['c'].dtype.value}") ++ if operands["d"].dtype != operands["d_col"].dtype: ++ raise ValueError("d and d_col must have matching dtypes") ++ if operands["d"].dtype not in ( ++ GroupedGemmElementType.FLOAT8_E4M3FN, ++ GroupedGemmElementType.FLOAT8_E5M2, ++ ): ++ raise ValueError(f"d/d_col must have an FP8 dtype, got {operands['d'].dtype.value}") ++ for role in ("alpha", "prob", "norm_const"): ++ if operands[role].dtype != GroupedGemmElementType.FLOAT32: ++ raise ValueError(f"{role}: expected float32") ++ if operands["padded_offsets"].dtype != GroupedGemmElementType.INT32: ++ raise ValueError("padded_offsets: expected int32") ++ ++ ++def _cache_key(operands: Mapping[str, GroupedGemmOperandDesc], config: GroupedGemmSwigluConfig) -> str: ++ payload = { ++ "operands": [ ++ { ++ "role": role, ++ "dtype": operands[role].dtype.value, ++ "shape": operands[role].shape, ++ "stride_order": operands[role].stride_order, ++ "align": operands[role].align, ++ } ++ for role in _ALL_ROLES ++ ], ++ "config": { ++ "sf_vec_size": config.sf_vec_size, ++ "acc_dtype": config.acc_dtype.value, ++ "mma_tiler_mn": config.mma_tiler_mn, ++ "cluster_shape_mn": config.cluster_shape_mn, ++ "vector_f32": config.vector_f32, ++ "discrete_col_sfd": config.discrete_col_sfd, ++ "use_mono_increase_expert_idx": config.use_mono_increase_expert_idx, ++ "num_cluster_overlap_margin": config.num_cluster_overlap_margin, ++ }, ++ } ++ digest = hashlib.sha256(json.dumps(payload, sort_keys=True).encode()).hexdigest()[:24] ++ return f"cudnn.grouped_gemm_swiglu.{digest}" ++ ++ ++def _make_fake(desc: GroupedGemmOperandDesc): ++ return cute.runtime.make_fake_compact_tensor( ++ dtype=_cutlass_dtype(desc.dtype), ++ shape=desc.shape, ++ stride_order=desc.stride_order, ++ assumed_align=desc.align, ++ ) ++ ++ ++def _compile_native( ++ operands: Mapping[str, GroupedGemmOperandDesc], ++ config: GroupedGemmSwigluConfig, ++ key: str, ++) -> tuple[Any, GroupedGemmAbiDescriptor]: ++ kernel = BlockScaledContiguousGroupedGemmKernel( ++ sf_vec_size=config.sf_vec_size, ++ acc_dtype=_cutlass_dtype(config.acc_dtype), ++ use_2cta_instrs=config.mma_tiler_mn[0] == 256, ++ mma_tiler_mn=config.mma_tiler_mn, ++ cluster_shape_mn=config.cluster_shape_mn, ++ vector_f32=config.vector_f32, ++ generate_sfd=True, ++ discrete_col_sfd=config.discrete_col_sfd, ++ expert_cnt=operands["padded_offsets"].shape[0], ++ use_mono_increase_expert_idx=config.use_mono_increase_expert_idx, ++ ) ++ hardware_info = cutlass.utils.HardwareInfo() ++ max_active_clusters = hardware_info.get_max_active_clusters(config.cluster_shape_mn[0] * config.cluster_shape_mn[1]) ++ max_active_clusters -= config.num_cluster_overlap_margin ++ if max_active_clusters <= 0: ++ raise ValueError("num_cluster_overlap_margin leaves no active clusters") ++ ++ # XLA's FFI buffer ABI does not carry strides, and jax-tvm-ffi therefore ++ # exposes every buffer to TVM FFI as compact row-major. B is logically ++ # [N,K,E] with K minor, which cannot be represented by a compact ++ # [N,K,E] buffer when E > 1. Make the public ABI's B slot the equivalent ++ # compact [E,N,K] storage and restore the logical CuTe view in the native ++ # launcher. This is a view only; no framework-specific code or copy is ++ # introduced on the execution path. ++ b_desc = operands["b"] ++ b_n, b_k, b_experts = b_desc.shape ++ ffi_operands = dict(operands) ++ ffi_operands["b"] = GroupedGemmOperandDesc( ++ role="b", ++ dtype=b_desc.dtype, ++ shape=(b_experts, b_n, b_k), ++ stride_order=(2, 1, 0), ++ align=b_desc.align, ++ ) ++ ++ # jax-tvm-ffi routes complete input and result groups. Compile a native ++ # reorder launcher with that stable grouping while keeping the volatile ++ # raw-kernel ABI private to cuDNN-FE. This remains native code: no Python ++ # callback is present on the execution path. ++ @cute.jit ++ def launch(a, b, sfa, sfb, padded_offsets, alpha, prob, norm_const, c, d, d_col, sfd_row, sfd_col, stream: cuda.CUstream): ++ b_logical = cute.make_tensor( ++ b.iterator, ++ cute.make_layout( ++ (b_n, b_k, b_experts), ++ stride=(b_k, 1, b_n * b_k), ++ ), ++ ) ++ kernel( ++ a, ++ b_logical, ++ c, ++ d, ++ d_col, ++ sfa, ++ sfb, ++ sfd_row, ++ sfd_col, ++ None, ++ norm_const, ++ padded_offsets, ++ alpha, ++ prob, ++ max_active_clusters, ++ stream, ++ ) ++ ++ fake_operands = [_make_fake(ffi_operands[role]) for role in _ALL_ROLES] ++ fake_stream = make_fake_stream(use_tvm_ffi_env_stream=False) ++ compiled = cute.compile(launch, *fake_operands, fake_stream, options="--enable-tvm-ffi") ++ entries = tuple( ++ GroupedGemmAbiEntry(role, "arg", ffi_operands[role].dtype, ffi_operands[role].shape, ffi_operands[role].stride_order) ++ for role in _INPUT_ROLES ++ ) + tuple( ++ GroupedGemmAbiEntry(role, "ret", ffi_operands[role].dtype, ffi_operands[role].shape, ffi_operands[role].stride_order) ++ for role in _OUTPUT_ROLES ++ ) + (GroupedGemmAbiEntry("stream", "stream", None),) ++ descriptor = GroupedGemmAbiDescriptor(entries=entries, key=key) ++ descriptor.arg_spec # Validate the grouped ABI before publishing it. ++ return compiled, descriptor ++ ++ ++def compile_grouped_gemm_swiglu( ++ *, ++ operands: Mapping[str, GroupedGemmOperandDesc], ++ config: GroupedGemmSwigluConfig | None = None, ++) -> tuple[Any, GroupedGemmAbiDescriptor]: ++ """Compile and cache a native grouped GEMM + SwiGLU TVM-FFI function. ++ ++ The call consumes metadata only. It does not allocate output buffers and ++ does not launch. The returned descriptor is the positional ABI authority ++ for the exact returned function. ++ """ ++ config = GroupedGemmSwigluConfig() if config is None else config ++ _validate_operands(operands, config) ++ key = _cache_key(operands, config) ++ with _COMPILE_LOCK: ++ cached = _COMPILE_CACHE.get(key) ++ if cached is None: ++ cached = _compile_native(operands, config, key) ++ _COMPILE_CACHE[key] = cached ++ return cached ++ ++ ++__all__ = [ ++ "GroupedGemmElementType", ++ "GroupedGemmOperandDesc", ++ "GroupedGemmSwigluConfig", ++ "GroupedGemmAbiEntry", ++ "GroupedGemmAbiDescriptor", ++ "compile_grouped_gemm_swiglu", ++] +diff --git a/python/cudnn/grouped_gemm/utils.py b/python/cudnn/grouped_gemm/utils.py +index cb60dc8..1287293 100644 +--- a/python/cudnn/grouped_gemm/utils.py ++++ b/python/cudnn/grouped_gemm/utils.py +@@ -33,7 +33,7 @@ This module contains the tile scheduler classes and helper functions used by bot + the forward (grouped_gemm_swiglu) and backward (grouped_gemm_dswiglu) kernels. + """ + +-from typing import Tuple, Union ++from typing import TYPE_CHECKING, Tuple, Union + + from cutlass.cutlass_dsl import ( + Boolean, +@@ -52,7 +52,6 @@ from cutlass.cutlass_dsl import T + from cutlass.cute.typing import Float32, Int32 + import cutlass.cute as cute + import cutlass +-import torch + import cutlass.pipeline as pipeline + from cutlass.pipeline import ( + Agent, +@@ -61,6 +60,9 @@ from cutlass.pipeline import ( + make_pipeline_state, + ) + ++if TYPE_CHECKING: ++ import torch ++ + ############################################################################## + # Helper functions + ############################################################################## +@@ -201,8 +203,10 @@ def fmax(a: Union[float, Float32], b: Union[float, Float32], *, loc=None, ip=Non + ) + + +-def logical_shape_fp4x2_aware(tensor: torch.Tensor) -> Tuple[int, ...]: ++def logical_shape_fp4x2_aware(tensor: "torch.Tensor") -> Tuple[int, ...]: + """Return correct shapes for NVFP4 tensor.""" ++ import torch ++ + if tensor.dtype == torch.float4_e2m1fn_x2: + innermost_dim_index = next((i for i, s in enumerate(tensor.stride()) if s == 1), None) + if innermost_dim_index is None: +diff --git a/test/python/fe_api/grouped_gemm/test_grouped_gemm_swiglu_compile.py b/test/python/fe_api/grouped_gemm/test_grouped_gemm_swiglu_compile.py +new file mode 100644 +index 0000000..82e1eee +--- /dev/null ++++ b/test/python/fe_api/grouped_gemm/test_grouped_gemm_swiglu_compile.py +@@ -0,0 +1,96 @@ ++# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. ++# SPDX-License-Identifier: MIT ++ ++"""Tests for the framework-neutral grouped SwiGLU compile API.""" ++ ++import subprocess ++import sys ++ ++import pytest ++import torch ++ ++ ++def _operands(m=256, k=128, n=256, experts=1): ++ from cudnn import GroupedGemmElementType as DType ++ from cudnn import GroupedGemmOperandDesc as Operand ++ ++ def scale_size(rows, cols, colwise=False): ++ if colwise: ++ rows, cols = cols, rows ++ ceil_div = lambda x, y: (x + y - 1) // y ++ return 32 * 4 * ceil_div(rows, 128) * 4 * ceil_div(ceil_div(cols, 32), 4) ++ ++ return { ++ "a": Operand("a", DType.FLOAT8_E4M3FN, (m, k, 1), (1, 0, 2)), ++ "b": Operand("b", DType.FLOAT8_E4M3FN, (n, k, experts), (1, 0, 2)), ++ "sfa": Operand("sfa", DType.FLOAT8_E8M0FNU, (scale_size(m, k),), (0,)), ++ "sfb": Operand("sfb", DType.FLOAT8_E8M0FNU, (scale_size(n, k) * experts,), (0,)), ++ "padded_offsets": Operand("padded_offsets", DType.INT32, (experts,), (0,)), ++ "alpha": Operand("alpha", DType.FLOAT32, (experts,), (0,)), ++ "prob": Operand("prob", DType.FLOAT32, (m, 1, 1), (1, 0, 2)), ++ "norm_const": Operand("norm_const", DType.FLOAT32, (1,), (0,)), ++ "c": Operand("c", DType.BFLOAT16, (m, n, 1), (1, 0, 2)), ++ "d": Operand("d", DType.FLOAT8_E4M3FN, (m, n // 2, 1), (1, 0, 2)), ++ "d_col": Operand("d_col", DType.FLOAT8_E4M3FN, (m, n // 2, 1), (1, 0, 2)), ++ "sfd_row": Operand("sfd_row", DType.FLOAT8_E8M0FNU, (scale_size(m, n // 2),), (0,)), ++ "sfd_col": Operand("sfd_col", DType.FLOAT8_E8M0FNU, (scale_size(m, n // 2, colwise=True),), (0,)), ++ } ++ ++ ++@pytest.mark.L0 ++def test_compile_api_imports_without_torch(): ++ script = r""" ++import importlib.abc ++import sys ++ ++class BlockTorch(importlib.abc.MetaPathFinder): ++ def find_spec(self, fullname, path=None, target=None): ++ if fullname == "torch" or fullname.startswith("torch."): ++ raise ModuleNotFoundError("torch intentionally blocked") ++ return None ++ ++sys.meta_path.insert(0, BlockTorch()) ++from cudnn import compile_grouped_gemm_swiglu ++assert callable(compile_grouped_gemm_swiglu) ++""" ++ subprocess.run([sys.executable, "-c", script], check=True) ++ ++ ++@pytest.mark.L0 ++def test_compile_api_rejects_unaligned_m_before_compile(): ++ from cudnn import compile_grouped_gemm_swiglu ++ ++ with pytest.raises(ValueError, match="M must be 256-aligned"): ++ compile_grouped_gemm_swiglu(operands=_operands(m=128)) ++ ++ ++@pytest.mark.L1 ++def test_compile_api_returns_native_function_and_self_describing_abi(): ++ if torch.cuda.get_device_capability()[0] < 10: ++ pytest.skip("Grouped GEMM SwiGLU compile requires SM100+") ++ tvm_ffi = pytest.importorskip("tvm_ffi") ++ from cudnn import compile_grouped_gemm_swiglu ++ ++ function, abi = compile_grouped_gemm_swiglu(operands=_operands()) ++ cached_function, cached_abi = compile_grouped_gemm_swiglu(operands=_operands()) ++ ++ assert isinstance(function, tvm_ffi.Function) ++ assert cached_function is function ++ assert cached_abi == abi ++ assert abi.arg_spec == ("args", "rets", "ctx.stream") ++ assert [(entry.role, entry.kind) for entry in abi.entries] == [ ++ ("a", "arg"), ++ ("b", "arg"), ++ ("sfa", "arg"), ++ ("sfb", "arg"), ++ ("padded_offsets", "arg"), ++ ("alpha", "arg"), ++ ("prob", "arg"), ++ ("norm_const", "arg"), ++ ("c", "ret"), ++ ("d", "ret"), ++ ("d_col", "ret"), ++ ("sfd_row", "ret"), ++ ("sfd_col", "ret"), ++ ("stream", "stream"), ++ ] +-- +2.50.0 diff --git a/tests/jax/requirements_cutedsl.txt b/tests/jax/requirements_cutedsl.txt index 9cf81761a31..ab925e61870 100644 --- a/tests/jax/requirements_cutedsl.txt +++ b/tests/jax/requirements_cutedsl.txt @@ -1,4 +1,8 @@ # Known-good baseline for the JAX CuTeDSL MoE pipeclean. jax[cuda13]>=0.9.1 -nvidia-cudnn-frontend==1.25.0 nvidia-cutlass-dsl[cu13]==4.5.2 +apache-tvm-ffi>=0.1.12 +jax-tvm-ffi>=0.1.3 + +# Install the sibling cuDNN-FE fork separately with its torch-free compiler extra: +# pip install -e "../cudnn-frontend[cutedsl-compile]" diff --git a/tests/jax/run_tvm_ffi_fused_moe_e2e.sh b/tests/jax/run_tvm_ffi_fused_moe_e2e.sh new file mode 100755 index 00000000000..48c6000176c --- /dev/null +++ b/tests/jax/run_tvm_ffi_fused_moe_e2e.sh @@ -0,0 +1,92 @@ +#!/usr/bin/env bash +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +# End-to-end remote SM100 validation for the cuDNN-FE compile-only API, +# JAX TVM-FFI forward primitive, fused backward, partitioning, and MoEBlock. + +set -euo pipefail + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +TE_ROOT="$(cd "$SCRIPT_DIR/../.." && pwd)" +CUDNN_FE_ROOT="${CUDNN_FE_ROOT:-$(cd "$TE_ROOT/.." && pwd)/cudnn-frontend}" +CUDNN_FE_REPOSITORY="${CUDNN_FE_REPOSITORY:-https://github.com/NVIDIA/cudnn-frontend.git}" +CUDNN_FE_BASE_REV="${CUDNN_FE_BASE_REV:-fbc713624ac403800ae286e422b5243d694abab2}" +CUDNN_FE_PATCH="$SCRIPT_DIR/patches/0001-Add-framework-neutral-grouped-SwiGLU-compiler.patch" +INSTALL_DEPS="${INSTALL_DEPS:-1}" +NUM_GPUS="${NUM_GPUS:-$(nvidia-smi -L | wc -l)}" +ARTIFACT_DIR="${E2E_ARTIFACT_DIR:-$TE_ROOT/tvm_ffi_e2e_$(date +%Y%m%d_%H%M%S)}" + +if [ ! -f "$CUDNN_FE_ROOT/pyproject.toml" ]; then + if [ ! -f "$CUDNN_FE_PATCH" ]; then + echo "Missing bundled cuDNN-FE patch: $CUDNN_FE_PATCH" >&2 + exit 1 + fi + mkdir -p "$(dirname "$CUDNN_FE_ROOT")" + git clone --filter=blob:none "$CUDNN_FE_REPOSITORY" "$CUDNN_FE_ROOT" + git -C "$CUDNN_FE_ROOT" checkout --detach "$CUDNN_FE_BASE_REV" +fi + +CUDNN_FE_COMPILE_API="$CUDNN_FE_ROOT/python/cudnn/grouped_gemm/grouped_gemm_swiglu/compile.py" +if [ ! -f "$CUDNN_FE_COMPILE_API" ]; then + if ! git -C "$CUDNN_FE_ROOT" apply --check "$CUDNN_FE_PATCH"; then + echo "The bundled compile-only patch does not apply cleanly to $CUDNN_FE_ROOT." >&2 + echo "Use the pinned base $CUDNN_FE_BASE_REV or provide an already-patched checkout." >&2 + exit 1 + fi + git -C "$CUDNN_FE_ROOT" apply "$CUDNN_FE_PATCH" +fi +if [ "$NUM_GPUS" -lt 2 ]; then + echo "The E2E suite requires at least two SM100 GPUs" >&2 + exit 1 +fi + +mkdir -p "$ARTIFACT_DIR" +export XLA_PYTHON_CLIENT_PREALLOCATE="${XLA_PYTHON_CLIENT_PREALLOCATE:-false}" +export XLA_PYTHON_CLIENT_MEM_FRACTION="${XLA_PYTHON_CLIENT_MEM_FRACTION:-0.5}" +export NVTE_CUDA_ARCHS="${NVTE_CUDA_ARCHS:-100a}" +export NVTE_CMAKE_BUILD_DIR="${NVTE_CMAKE_BUILD_DIR:-$TE_ROOT/build_jax}" + +if [ "$INSTALL_DEPS" = "1" ]; then + python3 -m pip install -U pip setuptools wheel cmake ninja pybind11 pytest + python3 -m pip install -r "$TE_ROOT/tests/jax/requirements_cutedsl.txt" + python3 -m pip install "$CUDNN_FE_ROOT[cutedsl-compile]" + NVTE_FRAMEWORK=jax python3 -m pip install --no-build-isolation -e "$TE_ROOT" +fi + +python3 - <<'PY' | tee "$ARTIFACT_DIR/preflight.log" +import inspect +import jax +import jax_tvm_ffi +import tvm_ffi +from cudnn import compile_grouped_gemm_swiglu + +print("jax", jax.__version__) +print("devices", jax.devices()) +print("cudnn_frontend", inspect.getfile(compile_grouped_gemm_swiglu)) +print("jax_tvm_ffi", inspect.getfile(jax_tvm_ffi)) +print("tvm_ffi", inspect.getfile(tvm_ffi)) +assert all(device.platform == "gpu" for device in jax.devices()) +PY + +python3 -m pytest -q "$TE_ROOT/tests/jax/test_cutedsl_moe.py" \ + -k "compile_only_api or swiglu_forward_fused_output_parity or dswiglu_backward_quantized_output_parity" \ + 2>&1 | tee "$ARTIFACT_DIR/standalone.log" + +NUM_GPUS=2 TVM_FFI_MP_LOG_DIR="$ARTIFACT_DIR/partitioning" \ + bash "$TE_ROOT/tests/jax/run_tvm_ffi_grouped_mlp_multiprocess.sh" \ + 2>&1 | tee "$ARTIFACT_DIR/partitioning.log" + +if [ "$NUM_GPUS" -ge 4 ]; then + NUM_GPUS=4 \ + TE_EP_MOE_MP_LOG_DIR="$ARTIFACT_DIR/moe" \ + NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION=1 \ + bash "$TE_ROOT/tests/jax/run_te_ep_moe.sh" -k TestTeEpMoeCudnnCutedslFusion \ + 2>&1 | tee "$ARTIFACT_DIR/moe.log" +else + echo "SKIPPED integrated MoEBlock test: four GPUs are required" \ + | tee "$ARTIFACT_DIR/moe.log" +fi + +echo "TVM-FFI fused MoE E2E validation passed. Artifacts: $ARTIFACT_DIR" diff --git a/tests/jax/run_tvm_ffi_grouped_mlp_multiprocess.sh b/tests/jax/run_tvm_ffi_grouped_mlp_multiprocess.sh new file mode 100755 index 00000000000..9ae6945f5e6 --- /dev/null +++ b/tests/jax/run_tvm_ffi_grouped_mlp_multiprocess.sh @@ -0,0 +1,62 @@ +#!/usr/bin/env bash +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +set -euo pipefail + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +TEST_FILE="$SCRIPT_DIR/tvm_ffi_grouped_mlp_multiprocess.py" +NUM_GPUS="${NUM_GPUS:-2}" +COORDINATOR="${TVM_FFI_COORDINATOR_ADDRESS:-127.0.0.1:13461}" +TIMEOUT_SECONDS="${TEST_TIMEOUT_S:-600}" +LOG_DIR="${TVM_FFI_MP_LOG_DIR:-$(mktemp -d -t te_tvm_ffi_mp_XXXXXX)}" + +if [ "$(nvidia-smi -L | wc -l)" -lt "$NUM_GPUS" ]; then + echo "Need at least $NUM_GPUS GPUs for the TVM-FFI multiprocess test" >&2 + exit 1 +fi + +export XLA_PYTHON_CLIENT_PREALLOCATE="${XLA_PYTHON_CLIENT_PREALLOCATE:-false}" +mkdir -p "$LOG_DIR" +PIDS=() + +cleanup() { + for pid in "${PIDS[@]:-}"; do + kill -TERM "$pid" 2>/dev/null || true + done +} +trap cleanup EXIT INT TERM + +for rank in $(seq 0 $((NUM_GPUS - 1))); do + log_file="$LOG_DIR/rank_${rank}.log" + if [ "$rank" -eq 0 ]; then + timeout --foreground --signal=KILL "$TIMEOUT_SECONDS" \ + python3 "$TEST_FILE" "$COORDINATOR" "$rank" "$NUM_GPUS" 2>&1 \ + | tee "$log_file" & + else + timeout --foreground --signal=KILL "$TIMEOUT_SECONDS" \ + python3 "$TEST_FILE" "$COORDINATOR" "$rank" "$NUM_GPUS" \ + >"$log_file" 2>&1 & + fi + PIDS+=("$!") +done + +failed=0 +for pid in "${PIDS[@]}"; do + if ! wait "$pid"; then + failed=1 + fi +done +PIDS=() + +if [ "$failed" -ne 0 ]; then + echo "TVM-FFI multiprocess test failed; logs: $LOG_DIR" >&2 + for log_file in "$LOG_DIR"/*.log; do + echo "--- $log_file ---" >&2 + tail -80 "$log_file" >&2 || true + done + exit 1 +fi + +echo "TVM-FFI multiprocess test passed; logs: $LOG_DIR" diff --git a/tests/jax/test_cutedsl_moe.py b/tests/jax/test_cutedsl_moe.py index b46aac78c7c..3b5a583aeed 100644 --- a/tests/jax/test_cutedsl_moe.py +++ b/tests/jax/test_cutedsl_moe.py @@ -3,6 +3,7 @@ # See LICENSE for license information. """Tests for the JAX binding to cuDNN frontend's CuTeDSL MoE kernel.""" +import importlib import sys import jax @@ -13,8 +14,6 @@ from transformer_engine.jax import cpp_extensions as tex from transformer_engine.jax.cutedsl_extensions.moe import ( grouped_gemm_dswiglu_mxfp8, - grouped_gemm_swiglu_mxfp8, - load_grouped_gemm_swiglu_kernel, pack_swiglu_pair, unpack_swiglu_pair, ) @@ -25,6 +24,8 @@ TensorUsage, ) +_RAGGED_EIGHT_EXPERT_GROUP_SIZES = (512, 256, 1024, 512, 768, 256, 256, 512) + def test_swiglu_block_pack_round_trip(): """Gate/up packing alternates 32-column blocks and is reversible.""" @@ -39,33 +40,38 @@ def test_swiglu_block_pack_round_trip(): np.testing.assert_array_equal(unpacked_up, up) -def test_direct_kernel_loader_bypasses_torch_api(): - """The direct source loader must not execute cuDNN's Torch API package.""" - kernel_cls = load_grouped_gemm_swiglu_kernel() - assert kernel_cls.__name__ == "BlockScaledContiguousGroupedGemmKernel" +def test_compile_only_api_bypasses_torch_wrapper(): + """The forward compiler must not execute cuDNN's Torch wrapper module.""" + from cudnn import compile_grouped_gemm_swiglu + + assert callable(compile_grouped_gemm_swiglu) assert "cudnn.grouped_gemm.grouped_gemm_swiglu.api" not in sys.modules - # cuDNN frontend 1.25's shared utility source imports torch only for an - # unused annotation. The loader supplies a non-executable sentinel rather - # than importing the Torch package. - torch_module = sys.modules.get("torch") - assert torch_module is None or getattr( - torch_module, "__transformer_engine_cutedsl_stub__", False - ) -def test_swiglu_forward_fused_output_parity(): +@pytest.mark.parametrize( + "group_sizes_tuple", + [ + pytest.param((256,), id="one-expert"), + pytest.param(_RAGGED_EIGHT_EXPERT_GROUP_SIZES, id="eight-expert-ragged"), + ], +) +def test_swiglu_forward_fused_output_parity(group_sizes_tuple): """The forward fused call matches TE projection plus JAX SwiGLU reference.""" try: from transformer_engine_jax import get_device_compute_capability if get_device_compute_capability(0) != 100: pytest.skip("cuDNN frontend grouped GEMM SwiGLU requires SM100") - load_grouped_gemm_swiglu_kernel() - import cutlass.jax # noqa: F401 # pylint: disable=unused-import,import-outside-toplevel + dependencies_available, dependency_error = ( + tex.grouped_gemm_swiglu_dependencies_available() + ) + if not dependencies_available: + pytest.skip(f"TVM-FFI JAX dependencies are unavailable: {dependency_error}") except (ImportError, RuntimeError) as exc: pytest.skip(f"CuTeDSL JAX dependencies are unavailable: {exc}") - experts, rows, hidden, intermediate = 1, 256, 128, 128 + experts = len(group_sizes_tuple) + rows, hidden, intermediate = sum(group_sizes_tuple), 128, 128 key = jax.random.PRNGKey(123) x = jax.random.normal(key, (rows, hidden), dtype=jnp.bfloat16) wi_0 = jax.random.normal( @@ -75,7 +81,9 @@ def test_swiglu_forward_fused_output_parity(): jax.random.fold_in(key, 2), (experts, hidden, intermediate), dtype=jnp.bfloat16 ) wi = pack_swiglu_pair(wi_0, wi_1) - group_sizes = jnp.asarray([rows], dtype=jnp.int32) + group_sizes = jnp.asarray(group_sizes_tuple, dtype=jnp.int32) + assert len(set(group_sizes_tuple)) > 1 or experts == 1 + assert all(offset % 256 == 0 for offset in np.cumsum(group_sizes_tuple)) quantizers = QuantizerFactory.create_set( scaling_mode=ScalingMode.MXFP8_1D_SCALING, fwd_dtype=jnp.float8_e4m3fn, @@ -97,8 +105,18 @@ def run(x_arg, wi_arg): casted_wi, contracting_dims=((1,), (1,)), ) + reference_gate, reference_up = unpack_swiglu_pair(reference) + swiglu_reference = jax.nn.silu(reference_gate) * reference_up + quantized_reference = tex.grouped_quantize( + swiglu_reference, + quantizers.x, + group_sizes, + flatten_axis=-1, + ) + reference_row = quantized_reference.get_tensor(TensorUsage.LHS) + reference_col = quantized_reference.get_tensor(TensorUsage.LHS_TRANS) physical_wi = casted_wi.data.reshape(experts, hidden, 2 * intermediate).transpose(0, 2, 1) - combined, swiglu_row, swiglu_col, scale_row, scale_col = grouped_gemm_swiglu_mxfp8( + combined, swiglu_row, swiglu_col, scale_row, scale_col = tex.grouped_gemm_swiglu( casted_x.data.reshape(rows, hidden, 1), physical_wi, casted_x.scale_inv, @@ -115,9 +133,34 @@ def run(x_arg, wi_arg): swiglu_col, scale_row, scale_col, + reference_row.data, + reference_row.scale_inv, + reference_col.data, + reference_col.scale_inv, ) - reference, combined, swiglu_row, swiglu_col, scale_row, scale_col = run(x, wi) + # Lowering is deliberately separate from execution: this proves that the + # cuDNN-FE compiler consumes only abstract descriptors and returns a native + # TVM-FFI target without requiring live device buffers. + run.lower(x, wi) + grouped_gemm_swiglu_module = importlib.import_module( + "transformer_engine.jax.cpp_extensions.grouped_gemm_swiglu" + ) + + assert grouped_gemm_swiglu_module._REGISTERED_TARGETS # pylint: disable=protected-access + + ( + reference, + combined, + swiglu_row, + swiglu_col, + scale_row, + scale_col, + reference_row, + reference_scale_row, + reference_col, + reference_scale_col, + ) = run(x, wi) jax.block_until_ready((reference, combined, swiglu_row, swiglu_col)) np.testing.assert_array_equal(combined, reference) assert swiglu_row.shape == swiglu_col.shape == (rows, intermediate, 1) @@ -126,10 +169,14 @@ def run(x_arg, wi_arg): gate, up = unpack_swiglu_pair(combined) swiglu_reference = np.asarray(jax.nn.silu(gate) * up, dtype=np.float32) - for is_colwise, payload, scale in ( - (False, swiglu_row, scale_row), - (True, swiglu_col, scale_col), + for is_colwise, payload, scale, reference_payload, reference_scale in ( + (False, swiglu_row, scale_row, reference_row, reference_scale_row), + (True, swiglu_col, scale_col, reference_col, reference_scale_col), ): + # The kernel emits the compact scale payload. TE grouped tensors use a + # larger metadata-compatible buffer and pad the unused tail in the + # production MoE path. + scale = jnp.pad(scale, (0, reference_scale.size - scale.size)) scaled = ScaledTensorFactory.create_1x( payload.reshape(-1), scale, @@ -142,12 +189,31 @@ def run(x_arg, wi_arg): original_shape=(rows, intermediate), pre_swizzled=True, ) - dequantized = np.asarray(scaled.dequantize()[0], dtype=np.float32) - assert np.all(np.isfinite(dequantized)) - relative_error = np.linalg.norm(dequantized - swiglu_reference) / np.linalg.norm( - swiglu_reference + reference_scaled = ScaledTensorFactory.create_1x( + reference_payload.reshape(-1), + reference_scale, + scaling_mode=ScalingMode.MXFP8_1D_SCALING, + dq_dtype=jnp.bfloat16, + is_colwise=is_colwise, + data_layout="N", + flatten_axis=1, + first_dims=group_sizes, + original_shape=(rows, intermediate), + pre_swizzled=True, + ) + dequantized = np.asarray(jnp.concatenate(scaled.dequantize(), axis=0), dtype=np.float32) + reference_dequantized = np.asarray( + jnp.concatenate(reference_scaled.dequantize(), axis=0), dtype=np.float32 ) - assert relative_error < 0.05 + assert np.all(np.isfinite(dequantized)) + swiglu_relative_error = np.linalg.norm( + dequantized - swiglu_reference + ) / np.linalg.norm(swiglu_reference) + quantization_relative_error = np.linalg.norm( + dequantized - reference_dequantized + ) / np.linalg.norm(reference_dequantized) + assert swiglu_relative_error < 0.05 + assert quantization_relative_error < 0.05 def test_dswiglu_backward_quantized_output_parity(): @@ -157,7 +223,6 @@ def test_dswiglu_backward_quantized_output_parity(): if get_device_compute_capability(0) != 100: pytest.skip("cuDNN frontend grouped GEMM dSwiGLU requires SM100") - load_grouped_gemm_swiglu_kernel() from transformer_engine.jax.cutedsl_extensions.moe import load_grouped_gemm_dswiglu_kernel load_grouped_gemm_dswiglu_kernel() diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index cfabb91aef5..ee4dcfbc9fe 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -741,6 +741,23 @@ def loss_fn(vars_arg, x_arg): assert np.all(np.isfinite(fused_np)) params_np = _params_global_numpy(variables, mesh) + # MoEBlock callers provide ordinary routing decisions, not the + # CuTeDSL kernel's 256-padded expert boundaries. Verify this seeded + # case is genuinely ragged before the internal dispatch/padding path. + routing_logits = jnp.asarray(jax.device_get(x)).reshape(-1, HIDDEN) @ jnp.asarray( + params_np["gate_kernel"] + ).astype(DTYPE) + _, routed_experts = jax.lax.top_k(routing_logits.astype(jnp.float32), k=TOPK) + raw_routing_counts = np.asarray( + jax.device_get(jnp.bincount(routed_experts.reshape(-1), length=NUM_EXPERTS)) + ) + assert np.unique(raw_routing_counts).size > 1, ( + f"expected unequal raw routing counts, got {raw_routing_counts}" + ) + assert np.any(raw_routing_counts % _CUDNN_CUTEDSL_ALIGN_SIZE != 0), ( + "MoEBlock regression must exercise unaligned caller-visible routing counts; " + f"got {raw_routing_counts}" + ) reference, _ = _pure_jax_moe_reference( jnp.asarray(jax.device_get(x)), jnp.asarray(params_np["gate_kernel"]), @@ -794,22 +811,34 @@ def reference_loss(params, x_arg): relative_error < 0.35 ), f"{name} fused MXFP8 VJP relative error {relative_error:.4f} exceeds 0.35" - # Exercise the failure mode seen in MaxText: apply one optimizer-like - # update and require the next forward pass to remain finite. - updated_variables = jax.tree_util.tree_map( - lambda param, grad: param - jnp.asarray(1e-3, param.dtype) * grad.astype(param.dtype), - variables, - grads, - ) + # Exercise the failure mode seen in MaxText over several optimizer-like + # steps. A one-step check can miss ABI/layout corruption that appears + # only after the newly quantized weights feed the next backward pass. + training_variables = variables + losses = [] with _ctx(mesh), autocast( enabled=True, recipe=MXFP8BlockScaling(), mesh_resource=mesh_resource, ): - updated_output, _ = jax.jit(block.apply)(updated_variables, _shard_inputs(x, mesh)) + compiled_loss_and_grad = jax.jit(jax.value_and_grad(loss_fn)) + for _ in range(3): + step_loss, step_grads = compiled_loss_and_grad(training_variables, x_sh) + training_variables = jax.tree_util.tree_map( + lambda param, grad: param + - jnp.asarray(1e-3, param.dtype) * grad.astype(param.dtype), + training_variables, + step_grads, + ) + losses.append(step_loss) + updated_output, _ = jax.jit(block.apply)(training_variables, _shard_inputs(x, mesh)) updated_output.block_until_ready() + jax.block_until_ready(losses) updated_output_np = _to_global_numpy(updated_output, mesh).astype(np.float32) - assert np.all(np.isfinite(updated_output_np)), "post-update output has NaN/Inf" + loss_values = np.asarray([jax.device_get(loss) for loss in losses], dtype=np.float32) + assert np.all(np.isfinite(loss_values)), f"training losses have NaN/Inf: {loss_values}" + assert loss_values[-1] <= 1.2 * loss_values[0], f"training loss diverged: {loss_values}" + assert np.all(np.isfinite(updated_output_np)), "post-training output has NaN/Inf" def test_maxtext_shape_vjp_update_stays_finite(self, mesh): """Regression for the NaN observed after MaxText's first update.""" diff --git a/tests/jax/tvm_ffi_grouped_mlp_multiprocess.py b/tests/jax/tvm_ffi_grouped_mlp_multiprocess.py new file mode 100644 index 00000000000..6536771a239 --- /dev/null +++ b/tests/jax/tvm_ffi_grouped_mlp_multiprocess.py @@ -0,0 +1,359 @@ +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. +"""Multiprocess partitioning and ragged-parity experiments for grouped SwiGLU.""" + +from __future__ import annotations + +import sys + +import jax +import jax.numpy as jnp +import numpy as np +from jax.experimental import multihost_utils +from jax.experimental.shard_map import shard_map +from jax.sharding import Mesh, NamedSharding, PartitionSpec as P + +_RAGGED_EIGHT_EXPERT_GROUP_SIZES = (512, 256, 1024, 512, 768, 256, 256, 512) + + +def _scale_size(rows: int, cols: int, *, colwise: bool = False) -> int: + if colwise: + rows, cols = cols, rows + ceil_div = lambda value, divisor: (value + divisor - 1) // divisor + return 32 * 4 * ceil_div(rows, 128) * 4 * ceil_div(ceil_div(cols, 32), 4) + + +def _global_array(local_array, mesh, spec): + mesh_devices = tuple(np.asarray(mesh.devices).reshape(-1)) + if all(device.process_index == jax.process_index() for device in mesh_devices): + return jax.device_put(local_array, NamedSharding(mesh, spec)) + return multihost_utils.host_local_array_to_global_array(local_array, mesh, spec) + + +def _run_matrix_cell(mesh: Mesh, label: str) -> None: + from transformer_engine.jax import cpp_extensions as tex + + local_rows, hidden, intermediate, experts = 256, 128, 128, 1 + combined = 2 * intermediate + process_key = jax.random.fold_in(jax.random.PRNGKey(2026), jax.process_index()) + x_local = jax.random.uniform( + process_key, + (local_rows, hidden, 1), + minval=-0.5, + maxval=0.5, + dtype=jnp.float32, + ).astype(jnp.float8_e4m3fn) + weight_local = jax.random.uniform( + jax.random.PRNGKey(17), + (experts, combined, hidden), + minval=-0.5, + maxval=0.5, + dtype=jnp.float32, + ).astype(jnp.float8_e4m3fn) + sfa_local = jnp.ones( + (_scale_size(local_rows, hidden),), dtype=jnp.float8_e8m0fnu + ) + sfb_local = jnp.ones( + (_scale_size(combined, hidden) * experts,), dtype=jnp.float8_e8m0fnu + ) + prob_local = jnp.ones((local_rows, 1, 1), dtype=jnp.float32) + + data_axis = mesh.axis_names[0] + x = _global_array(x_local, mesh, P(data_axis, None, None)) + weight = _global_array(weight_local, mesh, P()) + sfa = _global_array(sfa_local, mesh, P(data_axis)) + sfb = _global_array(sfb_local, mesh, P()) + prob = _global_array(prob_local, mesh, P(data_axis, None, None)) + padded_offsets = _global_array( + jnp.asarray([x.shape[0]], dtype=jnp.int32), mesh, P() + ) + + def fused(a, b, a_scale, b_scale, offsets, probabilities): + return tex.grouped_gemm_swiglu( + a, + b, + a_scale, + b_scale, + offsets, + probabilities, + compute_dtype=jnp.bfloat16, + output_dtype=jnp.float8_e4m3fn, + ) + + custom_outputs = jax.jit(fused)(x, weight, sfa, sfb, padded_offsets, prob) + + def local_fused(a, b, a_scale, b_scale, _offsets, probabilities): + local_offsets = jnp.asarray([a.shape[0]], dtype=jnp.int32) + return fused(a, b, a_scale, b_scale, local_offsets, probabilities) + + mapped_fused = shard_map( + local_fused, + mesh=mesh, + in_specs=( + P(data_axis, None, None), + P(), + P(data_axis), + P(), + P(), + P(data_axis, None, None), + ), + out_specs=[ + P(data_axis, None, None), + P(data_axis, None, None), + P(data_axis, None, None), + P(data_axis), + P(data_axis), + ], + check_rep=False, + ) + shard_map_outputs = jax.jit(mapped_fused)(x, weight, sfa, sfb, padded_offsets, prob) + jax.block_until_ready((custom_outputs, shard_map_outputs)) + + def metrics(custom, mapped, a, b): + reference = jnp.matmul( + a.reshape(a.shape[0], a.shape[1]).astype(jnp.bfloat16), + b[0].astype(jnp.bfloat16).T, + ).reshape(a.shape[0], b.shape[1], 1) + custom_c = custom[0].astype(jnp.float32) + reference = reference.astype(jnp.float32) + custom_vs_mapped = jnp.max( + jnp.stack( + [ + jnp.max(jnp.abs(lhs.astype(jnp.float32) - rhs.astype(jnp.float32))) + for lhs, rhs in zip(custom, mapped) + ] + ) + ) + relative_reference = jnp.linalg.norm(custom_c - reference) / jnp.maximum( + jnp.linalg.norm(reference), 1e-6 + ) + finite = jnp.all( + jnp.stack( + [jnp.all(jnp.isfinite(value.astype(jnp.float32))) for value in custom] + ) + ) + return custom_vs_mapped, relative_reference, finite + + custom_vs_mapped, relative_reference, finite = jax.jit(metrics)( + custom_outputs, shard_map_outputs, x, weight + ) + custom_vs_mapped = float(custom_vs_mapped) + relative_reference = float(relative_reference) + finite = bool(finite) + if custom_vs_mapped > 1e-3: + raise AssertionError(f"{label}: custom_partitioning vs shard_map max diff {custom_vs_mapped}") + if relative_reference > 0.05: + raise AssertionError(f"{label}: fused vs matmul relative error {relative_reference}") + if not finite: + raise AssertionError(f"{label}: non-finite fused output") + if jax.process_index() == 0: + print( + f"PASSED {label}: custom_vs_shard_map={custom_vs_mapped:.3e}, " + f"relative_reference={relative_reference:.3e}", + flush=True, + ) + + +def _run_ragged_multiprocess_cell(mesh: Mesh) -> None: + """Exercise eight unequal experts through a global multiprocess shard_map.""" + from transformer_engine.jax import cpp_extensions as tex + from transformer_engine.jax.cutedsl_extensions.moe import unpack_swiglu_pair + from transformer_engine.jax.quantize import ( + QuantizerFactory, + ScaledTensorFactory, + ScalingMode, + TensorUsage, + ) + + group_sizes_tuple = _RAGGED_EIGHT_EXPERT_GROUP_SIZES + group_sizes = jnp.asarray(group_sizes_tuple, dtype=jnp.int32) + local_rows = sum(group_sizes_tuple) + experts, hidden, intermediate = len(group_sizes_tuple), 128, 128 + combined = 2 * intermediate + assert experts >= 8 + assert len(set(group_sizes_tuple)) > 1 + assert all(offset % 256 == 0 for offset in np.cumsum(group_sizes_tuple)) + + process_key = jax.random.fold_in(jax.random.PRNGKey(20260723), jax.process_index()) + x_local = jax.random.uniform( + process_key, + (local_rows, hidden, 1), + minval=-0.5, + maxval=0.5, + dtype=jnp.float32, + ).astype(jnp.float8_e4m3fn) + weight_local = jax.random.uniform( + jax.random.PRNGKey(20260724), + (experts, combined, hidden), + minval=-0.5, + maxval=0.5, + dtype=jnp.float32, + ).astype(jnp.float8_e4m3fn) + sfa_local = jnp.ones( + (_scale_size(local_rows, hidden),), dtype=jnp.float8_e8m0fnu + ) + sfb_local = jnp.ones( + (_scale_size(combined, hidden) * experts,), dtype=jnp.float8_e8m0fnu + ) + prob_local = jnp.ones((local_rows, 1, 1), dtype=jnp.float32) + + data_axis = mesh.axis_names[0] + x = _global_array(x_local, mesh, P(data_axis, None, None)) + weight = _global_array(weight_local, mesh, P()) + sfa = _global_array(sfa_local, mesh, P(data_axis)) + sfb = _global_array(sfb_local, mesh, P()) + prob = _global_array(prob_local, mesh, P(data_axis, None, None)) + + def local_fused(a, b, a_scale, b_scale, probabilities): + return tex.grouped_gemm_swiglu( + a, + b, + a_scale, + b_scale, + jnp.cumsum(group_sizes), + probabilities, + compute_dtype=jnp.bfloat16, + output_dtype=jnp.float8_e4m3fn, + ) + + mapped_fused = shard_map( + local_fused, + mesh=mesh, + in_specs=( + P(data_axis, None, None), + P(), + P(data_axis), + P(), + P(data_axis, None, None), + ), + out_specs=[ + P(data_axis, None, None), + P(data_axis, None, None), + P(data_axis, None, None), + P(data_axis), + P(data_axis), + ], + check_rep=False, + ) + outputs = jax.jit(mapped_fused)(x, weight, sfa, sfb, prob) + jax.block_until_ready(outputs) + combined_local, row_local, col_local, row_scale_local, col_scale_local = ( + output.addressable_data(0) for output in outputs + ) + + boundaries = np.cumsum((0, *group_sizes_tuple)) + reference_parts = [ + jnp.matmul( + x_local[boundaries[i] : boundaries[i + 1], :, 0].astype(jnp.bfloat16), + weight_local[i].astype(jnp.bfloat16).T, + ) + for i in range(experts) + ] + combined_reference = jnp.concatenate(reference_parts, axis=0) + np.testing.assert_array_equal( + np.asarray(combined_local[:, :, 0]), + np.asarray(combined_reference), + ) + gate, up = unpack_swiglu_pair(combined_reference) + swiglu_reference = jax.nn.silu(gate) * up + + quantizers = QuantizerFactory.create_set( + scaling_mode=ScalingMode.MXFP8_1D_SCALING, + fwd_dtype=jnp.float8_e4m3fn, + bwd_dtype=jnp.float8_e5m2, + is_2x2x=True, + n_groups=experts, + ) + quantized_reference = tex.grouped_quantize( + swiglu_reference, + quantizers.x, + group_sizes, + flatten_axis=-1, + ) + reference_row = quantized_reference.get_tensor(TensorUsage.LHS) + reference_col = quantized_reference.get_tensor(TensorUsage.LHS_TRANS) + + errors = [] + for is_colwise, payload, scale, reference_tensor in ( + (False, row_local, row_scale_local, reference_row), + (True, col_local, col_scale_local, reference_col), + ): + expected_scale_size = ScalingMode.MXFP8_1D_SCALING.get_grouped_scale_shape( + (local_rows, intermediate), + experts, + is_colwise, + is_padded=True, + flatten_axis=1, + )[0] + scale = jnp.pad(scale, (0, expected_scale_size - scale.size)) + scaled = ScaledTensorFactory.create_1x( + payload.reshape(-1), + scale, + scaling_mode=ScalingMode.MXFP8_1D_SCALING, + dq_dtype=jnp.bfloat16, + is_colwise=is_colwise, + data_layout="N", + flatten_axis=1, + first_dims=group_sizes, + original_shape=(local_rows, intermediate), + pre_swizzled=True, + ) + dequantized = jnp.concatenate(scaled.dequantize(), axis=0).astype(jnp.float32) + reference_dequantized = jnp.concatenate( + reference_tensor.dequantize(), axis=0 + ).astype(jnp.float32) + swiglu_error = jnp.linalg.norm( + dequantized - swiglu_reference.astype(jnp.float32) + ) / jnp.linalg.norm(swiglu_reference.astype(jnp.float32)) + quantization_error = jnp.linalg.norm( + dequantized - reference_dequantized + ) / jnp.linalg.norm(reference_dequantized) + errors.append((float(swiglu_error), float(quantization_error))) + if errors[-1][0] >= 0.05 or errors[-1][1] >= 0.05: + raise AssertionError( + "eight-expert ragged fused quantization parity failed: " + f"is_colwise={is_colwise}, swiglu_error={errors[-1][0]:.4e}, " + f"quantization_error={errors[-1][1]:.4e}" + ) + + if jax.process_index() == 0: + print( + "PASSED multiprocess eight-expert ragged grouped GEMM/SwiGLU/MXFP8 parity: " + f"group_sizes={group_sizes_tuple}, row_errors={errors[0]}, col_errors={errors[1]}", + flush=True, + ) + + +def main() -> None: + if len(sys.argv) != 4: + raise SystemExit(f"usage: {sys.argv[0]} COORDINATOR PROCESS_ID NUM_PROCESSES") + coordinator, process_id, num_processes = sys.argv[1], int(sys.argv[2]), int(sys.argv[3]) + jax.distributed.initialize( + coordinator_address=coordinator, + num_processes=num_processes, + process_id=process_id, + local_device_ids=process_id, + ) + import transformer_engine # Load the TE core library before the raw JAX extension. + from transformer_engine_jax import get_device_compute_capability + + if get_device_compute_capability(0) != 100: + raise RuntimeError("TVM-FFI grouped SwiGLU experiment requires SM100") + local_mesh = Mesh(np.asarray(jax.local_devices()), ("data",)) + global_mesh = Mesh(np.asarray(jax.devices()), ("data",)) + with local_mesh: + _run_matrix_cell(local_mesh, "single-device shard_map/custom_partitioning") + multihost_utils.sync_global_devices("single-device-cell") + with global_mesh: + _run_matrix_cell(global_mesh, "multiprocess shard_map/custom_partitioning") + multihost_utils.sync_global_devices("multiprocess-cell") + with global_mesh: + _run_ragged_multiprocess_cell(global_mesh) + multihost_utils.sync_global_devices("multiprocess-eight-expert-ragged-cell") + if jax.process_index() == 0: + print("PASSED TVM-FFI grouped SwiGLU partitioning matrix", flush=True) + + +if __name__ == "__main__": + main() diff --git a/transformer_engine/jax/cpp_extensions/__init__.py b/transformer_engine/jax/cpp_extensions/__init__.py index c9647afb826..a78916ff95a 100644 --- a/transformer_engine/jax/cpp_extensions/__init__.py +++ b/transformer_engine/jax/cpp_extensions/__init__.py @@ -10,6 +10,7 @@ from .quantization import * from .softmax import * from .gemm import * +from .grouped_gemm_swiglu import * from .router import * from .ep import * from .topk import * diff --git a/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py b/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py new file mode 100644 index 00000000000..42d83933d81 --- /dev/null +++ b/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py @@ -0,0 +1,350 @@ +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. +"""TVM-FFI JAX primitive for cuDNN-FE's fused grouped SwiGLU kernel.""" + +from __future__ import annotations + +import threading +from typing import Any + +import jax +import jax.numpy as jnp +from jax.sharding import NamedSharding, PartitionSpec + +from .base import BasePrimitive, register_primitive +from .misc import get_padded_spec + +__all__ = ["grouped_gemm_swiglu", "grouped_gemm_swiglu_dependencies_available"] + + +_INPUT_ROLES = ("a", "b", "sfa", "sfb", "padded_offsets", "alpha", "prob", "norm_const") +_OUTPUT_ROLES = ("c", "d", "d_col", "sfd_row", "sfd_col") +_REGISTERED_TARGETS: dict[str, tuple[Any, Any]] = {} +_REGISTRATION_LOCK = threading.Lock() + + +def _ceil_div(value: int, divisor: int) -> int: + return (value + divisor - 1) // divisor + + +def _scale_size(rows: int, cols: int, *, colwise: bool) -> int: + if colwise: + rows, cols = cols, rows + return 32 * 4 * _ceil_div(rows, 128) * 4 * _ceil_div(_ceil_div(cols, 32), 4) + + +def _dtype_name(dtype) -> str: + dtype = jnp.dtype(dtype) + supported = { + jnp.dtype(jnp.float8_e4m3fn): "float8_e4m3fn", + jnp.dtype(jnp.float8_e5m2): "float8_e5m2", + jnp.dtype(jnp.float8_e8m0fnu): "float8_e8m0fnu", + jnp.dtype(jnp.float16): "float16", + jnp.dtype(jnp.bfloat16): "bfloat16", + jnp.dtype(jnp.float32): "float32", + jnp.dtype(jnp.int32): "int32", + } + try: + return supported[dtype] + except KeyError as exc: + raise ValueError(f"Unsupported grouped SwiGLU dtype {dtype}") from exc + + +def grouped_gemm_swiglu_dependencies_available() -> tuple[bool, str]: + """Check optional runtime dependencies without compiling a kernel.""" + try: + import jax_tvm_ffi # noqa: F401 # pylint: disable=unused-import,import-outside-toplevel + import tvm_ffi # noqa: F401 # pylint: disable=unused-import,import-outside-toplevel + from cudnn import compile_grouped_gemm_swiglu # noqa: F401 + except (ImportError, ModuleNotFoundError, RuntimeError, AttributeError) as exc: + return False, str(exc) + return True, "" + + +def _output_avals(a_aval, b_aval, compute_dtype, output_dtype): + if a_aval.ndim != 3 or a_aval.shape[-1] != 1: + raise ValueError(f"Expected A[M,K,1], got {a_aval.shape}") + if b_aval.ndim != 3: + raise ValueError(f"Expected physical B[E,N,K], got {b_aval.shape}") + m, k_a, _ = a_aval.shape + _, n, k_b = b_aval.shape + if k_a != k_b: + raise ValueError(f"A K={k_a} does not match B K={k_b}") + if m % 256: + raise ValueError(f"Padded activation rows M={m} must be divisible by 256") + if n % 2: + raise ValueError(f"Combined SwiGLU N={n} must be even") + intermediate = n // 2 + return ( + jax.core.ShapedArray((m, n, 1), compute_dtype), + jax.core.ShapedArray((m, intermediate, 1), output_dtype), + jax.core.ShapedArray((m, intermediate, 1), output_dtype), + jax.core.ShapedArray((_scale_size(m, intermediate, colwise=False),), jnp.float8_e8m0fnu), + jax.core.ShapedArray((_scale_size(m, intermediate, colwise=True),), jnp.float8_e8m0fnu), + ) + + +def _compile_and_register(avals_in, avals_out): + try: + import jax_tvm_ffi + import tvm_ffi + from cudnn import ( + GroupedGemmOperandDesc, + GroupedGemmSwigluConfig, + compile_grouped_gemm_swiglu, + ) + except (ImportError, ModuleNotFoundError, RuntimeError, AttributeError) as exc: + raise RuntimeError( + "TVM-FFI grouped SwiGLU requires the patched cuDNN-FE compiler package, " + "apache-tvm-ffi, and jax-tvm-ffi" + ) from exc + + layouts = { + "a": (1, 0, 2), + "b": (1, 0, 2), + "sfa": (0,), + "sfb": (0,), + "padded_offsets": (0,), + "alpha": (0,), + "prob": (1, 0, 2), + "norm_const": (0,), + "c": (1, 0, 2), + "d": (1, 0, 2), + "d_col": (1, 0, 2), + "sfd_row": (0,), + "sfd_col": (0,), + } + operands = {} + for role, aval in zip(_INPUT_ROLES, avals_in): + shape = aval.shape + if role == "b": + # The framework ABI is compact [E,N,K]. The cuDNN-FE compiler + # describes and restores the kernel's logical [N,K,E] view. + shape = (aval.shape[1], aval.shape[2], aval.shape[0]) + operands[role] = GroupedGemmOperandDesc( + role=role, + dtype=_dtype_name(aval.dtype), + shape=shape, + stride_order=layouts[role], + ) + for role, aval in zip(_OUTPUT_ROLES, avals_out): + operands[role] = GroupedGemmOperandDesc( + role=role, + dtype=_dtype_name(aval.dtype), + shape=aval.shape, + stride_order=layouts[role], + ) + + function, abi = compile_grouped_gemm_swiglu( + operands=operands, + config=GroupedGemmSwigluConfig( + sf_vec_size=32, + mma_tiler_mn=(256, 256), + cluster_shape_mn=(2, 1), + discrete_col_sfd=True, + ), + ) + if not isinstance(function, tvm_ffi.Function): + raise TypeError( + "cuDNN-FE compile_grouped_gemm_swiglu returned a Python wrapper; " + "the JAX execution target must be a native tvm_ffi.Function" + ) + expected = [(role, "arg") for role in _INPUT_ROLES] + [ + (role, "ret") for role in _OUTPUT_ROLES + ] + [("stream", "stream")] + actual = [(entry.role, entry.kind) for entry in abi.entries] + if actual != expected: + raise ValueError(f"Unsupported grouped SwiGLU ABI: expected {expected}, got {actual}") + + target_name = f"te_grouped_gemm_swiglu.{abi.key.rsplit('.', maxsplit=1)[-1]}" + with _REGISTRATION_LOCK: + if target_name not in _REGISTERED_TARGETS: + jax_tvm_ffi.register_ffi_target( + target_name, + function, + arg_spec=list(abi.arg_spec), + platform="gpu", + allow_cuda_graph=True, + ) + # Keep both the native function and its exact ABI alive for the + # lifetime of the process-global XLA target registration. + _REGISTERED_TARGETS[target_name] = (function, abi) + return target_name, abi + + +class GroupedGemmSwigluPrimitive(BasePrimitive): + """Fused grouped GEMM, SwiGLU, and dual MXFP8 quantization.""" + + name = "te_grouped_gemm_swiglu_tvm_ffi" + multiple_results = True + impl_static_args = (8, 9) + inner_primitive = None + outer_primitive = None + + @staticmethod + def abstract( + a_aval, + b_aval, + sfa_aval, + sfb_aval, + padded_offsets_aval, + alpha_aval, + prob_aval, + norm_const_aval, + *, + compute_dtype, + output_dtype, + ): + del sfa_aval, sfb_aval, alpha_aval, norm_const_aval + m = a_aval.shape[0] + experts = b_aval.shape[0] + if padded_offsets_aval.shape != (experts,): + raise ValueError( + f"Expected padded_offsets shape {(experts,)}, got {padded_offsets_aval.shape}" + ) + if prob_aval.shape != (m, 1, 1): + raise ValueError(f"Expected prob shape {(m, 1, 1)}, got {prob_aval.shape}") + return _output_avals(a_aval, b_aval, compute_dtype, output_dtype) + + @staticmethod + def lowering(ctx, *args, compute_dtype, output_dtype): + del compute_dtype, output_dtype + target_name, abi = _compile_and_register(ctx.avals_in, ctx.avals_out) + operand_layouts = [entry.stride_order for entry in abi.entries if entry.kind == "arg"] + result_layouts = [entry.stride_order for entry in abi.entries if entry.kind == "ret"] + return jax.ffi.ffi_lowering( + target_name, + operand_layouts=operand_layouts, + result_layouts=result_layouts, + )(ctx, *args) + + @staticmethod + def impl( + a, + b, + sfa, + sfb, + padded_offsets, + alpha, + prob, + norm_const, + compute_dtype, + output_dtype, + ): + if GroupedGemmSwigluPrimitive.inner_primitive is None: + raise RuntimeError("GroupedGemmSwigluPrimitive has not been registered") + return GroupedGemmSwigluPrimitive.inner_primitive.bind( + a, + b, + sfa, + sfb, + padded_offsets, + alpha, + prob, + norm_const, + compute_dtype=compute_dtype, + output_dtype=output_dtype, + ) + + @staticmethod + def batcher(batched_args, batch_dims, *, compute_dtype, output_dtype): + del batched_args, compute_dtype, output_dtype + raise NotImplementedError( + f"GroupedGemmSwigluPrimitive does not support vmap batch dimensions {batch_dims}" + ) + + @staticmethod + def partition(compute_dtype, output_dtype, mesh, arg_infos, result_infos): + del result_infos + a_spec = get_padded_spec(arg_infos[0]) + m_axis = a_spec[0] + experts = arg_infos[1].shape[0] + if experts != 1 and m_axis is not None: + raise NotImplementedError( + "Token-axis custom partitioning is currently supported only for the standalone " + "single-expert experiment; MoEBlock integration remains inside shard_map" + ) + + arg_shardings = ( + NamedSharding(mesh, PartitionSpec(m_axis, None, None)), + NamedSharding(mesh, PartitionSpec(None, None, None)), + NamedSharding(mesh, PartitionSpec(m_axis)), + NamedSharding(mesh, PartitionSpec(None)), + NamedSharding(mesh, PartitionSpec(None)), + NamedSharding(mesh, PartitionSpec(None)), + NamedSharding(mesh, PartitionSpec(m_axis, None, None)), + NamedSharding(mesh, PartitionSpec(None)), + ) + out_shardings = [ + NamedSharding(mesh, PartitionSpec(m_axis, None, None)), + NamedSharding(mesh, PartitionSpec(m_axis, None, None)), + NamedSharding(mesh, PartitionSpec(m_axis, None, None)), + NamedSharding(mesh, PartitionSpec(m_axis)), + NamedSharding(mesh, PartitionSpec(m_axis)), + ] + + def sharded_impl(a, b, sfa, sfb, padded_offsets, alpha, prob, norm_const): + if experts == 1: + padded_offsets = jnp.asarray([a.shape[0]], dtype=jnp.int32) + return GroupedGemmSwigluPrimitive.impl( + a, + b, + sfa, + sfb, + padded_offsets, + alpha, + prob, + norm_const, + compute_dtype=compute_dtype, + output_dtype=output_dtype, + ) + + return mesh, sharded_impl, out_shardings, arg_shardings + + @staticmethod + def shardy_sharding_rule(*args, **kwargs): + del args, kwargs + return ( + "m k one, experts n k, sfa, sfb, experts, experts, m prob_one_a prob_one_b, norm -> " + "m n one, m h one, m h one, sfd_row, sfd_col" + ) + + +register_primitive(GroupedGemmSwigluPrimitive) + + +def grouped_gemm_swiglu( + a: jax.Array, + b: jax.Array, + sfa: jax.Array, + sfb: jax.Array, + padded_offsets: jax.Array, + prob: jax.Array, + *, + compute_dtype, + output_dtype, +) -> tuple[jax.Array, jax.Array, jax.Array, jax.Array, jax.Array]: + """Run the TVM-FFI fused grouped SwiGLU primitive. + + ``b`` uses TE's compact physical ``[E,N,K]`` representation. The native + cuDNN-FE launcher restores the kernel's logical ``[N,K,E]`` view because + XLA FFI buffers do not expose strides to jax-tvm-ffi. + """ + if b.ndim != 3: + raise ValueError(f"Expected physical B[E,N,K], got {b.shape}") + experts = b.shape[0] + alpha = jnp.ones((experts,), dtype=jnp.float32) + norm_const = jnp.ones((1,), dtype=jnp.float32) + return GroupedGemmSwigluPrimitive.outer_primitive.bind( + a, + b, + sfa.reshape(-1), + sfb.reshape(-1), + padded_offsets.astype(jnp.int32), + alpha, + prob.astype(jnp.float32), + norm_const, + compute_dtype=jnp.dtype(compute_dtype), + output_dtype=jnp.dtype(output_dtype), + ) diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index 113472176a0..5529033937e 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -93,7 +93,6 @@ def _cudnn_cutedsl_fusion_rejection_reasons( from .cutedsl_extensions.moe import ( load_grouped_gemm_dswiglu_kernel, - load_grouped_gemm_swiglu_kernel, ) from .quantize import GroupedQuantizer, ScalingMode @@ -164,14 +163,14 @@ def _cudnn_cutedsl_fusion_rejection_reasons( except (AttributeError, RuntimeError) as exc: errors.append(f"CUTLASS JAX runtime check failed: {exc}") - for kernel_name, loader in ( - ("SwiGLU forward", load_grouped_gemm_swiglu_kernel), - ("dSwiGLU backward", load_grouped_gemm_dswiglu_kernel), - ): - try: - loader() - except (ImportError, ModuleNotFoundError, RuntimeError) as exc: - errors.append(f"could not load the cuDNN frontend {kernel_name} kernel: {exc}") + dependencies_available, dependency_error = tex.grouped_gemm_swiglu_dependencies_available() + if not dependencies_available: + errors.append(f"could not load the TVM-FFI CuTeDSL compiler/bridge: {dependency_error}") + + try: + load_grouped_gemm_dswiglu_kernel() + except (ImportError, ModuleNotFoundError, RuntimeError) as exc: + errors.append(f"could not load the cuDNN frontend dSwiGLU backward kernel: {exc}") return errors @@ -466,10 +465,7 @@ def _ffn_fwd_per_shard( casted_wi = tex.grouped_quantize(wi_combined, fc1_quantizer_set.kernel, flatten_axis=-1) casted_intermediate = None if use_cudnn_cutedsl_fusion: - from .cutedsl_extensions.moe import ( - grouped_gemm_swiglu_mxfp8, - unpack_swiglu_pair, - ) + from .cutedsl_extensions.moe import unpack_swiglu_pair casted_sorted_x_lhs = casted_sorted_x.get_tensor(usage=TensorUsage.LHS) casted_wi_rhs = casted_wi.get_tensor(usage=TensorUsage.RHS) @@ -485,7 +481,7 @@ def _ffn_fwd_per_shard( intermediate_col, intermediate_scale_row, intermediate_scale_col, - ) = grouped_gemm_swiglu_mxfp8( + ) = tex.grouped_gemm_swiglu( casted_sorted_x_lhs.data.reshape(sorted_x.shape[0], hidden, 1), # TE stores the colwise payload in logical [E,K,N] order even # though its values were quantized along K. Materialize the @@ -1622,7 +1618,7 @@ def moe( ``NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION=1`` requires the call to match the eligible MXFP8 SwiGLU surface, uses the cuDNN frontend CuTeDSL FC1 fusion, and raises dispatch alignment to the kernel-required 256 tokens. - Ineligible opt-in calls fail with the reasons they cannot use CuTeDSL. + Ineligible opt-in calls warn and fall back to the regular grouped-GEMM path. Axis-name parameters: @@ -1668,11 +1664,15 @@ def moe( ep_axis=ep_axis, ) if rejection_reasons: - raise ValueError( - f"{_CUDNN_CUTEDSL_ENV}=1 is unsupported for this moe() call:\n- " - + "\n- ".join(rejection_reasons) + warnings.warn( + f"{_CUDNN_CUTEDSL_ENV}=1 is unsupported for this moe() call; " + "falling back to the regular grouped-GEMM path:\n- " + + "\n- ".join(rejection_reasons), + RuntimeWarning, + stacklevel=2, ) - use_cudnn_cutedsl_fusion = True + else: + use_cudnn_cutedsl_fusion = True # Enforce ((outer_dp..., ep), None, None) on inbound activations. The # EP comm groups consecutive global ranks (dp_color = rank // ep_size), From f7abdef6eb78d69ed523eebdb06e34c25988664a Mon Sep 17 00:00:00 2001 From: tdophung Date: Fri, 31 Jul 2026 16:00:37 -0700 Subject: [PATCH 04/38] Use TVM-FFI for fused dSwiGLU backward --- docs/envvars.rst | 9 +- docs/examples/jax/cutedsl_moe_pipeclean.rst | 19 +- tests/jax/cutedsl_smoke.py | 86 --- ...work-neutral-grouped-SwiGLU-compiler.patch | 575 +++++++++++++++-- tests/jax/run_tvm_ffi_fused_moe_e2e.sh | 16 +- tests/jax/test_cutedsl_moe.py | 48 +- tests/jax/tvm_ffi_grouped_mlp_multiprocess.py | 3 +- .../jax/cpp_extensions/grouped_gemm_swiglu.py | 365 ++++++++++- .../jax/cutedsl_extensions/__init__.py | 4 - .../jax/cutedsl_extensions/moe.py | 576 ------------------ transformer_engine/jax/moe.py | 48 +- 11 files changed, 934 insertions(+), 815 deletions(-) delete mode 100644 tests/jax/cutedsl_smoke.py delete mode 100644 transformer_engine/jax/cutedsl_extensions/__init__.py delete mode 100644 transformer_engine/jax/cutedsl_extensions/moe.py diff --git a/docs/envvars.rst b/docs/envvars.rst index a14c9290b51..2ab4b3e82f4 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -503,10 +503,11 @@ JAX-Specific Variables :Type: ``int`` (0 or 1) :Default: ``0`` :Description: **(JAX only)** Enable the experimental cuDNN frontend CuTeDSL fusion - for MXFP8 MoE FC1 grouped GEMM, SwiGLU, and grouped quantization. Explicit - opt-in requires an eligible SM100 SwiGLU MXFP8 MoE call and the optional - CUTLASS/cuDNN frontend JAX runtime packages; unsupported calls fail with the - full validation reason list. + for MXFP8 MoE FC1 grouped GEMM, SwiGLU/dSwiGLU, and grouped quantization. + Forward and backward launch as native TVM-FFI functions through ``jax-tvm-ffi``. + Explicit opt-in requires an eligible SM100 SwiGLU MXFP8 MoE call and the + optional compiler/runtime packages; unsupported calls warn with the full + validation reason list and use the unfused path. JAX Triton Extensions ^^^^^^^^^^^^^^^^^^^^^ diff --git a/docs/examples/jax/cutedsl_moe_pipeclean.rst b/docs/examples/jax/cutedsl_moe_pipeclean.rst index 8df5ddec841..276ae281b21 100644 --- a/docs/examples/jax/cutedsl_moe_pipeclean.rst +++ b/docs/examples/jax/cutedsl_moe_pipeclean.rst @@ -4,9 +4,9 @@ CuTeDSL MXFP8 MoE pipeclean The JAX MoE path has an opt-in Blackwell fusion for FC1 grouped GEMM, SwiGLU, and rowwise/colwise MXFP8 quantization. The forward leaf is compiled from abstract tensor descriptors by cuDNN-FE and invoked as a native -``tvm_ffi.Function``. The surrounding EP path continues to use ``shard_map`` -and sees only shard-local tensors. The fused dSwiGLU backward leaf continues -to use CUTLASS DSL's JAX integration. +``tvm_ffi.Function``. The fused dSwiGLU backward leaf uses the same compile-only +cuDNN-FE and ``jax-tvm-ffi`` path. The surrounding EP path continues to use +``shard_map`` and sees only shard-local tensors. Environment baseline -------------------- @@ -39,16 +39,19 @@ shapes and biases, MXFP8 quantizers, JAX FFI availability, and both compiler paths before lowering. Unsupported configurations emit one warning with the complete validation list and use the unfused implementation. -The cuDNN-FE compile-only API is exported from:: +The cuDNN-FE compile-only APIs are exported from:: cudnn/grouped_gemm/grouped_gemm_swiglu/compile.py + cudnn/grouped_gemm/grouped_gemm_dswiglu/compile.py It accepts self-describing operand metadata, compiles without Torch or live buffers, and returns a native function plus an exact ABI descriptor. The -native launch wrapper reorders the eight arguments, five results, and stream -into the volatile raw kernel ABI. This is necessary because ``jax-tvm-ffi`` -currently describes only complete argument and result groups. The fused path -uses 256-token dispatch alignment; the unfused path remains at 128. +native launch wrappers reorder each operation's arguments, five results, and +stream into the volatile raw kernel ABIs. This is necessary because +``jax-tvm-ffi`` currently describes only complete argument and result groups. +Neither compile entrypoint imports the Torch wrapper APIs, and TE contains no +Torch stub or ``cutlass_call`` launcher. The fused path uses 256-token dispatch +alignment; the unfused path remains at 128. Custom partitioning migration design ------------------------------------ diff --git a/tests/jax/cutedsl_smoke.py b/tests/jax/cutedsl_smoke.py deleted file mode 100644 index e36806ede98..00000000000 --- a/tests/jax/cutedsl_smoke.py +++ /dev/null @@ -1,86 +0,0 @@ -# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# -# See LICENSE for license information. -"""Standalone CuTeDSL/JAX toolchain smoke test. - -Run this before the TE tests. It verifies CUTLASS FFI runtime discovery and -executes a trivial CuTeDSL kernel from inside ``jax.jit``. -""" - -from importlib.metadata import version - -import cuda.bindings.driver as cuda -import cutlass.cute as cute -import cutlass.jax as cjax -import jax -import jax.numpy as jnp -import numpy as np - - -BLOCK = 256 - - -@cute.kernel -def _vector_add_kernel(a: cute.Tensor, b: cute.Tensor, c: cute.Tensor): - tidx, _, _ = cute.arch.thread_idx() - bidx, _, _ = cute.arch.block_idx() - - frag_a = cute.make_rmem_tensor(cute.size(a, mode=[0]), a.element_type) - frag_b = cute.make_rmem_tensor(cute.size(b, mode=[0]), b.element_type) - frag_c = cute.make_rmem_tensor(cute.size(c, mode=[0]), c.element_type) - cute.autovec_copy(a[None, tidx, bidx], frag_a) - cute.autovec_copy(b[None, tidx, bidx], frag_b) - frag_c.store(frag_a.load() + frag_b.load()) - cute.autovec_copy(frag_c, c[None, tidx, bidx]) - - -@cute.jit -def _launch_vector_add( - stream: cuda.CUstream, - a: cute.Tensor, - b: cute.Tensor, - c: cute.Tensor, -): - _vector_add_kernel(a, b, c).launch( - grid=[a.shape[-1], 1, 1], - block=[a.shape[-2], 1, 1], - stream=stream, - ) - - -@jax.jit -def _cutlass_add(a, b): - size = a.shape[0] - padded_size = ((size + BLOCK - 1) // BLOCK) * BLOCK - a_3d = jnp.pad(a, (0, padded_size - size)).reshape(1, BLOCK, -1) - b_3d = jnp.pad(b, (0, padded_size - size)).reshape(1, BLOCK, -1) - call = cjax.cutlass_call( - _launch_vector_add, - output_shape_dtype=jax.ShapeDtypeStruct.like(a_3d), - use_static_tensors=True, - ) - return call(a_3d, b_3d).reshape(-1)[:size] - - -def main() -> None: - if not cjax.is_available(): - raise RuntimeError( - "cutlass.jax.is_available() is false; verify cute_dsl_runtime.so discovery" - ) - devices = jax.devices("gpu") - if not devices: - raise RuntimeError("No JAX GPU device is available") - - a = jnp.arange(1024, dtype=jnp.float32) - b = jnp.arange(1024, dtype=jnp.float32) * 2 - actual = _cutlass_add(a, b) - actual.block_until_ready() - np.testing.assert_array_equal(np.asarray(actual), np.asarray(a + b)) - print( - "CuTeDSL JAX smoke passed: " - f"jax={jax.__version__}, cutlass={version('nvidia-cutlass-dsl')}, device={devices[0]}" - ) - - -if __name__ == "__main__": - main() diff --git a/tests/jax/patches/0001-Add-framework-neutral-grouped-SwiGLU-compiler.patch b/tests/jax/patches/0001-Add-framework-neutral-grouped-SwiGLU-compiler.patch index d26ca1e4854..1fec3cc29d9 100644 --- a/tests/jax/patches/0001-Add-framework-neutral-grouped-SwiGLU-compiler.patch +++ b/tests/jax/patches/0001-Add-framework-neutral-grouped-SwiGLU-compiler.patch @@ -1,26 +1,8 @@ -From fe78f4c3f3050b537362ebca83b734adfbff1093 Mon Sep 17 00:00:00 2001 -From: tdophung -Date: Wed, 22 Jul 2026 12:15:31 -0700 -Subject: [PATCH] Add framework-neutral grouped SwiGLU compiler - ---- - .../gemm_fusions/grouped_gemm_swiglu.md | 41 ++ - pyproject.toml | 5 + - python/cudnn/__init__.py | 6 + - python/cudnn/grouped_gemm/__init__.py | 114 +++--- - .../grouped_gemm_swiglu/__init__.py | 32 +- - .../grouped_gemm_swiglu/compile.py | 366 ++++++++++++++++++ - python/cudnn/grouped_gemm/utils.py | 10 +- - .../test_grouped_gemm_swiglu_compile.py | 96 +++++ - 8 files changed, 597 insertions(+), 73 deletions(-) - create mode 100644 python/cudnn/grouped_gemm/grouped_gemm_swiglu/compile.py - create mode 100644 test/python/fe_api/grouped_gemm/test_grouped_gemm_swiglu_compile.py - diff --git a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_swiglu.md b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_swiglu.md -index 1ac83b7..bf73160 100644 +index 1ac83b7..359251d 100644 --- a/docs/fe-oss-apis/gemm_fusions/grouped_gemm_swiglu.md +++ b/docs/fe-oss-apis/gemm_fusions/grouped_gemm_swiglu.md -@@ -96,6 +96,47 @@ $$ +@@ -96,6 +96,64 @@ $$ ## API Usage @@ -64,9 +46,26 @@ index 1ac83b7..bf73160 100644 +CuTeDSL reorder launcher hides the volatile raw-kernel operand order, so no +Python callback runs during execution. The `amax` slot is disabled in this +five-output prototype. ++ ++The paired backward entrypoint follows the same framework-neutral contract: ++ ++```python ++from cudnn import GroupedGemmDswigluConfig, compile_grouped_gemm_dswiglu ++ ++native_backward, backward_abi = compile_grouped_gemm_dswiglu( ++ operands=backward_operands, ++ config=GroupedGemmDswigluConfig(), ++) ++``` ++ ++It groups ten inputs (`a`, `b`, the saved forward projection `c`, scale ++factors, grouped offsets, alpha/beta, probability, and normalization), five ++outputs (`d`, `d_col`, both output scale factors, and `dprob`), and the CUDA ++stream. Importing either compile-only entrypoint does not import the ++Torch-oriented `api.py` modules. + ### High-level Wrapper - + ```python diff --git a/pyproject.toml b/pyproject.toml index 544b28d..c68805d 100644 @@ -81,14 +80,14 @@ index 544b28d..c68805d 100644 + "cuda-python", + "apache-tvm-ffi>=0.1.12", +] - + [dependency-groups] dev = [ diff --git a/python/cudnn/__init__.py b/python/cudnn/__init__.py -index 6e8b998..e4493a4 100644 +index 6e8b998..15040b3 100644 --- a/python/cudnn/__init__.py +++ b/python/cudnn/__init__.py -@@ -298,6 +298,12 @@ _LAZY_OPTIONAL_IMPORTS = { +@@ -298,8 +298,16 @@ _LAZY_OPTIONAL_IMPORTS = { "grouped_gemm": (".grouped_gemm", None), "GroupedGemmSwigluSm100": (".grouped_gemm", "GroupedGemmSwigluSm100"), "grouped_gemm_swiglu_wrapper_sm100": (".grouped_gemm", "grouped_gemm_swiglu_wrapper_sm100"), @@ -100,16 +99,20 @@ index 6e8b998..e4493a4 100644 + "compile_grouped_gemm_swiglu": (".grouped_gemm", "compile_grouped_gemm_swiglu"), "GroupedGemmDswigluSm100": (".grouped_gemm", "GroupedGemmDswigluSm100"), "grouped_gemm_dswiglu_wrapper_sm100": (".grouped_gemm", "grouped_gemm_dswiglu_wrapper_sm100"), ++ "GroupedGemmDswigluConfig": (".grouped_gemm", "GroupedGemmDswigluConfig"), ++ "compile_grouped_gemm_dswiglu": (".grouped_gemm", "compile_grouped_gemm_dswiglu"), "GroupedGemmSreluSm100": (".grouped_gemm", "GroupedGemmSreluSm100"), + "grouped_gemm_srelu_wrapper_sm100": (".grouped_gemm", "grouped_gemm_srelu_wrapper_sm100"), + "GroupedGemmDsreluSm100": (".grouped_gemm", "GroupedGemmDsreluSm100"), diff --git a/python/cudnn/grouped_gemm/__init__.py b/python/cudnn/grouped_gemm/__init__.py -index 818a1ab..0f69996 100644 +index 818a1ab..98b9883 100644 --- a/python/cudnn/grouped_gemm/__init__.py +++ b/python/cudnn/grouped_gemm/__init__.py -@@ -1,68 +1,50 @@ +@@ -1,68 +1,52 @@ -# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: MIT - + -from .grouped_gemm_swiglu.api import ( - GroupedGemmSwigluSm100, - grouped_gemm_swiglu_wrapper_sm100, @@ -195,6 +198,8 @@ index 818a1ab..0f69996 100644 + "compile_grouped_gemm_swiglu": (".grouped_gemm_swiglu", "compile_grouped_gemm_swiglu"), + "GroupedGemmDswigluSm100": (".grouped_gemm_dswiglu.api", "GroupedGemmDswigluSm100"), + "grouped_gemm_dswiglu_wrapper_sm100": (".grouped_gemm_dswiglu.api", "grouped_gemm_dswiglu_wrapper_sm100"), ++ "GroupedGemmDswigluConfig": (".grouped_gemm_dswiglu", "GroupedGemmDswigluConfig"), ++ "compile_grouped_gemm_dswiglu": (".grouped_gemm_dswiglu", "compile_grouped_gemm_dswiglu"), + "GroupedGemmQuantSm100": (".grouped_gemm_quant.api", "GroupedGemmQuantSm100"), + "grouped_gemm_quant_wrapper_sm100": (".grouped_gemm_quant.api", "grouped_gemm_quant_wrapper_sm100"), + "GroupedGemmSreluSm100": (".grouped_gemm_srelu.api", "GroupedGemmSreluSm100"), @@ -222,6 +227,349 @@ index 818a1ab..0f69996 100644 + + +__all__ = list(_LAZY_IMPORTS) +diff --git a/python/cudnn/grouped_gemm/grouped_gemm_dswiglu/__init__.py b/python/cudnn/grouped_gemm/grouped_gemm_dswiglu/__init__.py +index f026195..335c26b 100644 +--- a/python/cudnn/grouped_gemm/grouped_gemm_dswiglu/__init__.py ++++ b/python/cudnn/grouped_gemm/grouped_gemm_dswiglu/__init__.py +@@ -1,12 +1,28 @@ + # Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + # SPDX-License-Identifier: MIT + +-from .api import ( +- GroupedGemmDswigluSm100, +- grouped_gemm_dswiglu_wrapper_sm100, +-) ++from importlib import import_module ++ ++ ++_LAZY_IMPORTS = { ++ "GroupedGemmDswigluSm100": (".api", "GroupedGemmDswigluSm100"), ++ "grouped_gemm_dswiglu_wrapper_sm100": (".api", "grouped_gemm_dswiglu_wrapper_sm100"), ++ "GroupedGemmDswigluConfig": (".compile", "GroupedGemmDswigluConfig"), ++ "compile_grouped_gemm_dswiglu": (".compile", "compile_grouped_gemm_dswiglu"), ++} ++ ++ ++def __getattr__(name): ++ if name not in _LAZY_IMPORTS: ++ raise AttributeError(name) ++ module_name, attr_name = _LAZY_IMPORTS[name] ++ value = getattr(import_module(module_name, package=__name__), attr_name) ++ globals()[name] = value ++ return value + + __all__ = [ + "GroupedGemmDswigluSm100", + "grouped_gemm_dswiglu_wrapper_sm100", ++ "GroupedGemmDswigluConfig", ++ "compile_grouped_gemm_dswiglu", + ] +diff --git a/python/cudnn/grouped_gemm/grouped_gemm_dswiglu/compile.py b/python/cudnn/grouped_gemm/grouped_gemm_dswiglu/compile.py +new file mode 100644 +index 0000000..46fa3d0 +--- /dev/null ++++ b/python/cudnn/grouped_gemm/grouped_gemm_dswiglu/compile.py +@@ -0,0 +1,300 @@ ++# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. ++# SPDX-License-Identifier: MIT ++ ++"""Framework-neutral compile-only API for grouped GEMM + dSwiGLU. ++ ++This module consumes tensor metadata, compiles a native TVM-FFI function, ++and returns a self-describing ABI. It has no PyTorch dependency, allocates no ++framework buffers, and never launches the kernel. ++""" ++ ++from __future__ import annotations ++ ++from dataclasses import dataclass ++import hashlib ++import json ++import threading ++from typing import Any, Mapping ++ ++from cuda.bindings import driver as cuda ++import cutlass ++import cutlass.cute as cute ++from cutlass.cute.runtime import make_fake_stream ++ ++from ..grouped_gemm_swiglu.compile import ( ++ GroupedGemmAbiDescriptor, ++ GroupedGemmAbiEntry, ++ GroupedGemmElementType, ++ GroupedGemmOperandDesc, ++ _cutlass_dtype, ++ _make_fake, ++ _normalize_dtype, ++ _scale_size, ++) ++from .grouped_gemm_dswiglu_quant import BlockScaledContiguousGroupedGemmKernel ++ ++ ++@dataclass(frozen=True) ++class GroupedGemmDswigluConfig: ++ """Compile-time configuration for the SM100 grouped dSwiGLU kernel.""" ++ ++ sf_vec_size: int = 32 ++ acc_dtype: GroupedGemmElementType | str = GroupedGemmElementType.FLOAT32 ++ mma_tiler_mn: tuple[int, int] = (256, 256) ++ cluster_shape_mn: tuple[int, int] = (2, 1) ++ vector_f32: bool = False ++ discrete_col_sfd: bool = True ++ use_mono_increase_expert_idx: bool = True ++ num_cluster_overlap_margin: int = 0 ++ ++ def __post_init__(self): ++ object.__setattr__(self, "acc_dtype", _normalize_dtype(self.acc_dtype)) ++ object.__setattr__(self, "mma_tiler_mn", tuple(self.mma_tiler_mn)) ++ object.__setattr__(self, "cluster_shape_mn", tuple(self.cluster_shape_mn)) ++ if self.sf_vec_size not in (16, 32): ++ raise ValueError(f"sf_vec_size must be 16 or 32, got {self.sf_vec_size}") ++ if self.acc_dtype != GroupedGemmElementType.FLOAT32: ++ raise ValueError("Grouped GEMM dSwiGLU requires a float32 accumulator") ++ ++ ++_INPUT_ROLES = ( ++ "a", ++ "b", ++ "c", ++ "sfa", ++ "sfb", ++ "padded_offsets", ++ "alpha", ++ "beta", ++ "prob", ++ "norm_const", ++) ++_OUTPUT_ROLES = ("d", "d_col", "sfd_row", "sfd_col", "dprob") ++_ALL_ROLES = _INPUT_ROLES + _OUTPUT_ROLES ++_COMPILE_CACHE: dict[str, tuple[Any, GroupedGemmAbiDescriptor]] = {} ++_COMPILE_LOCK = threading.Lock() ++ ++ ++def _expect_shape(operand: GroupedGemmOperandDesc, shape: tuple[int, ...]) -> None: ++ if operand.shape != shape: ++ raise ValueError(f"{operand.role}: expected shape {shape}, got {operand.shape}") ++ ++ ++def _validate_operands( ++ operands: Mapping[str, GroupedGemmOperandDesc], ++ config: GroupedGemmDswigluConfig, ++) -> None: ++ missing = sorted(set(_ALL_ROLES) - set(operands)) ++ extra = sorted(set(operands) - set(_ALL_ROLES)) ++ if missing or extra: ++ raise ValueError(f"Invalid operand roles: missing={missing}, extra={extra}") ++ for role, operand in operands.items(): ++ if operand.role != role: ++ raise ValueError( ++ f"Operand mapping key {role!r} does not match descriptor role {operand.role!r}" ++ ) ++ ++ a, b, c = operands["a"], operands["b"], operands["c"] ++ if len(a.shape) != 3 or a.shape[2] != 1: ++ raise ValueError(f"a: expected [M,K,1], got {a.shape}") ++ if len(b.shape) != 3: ++ raise ValueError(f"b: expected [N,K,E], got {b.shape}") ++ m, k, _ = a.shape ++ n, b_k, experts = b.shape ++ if k != b_k: ++ raise ValueError(f"A/B K mismatch: {k} != {b_k}") ++ if m % BlockScaledContiguousGroupedGemmKernel.FIX_PAD_SIZE: ++ raise ValueError( ++ f"M must be {BlockScaledContiguousGroupedGemmKernel.FIX_PAD_SIZE}-aligned, got {m}" ++ ) ++ if experts > 1024: ++ raise ValueError(f"Expert count must be <= 1024, got {experts}") ++ ++ _expect_shape(c, (m, 2 * n, 1)) ++ _expect_shape(operands["d"], (m, 2 * n, 1)) ++ _expect_shape(operands["d_col"], (m, 2 * n, 1)) ++ _expect_shape(operands["padded_offsets"], (experts,)) ++ _expect_shape(operands["alpha"], (experts,)) ++ _expect_shape(operands["beta"], (experts,)) ++ _expect_shape(operands["prob"], (m, 1, 1)) ++ _expect_shape(operands["norm_const"], (1,)) ++ _expect_shape(operands["dprob"], (m, 1, 1)) ++ ++ expected_sfa = _scale_size(m, k, config.sf_vec_size, colwise=False) ++ expected_sfb = _scale_size(n, k, config.sf_vec_size, colwise=False) * experts ++ if len(operands["sfa"].shape) != 1 or operands["sfa"].shape[0] < expected_sfa: ++ raise ValueError(f"sfa: expected a flat buffer with at least {expected_sfa} values") ++ if len(operands["sfb"].shape) != 1 or operands["sfb"].shape[0] < expected_sfb: ++ raise ValueError(f"sfb: expected a flat buffer with at least {expected_sfb} values") ++ _expect_shape( ++ operands["sfd_row"], ++ (_scale_size(m, 2 * n, config.sf_vec_size, colwise=False),), ++ ) ++ _expect_shape( ++ operands["sfd_col"], ++ (_scale_size(m, 2 * n, config.sf_vec_size, colwise=True),), ++ ) ++ ++ if a.dtype != b.dtype or a.dtype not in ( ++ GroupedGemmElementType.FLOAT8_E4M3FN, ++ GroupedGemmElementType.FLOAT8_E5M2, ++ ): ++ raise ValueError(f"A/B must have one matching FP8 dtype, got {a.dtype}/{b.dtype}") ++ if c.dtype not in (GroupedGemmElementType.BFLOAT16, GroupedGemmElementType.FLOAT16): ++ raise ValueError(f"c: expected bfloat16 or float16, got {c.dtype.value}") ++ if operands["d"].dtype != operands["d_col"].dtype: ++ raise ValueError("d and d_col must have matching dtypes") ++ if operands["d"].dtype not in ( ++ GroupedGemmElementType.FLOAT8_E4M3FN, ++ GroupedGemmElementType.FLOAT8_E5M2, ++ ): ++ raise ValueError(f"d/d_col must have an FP8 dtype, got {operands['d'].dtype.value}") ++ sf_dtype = GroupedGemmElementType.FLOAT8_E8M0FNU ++ for role in ("sfa", "sfb", "sfd_row", "sfd_col"): ++ if operands[role].dtype != sf_dtype: ++ raise ValueError(f"{role}: expected {sf_dtype.value}, got {operands[role].dtype.value}") ++ for role in ("alpha", "beta", "prob", "norm_const", "dprob"): ++ if operands[role].dtype != GroupedGemmElementType.FLOAT32: ++ raise ValueError(f"{role}: expected float32") ++ if operands["padded_offsets"].dtype != GroupedGemmElementType.INT32: ++ raise ValueError("padded_offsets: expected int32") ++ ++ ++def _cache_key( ++ operands: Mapping[str, GroupedGemmOperandDesc], ++ config: GroupedGemmDswigluConfig, ++) -> str: ++ payload = { ++ "operands": [ ++ { ++ "role": role, ++ "dtype": operands[role].dtype.value, ++ "shape": operands[role].shape, ++ "stride_order": operands[role].stride_order, ++ "align": operands[role].align, ++ } ++ for role in _ALL_ROLES ++ ], ++ "config": { ++ "sf_vec_size": config.sf_vec_size, ++ "acc_dtype": config.acc_dtype.value, ++ "mma_tiler_mn": config.mma_tiler_mn, ++ "cluster_shape_mn": config.cluster_shape_mn, ++ "vector_f32": config.vector_f32, ++ "discrete_col_sfd": config.discrete_col_sfd, ++ "use_mono_increase_expert_idx": config.use_mono_increase_expert_idx, ++ "num_cluster_overlap_margin": config.num_cluster_overlap_margin, ++ }, ++ } ++ digest = hashlib.sha256(json.dumps(payload, sort_keys=True).encode()).hexdigest()[:24] ++ return f"cudnn.grouped_gemm_dswiglu.{digest}" ++ ++ ++def _compile_native( ++ operands: Mapping[str, GroupedGemmOperandDesc], ++ config: GroupedGemmDswigluConfig, ++ key: str, ++) -> tuple[Any, GroupedGemmAbiDescriptor]: ++ kernel = BlockScaledContiguousGroupedGemmKernel( ++ sf_vec_size=config.sf_vec_size, ++ acc_dtype=_cutlass_dtype(config.acc_dtype), ++ use_2cta_instrs=config.mma_tiler_mn[0] == 256, ++ mma_tiler_mn=config.mma_tiler_mn, ++ cluster_shape_mn=config.cluster_shape_mn, ++ vectorized_f32=config.vector_f32, ++ discrete_col_sfd=config.discrete_col_sfd, ++ expert_cnt=operands["padded_offsets"].shape[0], ++ use_mono_increase_expert_idx=config.use_mono_increase_expert_idx, ++ ) ++ max_active_clusters = cutlass.utils.HardwareInfo().get_max_active_clusters( ++ config.cluster_shape_mn[0] * config.cluster_shape_mn[1] ++ ) ++ max_active_clusters -= config.num_cluster_overlap_margin ++ if max_active_clusters <= 0: ++ raise ValueError("num_cluster_overlap_margin leaves no active clusters") ++ ++ @cute.jit ++ def launch( ++ a, ++ b, ++ c, ++ sfa, ++ sfb, ++ padded_offsets, ++ alpha, ++ beta, ++ prob, ++ norm_const, ++ d, ++ d_col, ++ sfd_row, ++ sfd_col, ++ dprob, ++ stream: cuda.CUstream, ++ ): ++ kernel( ++ a, ++ b, ++ c, ++ d, ++ d_col, ++ sfa, ++ sfb, ++ sfd_row, ++ sfd_col, ++ None, ++ norm_const, ++ padded_offsets, ++ alpha, ++ beta, ++ prob, ++ dprob, ++ max_active_clusters, ++ stream, ++ ) ++ ++ fake_operands = [_make_fake(operands[role]) for role in _ALL_ROLES] ++ fake_stream = make_fake_stream(use_tvm_ffi_env_stream=False) ++ compiled = cute.compile(launch, *fake_operands, fake_stream, options="--enable-tvm-ffi") ++ entries = tuple( ++ GroupedGemmAbiEntry( ++ role, ++ "arg", ++ operands[role].dtype, ++ operands[role].shape, ++ operands[role].stride_order, ++ ) ++ for role in _INPUT_ROLES ++ ) + tuple( ++ GroupedGemmAbiEntry( ++ role, ++ "ret", ++ operands[role].dtype, ++ operands[role].shape, ++ operands[role].stride_order, ++ ) ++ for role in _OUTPUT_ROLES ++ ) + (GroupedGemmAbiEntry("stream", "stream", None),) ++ descriptor = GroupedGemmAbiDescriptor(entries=entries, key=key) ++ descriptor.arg_spec ++ return compiled, descriptor ++ ++ ++def compile_grouped_gemm_dswiglu( ++ *, ++ operands: Mapping[str, GroupedGemmOperandDesc], ++ config: GroupedGemmDswigluConfig | None = None, ++) -> tuple[Any, GroupedGemmAbiDescriptor]: ++ """Compile and cache a native grouped GEMM + dSwiGLU TVM-FFI function.""" ++ config = GroupedGemmDswigluConfig() if config is None else config ++ _validate_operands(operands, config) ++ key = _cache_key(operands, config) ++ with _COMPILE_LOCK: ++ cached = _COMPILE_CACHE.get(key) ++ if cached is None: ++ cached = _compile_native(operands, config, key) ++ _COMPILE_CACHE[key] = cached ++ return cached ++ ++ ++__all__ = ["GroupedGemmDswigluConfig", "compile_grouped_gemm_dswiglu"] diff --git a/python/cudnn/grouped_gemm/grouped_gemm_swiglu/__init__.py b/python/cudnn/grouped_gemm/grouped_gemm_swiglu/__init__.py index 6bcd935..77483cf 100644 --- a/python/cudnn/grouped_gemm/grouped_gemm_swiglu/__init__.py @@ -229,7 +577,7 @@ index 6bcd935..77483cf 100644 @@ -8,12 +8,36 @@ This module provides the forward grouped GEMM with SwiGLU activation for MoE (Mixture of Experts) workloads on SM100+ GPUs. """ - + -from .api import ( - GroupedGemmSwigluSm100, - grouped_gemm_swiglu_wrapper_sm100, @@ -256,7 +604,7 @@ index 6bcd935..77483cf 100644 + value = getattr(import_module(module_name, package=__name__), attr_name) + globals()[name] = value + return value - + __all__ = [ "GroupedGemmSwigluSm100", "grouped_gemm_swiglu_wrapper_sm100", @@ -272,7 +620,7 @@ new file mode 100644 index 0000000..04f5d30 --- /dev/null +++ b/python/cudnn/grouped_gemm/grouped_gemm_swiglu/compile.py -@@ -0,0 +1,391 @@ +@@ -0,0 +1,366 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: MIT + @@ -569,40 +917,15 @@ index 0000000..04f5d30 + if max_active_clusters <= 0: + raise ValueError("num_cluster_overlap_margin leaves no active clusters") + -+ # XLA's FFI buffer ABI does not carry strides, and jax-tvm-ffi therefore -+ # exposes every buffer to TVM FFI as compact row-major. B is logically -+ # [N,K,E] with K minor, which cannot be represented by a compact -+ # [N,K,E] buffer when E > 1. Make the public ABI's B slot the equivalent -+ # compact [E,N,K] storage and restore the logical CuTe view in the native -+ # launcher. This is a view only; no framework-specific code or copy is -+ # introduced on the execution path. -+ b_desc = operands["b"] -+ b_n, b_k, b_experts = b_desc.shape -+ ffi_operands = dict(operands) -+ ffi_operands["b"] = GroupedGemmOperandDesc( -+ role="b", -+ dtype=b_desc.dtype, -+ shape=(b_experts, b_n, b_k), -+ stride_order=(2, 1, 0), -+ align=b_desc.align, -+ ) -+ + # jax-tvm-ffi routes complete input and result groups. Compile a native + # reorder launcher with that stable grouping while keeping the volatile + # raw-kernel ABI private to cuDNN-FE. This remains native code: no Python + # callback is present on the execution path. + @cute.jit + def launch(a, b, sfa, sfb, padded_offsets, alpha, prob, norm_const, c, d, d_col, sfd_row, sfd_col, stream: cuda.CUstream): -+ b_logical = cute.make_tensor( -+ b.iterator, -+ cute.make_layout( -+ (b_n, b_k, b_experts), -+ stride=(b_k, 1, b_n * b_k), -+ ), -+ ) + kernel( + a, -+ b_logical, ++ b, + c, + d, + d_col, @@ -619,14 +942,14 @@ index 0000000..04f5d30 + stream, + ) + -+ fake_operands = [_make_fake(ffi_operands[role]) for role in _ALL_ROLES] ++ fake_operands = [_make_fake(operands[role]) for role in _ALL_ROLES] + fake_stream = make_fake_stream(use_tvm_ffi_env_stream=False) + compiled = cute.compile(launch, *fake_operands, fake_stream, options="--enable-tvm-ffi") + entries = tuple( -+ GroupedGemmAbiEntry(role, "arg", ffi_operands[role].dtype, ffi_operands[role].shape, ffi_operands[role].stride_order) ++ GroupedGemmAbiEntry(role, "arg", operands[role].dtype, operands[role].shape, operands[role].stride_order) + for role in _INPUT_ROLES + ) + tuple( -+ GroupedGemmAbiEntry(role, "ret", ffi_operands[role].dtype, ffi_operands[role].shape, ffi_operands[role].stride_order) ++ GroupedGemmAbiEntry(role, "ret", operands[role].dtype, operands[role].shape, operands[role].stride_order) + for role in _OUTPUT_ROLES + ) + (GroupedGemmAbiEntry("stream", "stream", None),) + descriptor = GroupedGemmAbiDescriptor(entries=entries, key=key) @@ -671,10 +994,10 @@ index cb60dc8..1287293 100644 @@ -33,7 +33,7 @@ This module contains the tile scheduler classes and helper functions used by bot the forward (grouped_gemm_swiglu) and backward (grouped_gemm_dswiglu) kernels. """ - + -from typing import Tuple, Union +from typing import TYPE_CHECKING, Tuple, Union - + from cutlass.cutlass_dsl import ( Boolean, @@ -52,7 +52,6 @@ from cutlass.cutlass_dsl import T @@ -688,7 +1011,7 @@ index cb60dc8..1287293 100644 @@ -61,6 +60,9 @@ from cutlass.pipeline import ( make_pipeline_state, ) - + +if TYPE_CHECKING: + import torch + @@ -697,8 +1020,8 @@ index cb60dc8..1287293 100644 ############################################################################## @@ -201,8 +203,10 @@ def fmax(a: Union[float, Float32], b: Union[float, Float32], *, loc=None, ip=Non ) - - + + -def logical_shape_fp4x2_aware(tensor: torch.Tensor) -> Tuple[int, ...]: +def logical_shape_fp4x2_aware(tensor: "torch.Tensor") -> Tuple[int, ...]: """Return correct shapes for NVFP4 tensor.""" @@ -707,6 +1030,128 @@ index cb60dc8..1287293 100644 if tensor.dtype == torch.float4_e2m1fn_x2: innermost_dim_index = next((i for i, s in enumerate(tensor.stride()) if s == 1), None) if innermost_dim_index is None: +diff --git a/test/python/fe_api/grouped_gemm/test_grouped_gemm_dswiglu_compile.py b/test/python/fe_api/grouped_gemm/test_grouped_gemm_dswiglu_compile.py +new file mode 100644 +index 0000000..d3afc1e +--- /dev/null ++++ b/test/python/fe_api/grouped_gemm/test_grouped_gemm_dswiglu_compile.py +@@ -0,0 +1,116 @@ ++# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. ++# SPDX-License-Identifier: MIT ++ ++"""Tests for the framework-neutral grouped dSwiGLU compile API.""" ++ ++import subprocess ++import sys ++ ++import pytest ++import torch ++ ++ ++def _operands(m=256, k=128, n=128, experts=1): ++ from cudnn import GroupedGemmElementType as DType ++ from cudnn import GroupedGemmOperandDesc as Operand ++ ++ def scale_size(rows, cols, colwise=False): ++ if colwise: ++ rows, cols = cols, rows ++ ceil_div = lambda x, y: (x + y - 1) // y ++ return 32 * 4 * ceil_div(rows, 128) * 4 * ceil_div(ceil_div(cols, 32), 4) ++ ++ return { ++ "a": Operand("a", DType.FLOAT8_E4M3FN, (m, k, 1), (1, 0, 2)), ++ "b": Operand("b", DType.FLOAT8_E4M3FN, (n, k, experts), (1, 0, 2)), ++ "c": Operand("c", DType.BFLOAT16, (m, 2 * n, 1), (1, 0, 2)), ++ "sfa": Operand("sfa", DType.FLOAT8_E8M0FNU, (scale_size(m, k),), (0,)), ++ "sfb": Operand( ++ "sfb", ++ DType.FLOAT8_E8M0FNU, ++ (scale_size(n, k) * experts,), ++ (0,), ++ ), ++ "padded_offsets": Operand("padded_offsets", DType.INT32, (experts,), (0,)), ++ "alpha": Operand("alpha", DType.FLOAT32, (experts,), (0,)), ++ "beta": Operand("beta", DType.FLOAT32, (experts,), (0,)), ++ "prob": Operand("prob", DType.FLOAT32, (m, 1, 1), (1, 0, 2)), ++ "norm_const": Operand("norm_const", DType.FLOAT32, (1,), (0,)), ++ "d": Operand("d", DType.FLOAT8_E4M3FN, (m, 2 * n, 1), (1, 0, 2)), ++ "d_col": Operand("d_col", DType.FLOAT8_E4M3FN, (m, 2 * n, 1), (1, 0, 2)), ++ "sfd_row": Operand( ++ "sfd_row", ++ DType.FLOAT8_E8M0FNU, ++ (scale_size(m, 2 * n),), ++ (0,), ++ ), ++ "sfd_col": Operand( ++ "sfd_col", ++ DType.FLOAT8_E8M0FNU, ++ (scale_size(m, 2 * n, colwise=True),), ++ (0,), ++ ), ++ "dprob": Operand("dprob", DType.FLOAT32, (m, 1, 1), (1, 0, 2)), ++ } ++ ++ ++@pytest.mark.L0 ++def test_compile_api_imports_without_torch(): ++ script = r""" ++import importlib.abc ++import sys ++ ++class BlockTorch(importlib.abc.MetaPathFinder): ++ def find_spec(self, fullname, path=None, target=None): ++ if fullname == "torch" or fullname.startswith("torch."): ++ raise ModuleNotFoundError("torch intentionally blocked") ++ return None ++ ++sys.meta_path.insert(0, BlockTorch()) ++from cudnn import compile_grouped_gemm_dswiglu ++assert callable(compile_grouped_gemm_dswiglu) ++assert "torch" not in sys.modules ++""" ++ subprocess.run([sys.executable, "-c", script], check=True) ++ ++ ++@pytest.mark.L0 ++def test_compile_api_rejects_unaligned_m_before_compile(): ++ from cudnn import compile_grouped_gemm_dswiglu ++ ++ with pytest.raises(ValueError, match="M must be 256-aligned"): ++ compile_grouped_gemm_dswiglu(operands=_operands(m=128)) ++ ++ ++@pytest.mark.L1 ++def test_compile_api_returns_native_function_and_self_describing_abi(): ++ if torch.cuda.get_device_capability()[0] < 10: ++ pytest.skip("Grouped GEMM dSwiGLU compile requires SM100+") ++ tvm_ffi = pytest.importorskip("tvm_ffi") ++ from cudnn import compile_grouped_gemm_dswiglu ++ ++ function, abi = compile_grouped_gemm_dswiglu(operands=_operands()) ++ cached_function, cached_abi = compile_grouped_gemm_dswiglu(operands=_operands()) ++ ++ assert isinstance(function, tvm_ffi.Function) ++ assert cached_function is function ++ assert cached_abi == abi ++ assert abi.arg_spec == ("args", "rets", "ctx.stream") ++ assert [(entry.role, entry.kind) for entry in abi.entries] == [ ++ ("a", "arg"), ++ ("b", "arg"), ++ ("c", "arg"), ++ ("sfa", "arg"), ++ ("sfb", "arg"), ++ ("padded_offsets", "arg"), ++ ("alpha", "arg"), ++ ("beta", "arg"), ++ ("prob", "arg"), ++ ("norm_const", "arg"), ++ ("d", "ret"), ++ ("d_col", "ret"), ++ ("sfd_row", "ret"), ++ ("sfd_col", "ret"), ++ ("dprob", "ret"), ++ ("stream", "stream"), ++ ] diff --git a/test/python/fe_api/grouped_gemm/test_grouped_gemm_swiglu_compile.py b/test/python/fe_api/grouped_gemm/test_grouped_gemm_swiglu_compile.py new file mode 100644 index 0000000..82e1eee @@ -809,5 +1254,3 @@ index 0000000..82e1eee + ("sfd_col", "ret"), + ("stream", "stream"), + ] --- -2.50.0 diff --git a/tests/jax/run_tvm_ffi_fused_moe_e2e.sh b/tests/jax/run_tvm_ffi_fused_moe_e2e.sh index 48c6000176c..dc815c82e93 100755 --- a/tests/jax/run_tvm_ffi_fused_moe_e2e.sh +++ b/tests/jax/run_tvm_ffi_fused_moe_e2e.sh @@ -4,7 +4,7 @@ # See LICENSE for license information. # End-to-end remote SM100 validation for the cuDNN-FE compile-only API, -# JAX TVM-FFI forward primitive, fused backward, partitioning, and MoEBlock. +# JAX TVM-FFI forward/backward primitives, partitioning, and MoEBlock. set -euo pipefail @@ -28,14 +28,19 @@ if [ ! -f "$CUDNN_FE_ROOT/pyproject.toml" ]; then git -C "$CUDNN_FE_ROOT" checkout --detach "$CUDNN_FE_BASE_REV" fi -CUDNN_FE_COMPILE_API="$CUDNN_FE_ROOT/python/cudnn/grouped_gemm/grouped_gemm_swiglu/compile.py" -if [ ! -f "$CUDNN_FE_COMPILE_API" ]; then +CUDNN_FE_FWD_COMPILE_API="$CUDNN_FE_ROOT/python/cudnn/grouped_gemm/grouped_gemm_swiglu/compile.py" +CUDNN_FE_BWD_COMPILE_API="$CUDNN_FE_ROOT/python/cudnn/grouped_gemm/grouped_gemm_dswiglu/compile.py" +if [ ! -f "$CUDNN_FE_FWD_COMPILE_API" ] && [ ! -f "$CUDNN_FE_BWD_COMPILE_API" ]; then if ! git -C "$CUDNN_FE_ROOT" apply --check "$CUDNN_FE_PATCH"; then echo "The bundled compile-only patch does not apply cleanly to $CUDNN_FE_ROOT." >&2 echo "Use the pinned base $CUDNN_FE_BASE_REV or provide an already-patched checkout." >&2 exit 1 fi git -C "$CUDNN_FE_ROOT" apply "$CUDNN_FE_PATCH" +elif [ ! -f "$CUDNN_FE_FWD_COMPILE_API" ] || [ ! -f "$CUDNN_FE_BWD_COMPILE_API" ]; then + echo "The cuDNN-FE checkout has only one of the required forward/backward compile APIs." >&2 + echo "Use the updated compile-only branch or a fresh checkout so both APIs come from one patch." >&2 + exit 1 fi if [ "$NUM_GPUS" -lt 2 ]; then echo "The E2E suite requires at least two SM100 GPUs" >&2 @@ -60,11 +65,12 @@ import inspect import jax import jax_tvm_ffi import tvm_ffi -from cudnn import compile_grouped_gemm_swiglu +from cudnn import compile_grouped_gemm_dswiglu, compile_grouped_gemm_swiglu print("jax", jax.__version__) print("devices", jax.devices()) -print("cudnn_frontend", inspect.getfile(compile_grouped_gemm_swiglu)) +print("cudnn_frontend_fwd", inspect.getfile(compile_grouped_gemm_swiglu)) +print("cudnn_frontend_bwd", inspect.getfile(compile_grouped_gemm_dswiglu)) print("jax_tvm_ffi", inspect.getfile(jax_tvm_ffi)) print("tvm_ffi", inspect.getfile(tvm_ffi)) assert all(device.platform == "gpu" for device in jax.devices()) diff --git a/tests/jax/test_cutedsl_moe.py b/tests/jax/test_cutedsl_moe.py index 3b5a583aeed..4b7bb5fb19b 100644 --- a/tests/jax/test_cutedsl_moe.py +++ b/tests/jax/test_cutedsl_moe.py @@ -12,11 +12,6 @@ import pytest from transformer_engine.jax import cpp_extensions as tex -from transformer_engine.jax.cutedsl_extensions.moe import ( - grouped_gemm_dswiglu_mxfp8, - pack_swiglu_pair, - unpack_swiglu_pair, -) from transformer_engine.jax.quantize import ( QuantizerFactory, ScaledTensorFactory, @@ -31,21 +26,23 @@ def test_swiglu_block_pack_round_trip(): """Gate/up packing alternates 32-column blocks and is reversible.""" gate = jnp.arange(2 * 3 * 64, dtype=jnp.float32).reshape(2, 3, 64) up = gate + 1000 - interleaved = pack_swiglu_pair(gate, up) + interleaved = tex.pack_swiglu_pair(gate, up) np.testing.assert_array_equal(interleaved[..., :32], gate[..., :32]) np.testing.assert_array_equal(interleaved[..., 32:64], up[..., :32]) - unpacked_gate, unpacked_up = unpack_swiglu_pair(interleaved) + unpacked_gate, unpacked_up = tex.unpack_swiglu_pair(interleaved) np.testing.assert_array_equal(unpacked_gate, gate) np.testing.assert_array_equal(unpacked_up, up) -def test_compile_only_api_bypasses_torch_wrapper(): - """The forward compiler must not execute cuDNN's Torch wrapper module.""" - from cudnn import compile_grouped_gemm_swiglu +def test_compile_only_api_bypasses_torch_wrappers(): + """Neither compiler executes a cuDNN Torch wrapper module.""" + from cudnn import compile_grouped_gemm_dswiglu, compile_grouped_gemm_swiglu assert callable(compile_grouped_gemm_swiglu) + assert callable(compile_grouped_gemm_dswiglu) assert "cudnn.grouped_gemm.grouped_gemm_swiglu.api" not in sys.modules + assert "cudnn.grouped_gemm.grouped_gemm_dswiglu.api" not in sys.modules @pytest.mark.parametrize( @@ -80,7 +77,7 @@ def test_swiglu_forward_fused_output_parity(group_sizes_tuple): wi_1 = jax.random.normal( jax.random.fold_in(key, 2), (experts, hidden, intermediate), dtype=jnp.bfloat16 ) - wi = pack_swiglu_pair(wi_0, wi_1) + wi = tex.pack_swiglu_pair(wi_0, wi_1) group_sizes = jnp.asarray(group_sizes_tuple, dtype=jnp.int32) assert len(set(group_sizes_tuple)) > 1 or experts == 1 assert all(offset % 256 == 0 for offset in np.cumsum(group_sizes_tuple)) @@ -105,7 +102,7 @@ def run(x_arg, wi_arg): casted_wi, contracting_dims=((1,), (1,)), ) - reference_gate, reference_up = unpack_swiglu_pair(reference) + reference_gate, reference_up = tex.unpack_swiglu_pair(reference) swiglu_reference = jax.nn.silu(reference_gate) * reference_up quantized_reference = tex.grouped_quantize( swiglu_reference, @@ -167,7 +164,7 @@ def run(x_arg, wi_arg): assert swiglu_row.dtype == swiglu_col.dtype == jnp.float8_e4m3fn assert scale_row.dtype == scale_col.dtype == jnp.float8_e8m0fnu - gate, up = unpack_swiglu_pair(combined) + gate, up = tex.unpack_swiglu_pair(combined) swiglu_reference = np.asarray(jax.nn.silu(gate) * up, dtype=np.float32) for is_colwise, payload, scale, reference_payload, reference_scale in ( (False, swiglu_row, scale_row, reference_row, reference_scale_row), @@ -223,10 +220,9 @@ def test_dswiglu_backward_quantized_output_parity(): if get_device_compute_capability(0) != 100: pytest.skip("cuDNN frontend grouped GEMM dSwiGLU requires SM100") - from transformer_engine.jax.cutedsl_extensions.moe import load_grouped_gemm_dswiglu_kernel - - load_grouped_gemm_dswiglu_kernel() - import cutlass.jax # noqa: F401 # pylint: disable=unused-import,import-outside-toplevel + dependencies_available, dependency_error = tex.grouped_gemm_swiglu_dependencies_available() + if not dependencies_available: + pytest.skip(f"TVM-FFI JAX dependencies are unavailable: {dependency_error}") except (ImportError, RuntimeError) as exc: pytest.skip(f"CuTeDSL JAX dependencies are unavailable: {exc}") @@ -238,7 +234,7 @@ def test_dswiglu_backward_quantized_output_parity(): ) gate = jax.random.normal(jax.random.fold_in(key, 2), (rows, intermediate), dtype=jnp.bfloat16) up = jax.random.normal(jax.random.fold_in(key, 3), (rows, intermediate), dtype=jnp.bfloat16) - packed_forward = pack_swiglu_pair(gate, up) + packed_forward = tex.pack_swiglu_pair(gate, up) group_sizes = jnp.asarray([rows], dtype=jnp.int32) quantizers = QuantizerFactory.create_set( scaling_mode=ScalingMode.MXFP8_1D_SCALING, @@ -261,9 +257,9 @@ def run(d_eo_arg, wo_arg, packed_arg): casted_wo, contracting_dims=((1,), (2,)), ) - d_row, d_col, scale_row, scale_col, dprob = grouped_gemm_dswiglu_mxfp8( + d_row, d_col, scale_row, scale_col, dprob = tex.grouped_gemm_dswiglu( casted_d_eo.data.reshape(rows, hidden, 1), - casted_wo.data.reshape(experts, intermediate, hidden).transpose(1, 2, 0), + casted_wo.data.reshape(experts, intermediate, hidden), packed_arg.reshape(rows, 2 * intermediate, 1), casted_d_eo.scale_inv, casted_wo.scale_inv, @@ -273,6 +269,16 @@ def run(d_eo_arg, wo_arg, packed_arg): ) return reference_intermediate, d_row, d_col, scale_row, scale_col, dprob + run.lower(d_eo, wo, packed_forward) + grouped_gemm_module = importlib.import_module( + "transformer_engine.jax.cpp_extensions.grouped_gemm_swiglu" + ) + + assert any( + name.startswith("te_grouped_gemm_dswiglu.") + for name in grouped_gemm_module._REGISTERED_TARGETS # pylint: disable=protected-access + ) + reference_intermediate, d_row, d_col, scale_row, scale_col, dprob = run( d_eo, wo, packed_forward ) @@ -287,7 +293,7 @@ def run(d_eo_arg, wo_arg, packed_arg): ref = reference_intermediate.astype(jnp.float32) d_up = ref * swish d_gate = ref * up.astype(jnp.float32) * sigmoid * (1 + gate.astype(jnp.float32) * (1 - sigmoid)) - packed_reference = np.asarray(pack_swiglu_pair(d_gate, d_up), dtype=np.float32) + packed_reference = np.asarray(tex.pack_swiglu_pair(d_gate, d_up), dtype=np.float32) for is_colwise, payload, scale in ( (False, d_row, scale_row), diff --git a/tests/jax/tvm_ffi_grouped_mlp_multiprocess.py b/tests/jax/tvm_ffi_grouped_mlp_multiprocess.py index 6536771a239..fffc8d13e5f 100644 --- a/tests/jax/tvm_ffi_grouped_mlp_multiprocess.py +++ b/tests/jax/tvm_ffi_grouped_mlp_multiprocess.py @@ -158,7 +158,6 @@ def metrics(custom, mapped, a, b): def _run_ragged_multiprocess_cell(mesh: Mesh) -> None: """Exercise eight unequal experts through a global multiprocess shard_map.""" from transformer_engine.jax import cpp_extensions as tex - from transformer_engine.jax.cutedsl_extensions.moe import unpack_swiglu_pair from transformer_engine.jax.quantize import ( QuantizerFactory, ScaledTensorFactory, @@ -255,7 +254,7 @@ def local_fused(a, b, a_scale, b_scale, probabilities): np.asarray(combined_local[:, :, 0]), np.asarray(combined_reference), ) - gate, up = unpack_swiglu_pair(combined_reference) + gate, up = tex.unpack_swiglu_pair(combined_reference) swiglu_reference = jax.nn.silu(gate) * up quantizers = QuantizerFactory.create_set( diff --git a/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py b/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py index 42d83933d81..68e24469e7f 100644 --- a/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py +++ b/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py @@ -15,7 +15,13 @@ from .base import BasePrimitive, register_primitive from .misc import get_padded_spec -__all__ = ["grouped_gemm_swiglu", "grouped_gemm_swiglu_dependencies_available"] +__all__ = [ + "grouped_gemm_dswiglu", + "grouped_gemm_swiglu", + "grouped_gemm_swiglu_dependencies_available", + "pack_swiglu_pair", + "unpack_swiglu_pair", +] _INPUT_ROLES = ("a", "b", "sfa", "sfb", "padded_offsets", "alpha", "prob", "norm_const") @@ -24,6 +30,37 @@ _REGISTRATION_LOCK = threading.Lock() +def pack_swiglu_pair(gate: jax.Array, up: jax.Array) -> jax.Array: + """Interleave 32-column gate/up blocks as required by the kernels.""" + if gate.shape != up.shape: + raise ValueError(f"gate shape {gate.shape} must match up shape {up.shape}") + if gate.shape[-1] % 32: + raise ValueError(f"SwiGLU intermediate dimension {gate.shape[-1]} must be divisible by 32") + blocks = gate.shape[-1] // 32 + return jnp.stack( + ( + gate.reshape(*gate.shape[:-1], blocks, 32), + up.reshape(*up.shape[:-1], blocks, 32), + ), + axis=-2, + ).reshape(*gate.shape[:-1], 2 * gate.shape[-1]) + + +def unpack_swiglu_pair(interleaved: jax.Array) -> tuple[jax.Array, jax.Array]: + """Undo :func:`pack_swiglu_pair`.""" + if interleaved.shape[-1] % 64: + raise ValueError( + f"Interleaved SwiGLU dimension {interleaved.shape[-1]} must be divisible by 64" + ) + intermediate = interleaved.shape[-1] // 2 + blocks = intermediate // 32 + paired = interleaved.reshape(*interleaved.shape[:-1], blocks, 2, 32) + return ( + paired[..., 0, :].reshape(*interleaved.shape[:-1], intermediate), + paired[..., 1, :].reshape(*interleaved.shape[:-1], intermediate), + ) + + def _ceil_div(value: int, divisor: int) -> int: return (value + divisor - 1) // divisor @@ -56,7 +93,10 @@ def grouped_gemm_swiglu_dependencies_available() -> tuple[bool, str]: try: import jax_tvm_ffi # noqa: F401 # pylint: disable=unused-import,import-outside-toplevel import tvm_ffi # noqa: F401 # pylint: disable=unused-import,import-outside-toplevel - from cudnn import compile_grouped_gemm_swiglu # noqa: F401 + from cudnn import ( # noqa: F401 + compile_grouped_gemm_dswiglu, + compile_grouped_gemm_swiglu, + ) except (ImportError, ModuleNotFoundError, RuntimeError, AttributeError) as exc: return False, str(exc) return True, "" @@ -173,6 +213,131 @@ def _compile_and_register(avals_in, avals_out): return target_name, abi +_DSWIGLU_INPUT_ROLES = ( + "a", + "b", + "c", + "sfa", + "sfb", + "padded_offsets", + "alpha", + "beta", + "prob", + "norm_const", +) +_DSWIGLU_OUTPUT_ROLES = ("d", "d_col", "sfd_row", "sfd_col", "dprob") + + +def _dswiglu_output_avals(a_aval, b_aval, c_aval, output_dtype): + if a_aval.ndim != 3 or a_aval.shape[-1] != 1: + raise ValueError(f"Expected A[M,K,1], got {a_aval.shape}") + if b_aval.ndim != 3: + raise ValueError(f"Expected physical B[E,N,K], got {b_aval.shape}") + m, k_a, _ = a_aval.shape + _, n, k_b = b_aval.shape + if k_a != k_b: + raise ValueError(f"A K={k_a} does not match B K={k_b}") + if c_aval.shape != (m, 2 * n, 1): + raise ValueError(f"Expected C shape {(m, 2 * n, 1)}, got {c_aval.shape}") + if m % 256: + raise ValueError(f"Padded activation rows M={m} must be divisible by 256") + return ( + jax.core.ShapedArray((m, 2 * n, 1), output_dtype), + jax.core.ShapedArray((m, 2 * n, 1), output_dtype), + jax.core.ShapedArray((_scale_size(m, 2 * n, colwise=False),), jnp.float8_e8m0fnu), + jax.core.ShapedArray((_scale_size(m, 2 * n, colwise=True),), jnp.float8_e8m0fnu), + jax.core.ShapedArray((m, 1, 1), jnp.float32), + ) + + +def _compile_and_register_dswiglu(avals_in, avals_out): + try: + import jax_tvm_ffi + import tvm_ffi + from cudnn import ( + GroupedGemmDswigluConfig, + GroupedGemmOperandDesc, + compile_grouped_gemm_dswiglu, + ) + except (ImportError, ModuleNotFoundError, RuntimeError, AttributeError) as exc: + raise RuntimeError( + "TVM-FFI grouped dSwiGLU requires the patched cuDNN-FE compiler package, " + "apache-tvm-ffi, and jax-tvm-ffi" + ) from exc + + layouts = { + "a": (1, 0, 2), + "b": (1, 0, 2), + "c": (1, 0, 2), + "sfa": (0,), + "sfb": (0,), + "padded_offsets": (0,), + "alpha": (0,), + "beta": (0,), + "prob": (1, 0, 2), + "norm_const": (0,), + "d": (1, 0, 2), + "d_col": (1, 0, 2), + "sfd_row": (0,), + "sfd_col": (0,), + "dprob": (1, 0, 2), + } + operands = {} + for role, aval in zip(_DSWIGLU_INPUT_ROLES, avals_in): + shape = aval.shape + if role == "b": + # Present the compact [E,N,K] XLA buffer to the compiler as the + # kernel's logical [N,K,E] tensor without a runtime transpose. + shape = (aval.shape[1], aval.shape[2], aval.shape[0]) + operands[role] = GroupedGemmOperandDesc( + role=role, + dtype=_dtype_name(aval.dtype), + shape=shape, + stride_order=layouts[role], + ) + for role, aval in zip(_DSWIGLU_OUTPUT_ROLES, avals_out): + operands[role] = GroupedGemmOperandDesc( + role=role, + dtype=_dtype_name(aval.dtype), + shape=aval.shape, + stride_order=layouts[role], + ) + + function, abi = compile_grouped_gemm_dswiglu( + operands=operands, + config=GroupedGemmDswigluConfig( + sf_vec_size=32, + mma_tiler_mn=(256, 256), + cluster_shape_mn=(2, 1), + discrete_col_sfd=True, + ), + ) + if not isinstance(function, tvm_ffi.Function): + raise TypeError( + "cuDNN-FE compile_grouped_gemm_dswiglu returned a Python wrapper; " + "the JAX execution target must be a native tvm_ffi.Function" + ) + expected = [(role, "arg") for role in _DSWIGLU_INPUT_ROLES] + [ + (role, "ret") for role in _DSWIGLU_OUTPUT_ROLES + ] + [("stream", "stream")] + actual = [(entry.role, entry.kind) for entry in abi.entries] + if actual != expected: + raise ValueError(f"Unsupported grouped dSwiGLU ABI: expected {expected}, got {actual}") + + target_name = f"te_grouped_gemm_dswiglu.{abi.key.rsplit('.', maxsplit=1)[-1]}" + with _REGISTRATION_LOCK: + if target_name not in _REGISTERED_TARGETS: + jax_tvm_ffi.register_ffi_target( + target_name, + function, + arg_spec=list(abi.arg_spec), + platform="gpu", + allow_cuda_graph=True, + ) + _REGISTERED_TARGETS[target_name] = (function, abi) + return target_name, abi + + class GroupedGemmSwigluPrimitive(BasePrimitive): """Fused grouped GEMM, SwiGLU, and dual MXFP8 quantization.""" @@ -314,6 +479,164 @@ def shardy_sharding_rule(*args, **kwargs): register_primitive(GroupedGemmSwigluPrimitive) +class GroupedGemmDswigluPrimitive(BasePrimitive): + """Fused grouped GEMM, dSwiGLU, and dual MXFP8 quantization.""" + + name = "te_grouped_gemm_dswiglu_tvm_ffi" + multiple_results = True + impl_static_args = (10,) + inner_primitive = None + outer_primitive = None + + @staticmethod + def abstract( + a_aval, + b_aval, + c_aval, + sfa_aval, + sfb_aval, + padded_offsets_aval, + alpha_aval, + beta_aval, + prob_aval, + norm_const_aval, + *, + output_dtype, + ): + del sfa_aval, sfb_aval, alpha_aval, beta_aval, norm_const_aval + m = a_aval.shape[0] + experts = b_aval.shape[0] + if padded_offsets_aval.shape != (experts,): + raise ValueError( + f"Expected padded_offsets shape {(experts,)}, got {padded_offsets_aval.shape}" + ) + if prob_aval.shape != (m, 1, 1): + raise ValueError(f"Expected prob shape {(m, 1, 1)}, got {prob_aval.shape}") + return _dswiglu_output_avals(a_aval, b_aval, c_aval, output_dtype) + + @staticmethod + def lowering(ctx, *args, output_dtype): + del output_dtype + target_name, abi = _compile_and_register_dswiglu(ctx.avals_in, ctx.avals_out) + operand_layouts = [entry.stride_order for entry in abi.entries if entry.kind == "arg"] + result_layouts = [entry.stride_order for entry in abi.entries if entry.kind == "ret"] + return jax.ffi.ffi_lowering( + target_name, + operand_layouts=operand_layouts, + result_layouts=result_layouts, + )(ctx, *args) + + @staticmethod + def impl( + a, + b, + c, + sfa, + sfb, + padded_offsets, + alpha, + beta, + prob, + norm_const, + output_dtype, + ): + if GroupedGemmDswigluPrimitive.inner_primitive is None: + raise RuntimeError("GroupedGemmDswigluPrimitive has not been registered") + return GroupedGemmDswigluPrimitive.inner_primitive.bind( + a, + b, + c, + sfa, + sfb, + padded_offsets, + alpha, + beta, + prob, + norm_const, + output_dtype=output_dtype, + ) + + @staticmethod + def batcher(batched_args, batch_dims, *, output_dtype): + del batched_args, output_dtype + raise NotImplementedError( + f"GroupedGemmDswigluPrimitive does not support vmap batch dimensions {batch_dims}" + ) + + @staticmethod + def partition(output_dtype, mesh, arg_infos, result_infos): + del result_infos + a_spec = get_padded_spec(arg_infos[0]) + m_axis = a_spec[0] + experts = arg_infos[1].shape[0] + if experts != 1 and m_axis is not None: + raise NotImplementedError( + "Token-axis custom partitioning is currently supported only for the standalone " + "single-expert experiment; MoEBlock integration remains inside shard_map" + ) + arg_shardings = ( + NamedSharding(mesh, PartitionSpec(m_axis, None, None)), + NamedSharding(mesh, PartitionSpec(None, None, None)), + NamedSharding(mesh, PartitionSpec(m_axis, None, None)), + NamedSharding(mesh, PartitionSpec(m_axis)), + NamedSharding(mesh, PartitionSpec(None)), + NamedSharding(mesh, PartitionSpec(None)), + NamedSharding(mesh, PartitionSpec(None)), + NamedSharding(mesh, PartitionSpec(None)), + NamedSharding(mesh, PartitionSpec(m_axis, None, None)), + NamedSharding(mesh, PartitionSpec(None)), + ) + out_shardings = ( + NamedSharding(mesh, PartitionSpec(m_axis, None, None)), + NamedSharding(mesh, PartitionSpec(m_axis, None, None)), + NamedSharding(mesh, PartitionSpec(m_axis)), + NamedSharding(mesh, PartitionSpec(m_axis)), + NamedSharding(mesh, PartitionSpec(m_axis, None, None)), + ) + + def sharded_impl( + a, + b, + c, + sfa, + sfb, + padded_offsets, + alpha, + beta, + prob, + norm_const, + ): + if experts == 1: + padded_offsets = jnp.asarray([a.shape[0]], dtype=jnp.int32) + return GroupedGemmDswigluPrimitive.impl( + a, + b, + c, + sfa, + sfb, + padded_offsets, + alpha, + beta, + prob, + norm_const, + output_dtype=output_dtype, + ) + + return mesh, sharded_impl, out_shardings, arg_shardings + + @staticmethod + def shardy_sharding_rule(*args, **kwargs): + del args, kwargs + return ( + "m k one, experts n k, m two_n one, sfa, sfb, experts, experts, experts, " + "m prob_one_a prob_one_b, norm -> " + "m two_n one, m two_n one, sfd_row, sfd_col, m prob_one_a prob_one_b" + ) + + +register_primitive(GroupedGemmDswigluPrimitive) + + def grouped_gemm_swiglu( a: jax.Array, b: jax.Array, @@ -348,3 +671,41 @@ def grouped_gemm_swiglu( compute_dtype=jnp.dtype(compute_dtype), output_dtype=jnp.dtype(output_dtype), ) + + +def grouped_gemm_dswiglu( + a: jax.Array, + b: jax.Array, + c: jax.Array, + sfa: jax.Array, + sfb: jax.Array, + padded_offsets: jax.Array, + prob: jax.Array, + *, + output_dtype, +) -> tuple[jax.Array, jax.Array, jax.Array, jax.Array, jax.Array]: + """Run the TVM-FFI fused grouped dSwiGLU primitive. + + ``b`` uses TE's compact physical ``[E,N,K]`` representation. The returned + dprob buffer is not consumed because the kernel accumulates into it + atomically; the MoE VJP computes the routing cotangent explicitly. + """ + if b.ndim != 3: + raise ValueError(f"Expected physical B[E,N,K], got {b.shape}") + experts = b.shape[0] + alpha = jnp.ones((experts,), dtype=jnp.float32) + beta = jnp.ones((experts,), dtype=jnp.float32) + norm_const = jnp.ones((1,), dtype=jnp.float32) + return GroupedGemmDswigluPrimitive.outer_primitive.bind( + a, + b, + c, + sfa.reshape(-1), + sfb.reshape(-1), + padded_offsets.astype(jnp.int32), + alpha, + beta, + prob.astype(jnp.float32), + norm_const, + output_dtype=jnp.dtype(output_dtype), + ) diff --git a/transformer_engine/jax/cutedsl_extensions/__init__.py b/transformer_engine/jax/cutedsl_extensions/__init__.py deleted file mode 100644 index 27730a4e555..00000000000 --- a/transformer_engine/jax/cutedsl_extensions/__init__.py +++ /dev/null @@ -1,4 +0,0 @@ -# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# -# See LICENSE for license information. -"""JAX wrappers for kernels authored with the CUTLASS CuTe DSL.""" diff --git a/transformer_engine/jax/cutedsl_extensions/moe.py b/transformer_engine/jax/cutedsl_extensions/moe.py deleted file mode 100644 index b51b2892e6c..00000000000 --- a/transformer_engine/jax/cutedsl_extensions/moe.py +++ /dev/null @@ -1,576 +0,0 @@ -# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# -# See LICENSE for license information. -"""JAX binding for cuDNN frontend's MXFP8 grouped-GEMM SwiGLU kernel. - -The kernel is intentionally imported from the installed -``nvidia-cudnn-frontend`` distribution. The public cuDNN frontend package -initializers import the Torch API wrappers eagerly, so this module loads the -kernel source directly under a private namespace. No kernel implementation is -vendored into Transformer Engine. -""" - -from __future__ import annotations - -from functools import lru_cache -import importlib.metadata -import importlib.util -import sys -import threading -import types -from typing import Any - -import jax -import jax.numpy as jnp - - -_CUDNN_FRONTEND_DISTRIBUTION = "nvidia-cudnn-frontend" -_PRIVATE_PACKAGE = "_transformer_engine_cudnn_grouped_gemm" -_LOAD_LOCK = threading.Lock() - - -def _namespace_package(name: str, path) -> types.ModuleType: - module = types.ModuleType(name) - module.__package__ = name - module.__path__ = [str(path)] - sys.modules[name] = module - return module - - -def _load_source_module(name: str, path): - spec = importlib.util.spec_from_file_location(name, path) - if spec is None or spec.loader is None: - raise ImportError(f"Could not create an import spec for {path}") - module = importlib.util.module_from_spec(spec) - sys.modules[name] = module - try: - spec.loader.exec_module(module) - except Exception: - sys.modules.pop(name, None) - raise - return module - - -def _patch_atomic_add_float32_for_cutlass_dsl(utils_module) -> None: - """Patch cuDNN frontend utility source for newer CUTLASS DSL bindings.""" - try: - import cutlass - from cutlass._mlir.dialects import nvvm - from cutlass.cute.typing import Float32 - except ImportError: - return - - def atomic_add_float32(ptr, value: Float32, *, loc=None, ip=None) -> Float32: - a = value.ir_value(loc=loc, ip=ip) - try: - old_value = nvvm.atomicrmw( - op=cutlass._mlir.dialects.nvvm.AtomicOpKind.FADD, - ptr=ptr, - a=a, - res=a.type, - loc=loc, - ip=ip, - ) - except TypeError as exc: - if "res" not in str(exc): - raise - old_value = nvvm.atomicrmw( - op=cutlass._mlir.dialects.nvvm.AtomicOpKind.FADD, - ptr=ptr, - a=a, - loc=loc, - ip=ip, - ) - return Float32(old_value) - - utils_module.atomic_add_float32 = atomic_add_float32 - - -def _ensure_torch_annotation_stub() -> bool: - # Do not shadow a real Torch installation. If it is installed, the - # unmodified cuDNN frontend source can import it normally; the sentinel is - # only for JAX-only environments where no Torch module exists. - try: - importlib.metadata.distribution("torch") - torch_installed = True - except importlib.metadata.PackageNotFoundError: - torch_installed = False - torch_stubbed = "torch" not in sys.modules and not torch_installed - if torch_stubbed: - torch_stub = types.ModuleType("torch") - torch_stub.Tensor = object - torch_stub.float4_e2m1fn_x2 = object() - torch_stub.__transformer_engine_cutedsl_stub__ = True - sys.modules["torch"] = torch_stub - return torch_stubbed - - -def _load_grouped_gemm_kernel_module(subpackage: str, filename: str): - distribution = importlib.metadata.distribution(_CUDNN_FRONTEND_DISTRIBUTION) - grouped_root = distribution.locate_file("cudnn/grouped_gemm") - utils_path = grouped_root / "utils.py" - kernel_root = grouped_root / subpackage - kernel_path = kernel_root / filename - if not utils_path.is_file() or not kernel_path.is_file(): - raise ImportError( - f"{_CUDNN_FRONTEND_DISTRIBUTION} {distribution.version} does not contain " - f"the {subpackage} CuTeDSL kernel sources" - ) - - _namespace_package(_PRIVATE_PACKAGE, grouped_root) - private_subpackage = f"{_PRIVATE_PACKAGE}.{subpackage}" - _namespace_package(private_subpackage, kernel_root) - - torch_stubbed = _ensure_torch_annotation_stub() - try: - utils_module = _load_source_module(f"{_PRIVATE_PACKAGE}.utils", utils_path) - _patch_atomic_add_float32_for_cutlass_dsl(utils_module) - module_name = f"{private_subpackage}.{filename.removesuffix('.py')}" - return _load_source_module(module_name, kernel_path) - except Exception: - if torch_stubbed: - sys.modules.pop("torch", None) - raise - - -@lru_cache(maxsize=1) -def load_grouped_gemm_swiglu_kernel(): - """Load the cuDNN frontend forward kernel class without its Torch API.""" - with _LOAD_LOCK: - kernel_module = _load_grouped_gemm_kernel_module( - "grouped_gemm_swiglu", "grouped_gemm_swiglu_quant.py" - ) - try: - return kernel_module.BlockScaledContiguousGroupedGemmKernel - except AttributeError as exc: - distribution = importlib.metadata.distribution(_CUDNN_FRONTEND_DISTRIBUTION) - raise ImportError( - f"{_CUDNN_FRONTEND_DISTRIBUTION} {distribution.version} has an " - "incompatible grouped-GEMM SwiGLU kernel API" - ) from exc - - -@lru_cache(maxsize=1) -def load_grouped_gemm_dswiglu_kernel(): - """Load the cuDNN frontend dSwiGLU backward kernel class.""" - with _LOAD_LOCK: - kernel_module = _load_grouped_gemm_kernel_module( - "grouped_gemm_dswiglu", "grouped_gemm_dswiglu_quant.py" - ) - try: - return kernel_module.BlockScaledContiguousGroupedGemmKernel - except AttributeError as exc: - distribution = importlib.metadata.distribution(_CUDNN_FRONTEND_DISTRIBUTION) - raise ImportError( - f"{_CUDNN_FRONTEND_DISTRIBUTION} {distribution.version} has an " - "incompatible grouped-GEMM dSwiGLU kernel API" - ) from exc - - -def pack_swiglu_pair(gate: jax.Array, up: jax.Array) -> jax.Array: - """Interleave 32-column gate/up blocks as required by the kernel.""" - if gate.shape != up.shape: - raise ValueError(f"gate shape {gate.shape} must match up shape {up.shape}") - if gate.shape[-1] % 32: - raise ValueError(f"SwiGLU intermediate dimension {gate.shape[-1]} must be divisible by 32") - blocks = gate.shape[-1] // 32 - return jnp.stack( - ( - gate.reshape(*gate.shape[:-1], blocks, 32), - up.reshape(*up.shape[:-1], blocks, 32), - ), - axis=-2, - ).reshape(*gate.shape[:-1], 2 * gate.shape[-1]) - - -def unpack_swiglu_pair(interleaved: jax.Array) -> tuple[jax.Array, jax.Array]: - """Undo :func:`pack_swiglu_pair`.""" - if interleaved.shape[-1] % 64: - raise ValueError( - f"Interleaved SwiGLU dimension {interleaved.shape[-1]} must be divisible by 64" - ) - intermediate = interleaved.shape[-1] // 2 - blocks = intermediate // 32 - paired = interleaved.reshape(*interleaved.shape[:-1], blocks, 2, 32) - return ( - paired[..., 0, :].reshape(*interleaved.shape[:-1], intermediate), - paired[..., 1, :].reshape(*interleaved.shape[:-1], intermediate), - ) - - -def _ceil_div(x: int, y: int) -> int: - return (x + y - 1) // y - - -@lru_cache(maxsize=None) -def _make_launcher( - expert_count: int, - sf_vec_size: int, - mma_tiler_m: int, - mma_tiler_n: int, -): - try: - import cutlass - from cutlass import cute - except ImportError as exc: - raise ImportError( - "NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION=1 requires nvidia-cutlass-dsl" - ) from exc - - kernel_cls = load_grouped_gemm_swiglu_kernel() - use_2cta_instrs = mma_tiler_m == 256 - cluster_shape = (2, 1) if use_2cta_instrs else (1, 1) - kernel = kernel_cls( - sf_vec_size=sf_vec_size, - acc_dtype=cutlass.Float32, - use_2cta_instrs=use_2cta_instrs, - mma_tiler_mn=(mma_tiler_m, mma_tiler_n), - cluster_shape_mn=cluster_shape, - vector_f32=False, - generate_sfd=True, - # TE's grouped colwise tensor stores an independently padded scale - # segment for every expert. The kernel's default colwise SFD layout - # treats the concatenated M dimension as one matrix, which makes the - # FC2 wgrad read another expert's scales (and eventually emit NaNs). - discrete_col_sfd=True, - expert_cnt=expert_count, - use_mono_increase_expert_idx=True, - ) - max_active_clusters = cutlass.utils.HardwareInfo().get_max_active_clusters( - cluster_shape[0] * cluster_shape[1] - ) - - @cute.jit - def launch( - stream, - a, - b, - sfa, - sfb, - padded_offsets, - alpha, - prob, - norm_const, - c, - d, - d_col, - sfd_row, - sfd_col, - ): - kernel( - a, - b, - c, - d, - d_col, - sfa, - sfb, - sfd_row, - sfd_col, - None, - norm_const, - padded_offsets, - alpha, - prob, - max_active_clusters, - stream, - ) - - return launch - - -def grouped_gemm_swiglu_mxfp8( - a: jax.Array, - b: jax.Array, - sfa: jax.Array, - sfb: jax.Array, - padded_offsets: jax.Array, - prob: jax.Array, - *, - compute_dtype: Any, - output_dtype: Any, -) -> tuple[jax.Array, jax.Array, jax.Array, jax.Array, jax.Array]: - """Call cuDNN frontend's grouped MXFP8 GEMM+SwiGLU+quant kernel. - - Args: - a: Quantized activation payload with physical shape ``[M, K, 1]``. - b: Quantized, block-interleaved colwise weights with physical shape - ``[E, N, K]``. - sfa/sfb: Pre-swizzled E8M0 inverse-scale buffers. - padded_offsets: Exclusive, 256-aligned end offset for each expert. - prob: Per-row multiplier, physical shape ``[M, 1, 1]``. - compute_dtype: Unquantized combined-projection dtype. - output_dtype: MXFP8 payload dtype for the quantized SwiGLU output. - - Returns: - Raw combined projection, rowwise payload, colwise payload, rowwise - inverse scales, and colwise inverse scales. - """ - try: - from cutlass.jax import TensorSpec, cutlass_call - except ImportError as exc: - raise ImportError( - "NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION=1 requires CUTLASS JAX bindings" - ) from exc - - if a.ndim != 3 or a.shape[-1] != 1: - raise ValueError(f"Expected A[M,K,1], got {a.shape}") - if b.ndim != 3: - raise ValueError(f"Expected physical B[E,N,K], got {b.shape}") - expert_count, n, k_b = b.shape - m, k_a, _ = a.shape - if k_a != k_b: - raise ValueError(f"A K={k_a} does not match B K={k_b}") - if n % 2: - raise ValueError(f"Combined SwiGLU N={n} must be even") - intermediate = n // 2 - - # The public kernel requires 256-aligned expert ranges. M is static at - # trace time, while the individual offsets remain runtime values. - if m % 256: - raise ValueError(f"Padded activation rows M={m} must be divisible by 256") - - sf_vec_size = 32 - row_scale_size = ( - 32 * 4 * _ceil_div(m, 128) * 4 * _ceil_div(_ceil_div(intermediate, sf_vec_size), 4) - ) - col_scale_size = ( - 32 * 4 * _ceil_div(intermediate, 128) * 4 * _ceil_div(_ceil_div(m, sf_vec_size), 4) - ) - outputs = ( - jax.ShapeDtypeStruct((m, n, 1), compute_dtype), - jax.ShapeDtypeStruct((m, intermediate, 1), output_dtype), - jax.ShapeDtypeStruct((m, intermediate, 1), output_dtype), - jax.ShapeDtypeStruct((row_scale_size,), jnp.float8_e8m0fnu), - jax.ShapeDtypeStruct((col_scale_size,), jnp.float8_e8m0fnu), - ) - launcher = _make_launcher(expert_count, sf_vec_size, 256, 256) - call = cutlass_call( - launcher, - output_shape_dtype=outputs, - input_spec=( - # Singleton L must be outermost so K, rather than L, is the - # leading (stride-1) dimension seen by the kernel. - TensorSpec(layout=(1, 0, 2)), - # physical colwise [E,N,K] -> logical K-major [N,K,E] - TensorSpec(mode=(1, 2, 0)), - TensorSpec(), - TensorSpec(), - TensorSpec(), - TensorSpec(), - TensorSpec(), - TensorSpec(), - ), - output_spec=( - TensorSpec(layout=(1, 0, 2)), - TensorSpec(layout=(1, 0, 2)), - TensorSpec(layout=(1, 0, 2)), - TensorSpec(), - TensorSpec(), - ), - allow_cuda_graph=True, - ) - alpha = jnp.ones((expert_count,), dtype=jnp.float32) - # With norm_const=1 the generated E8M0 factors are directly the - # scale-inverse values expected by TE's MXFP8 tensor representation. - norm_const = jnp.ones((1,), dtype=jnp.float32) - return call( - a, - b, - sfa.reshape(-1), - sfb.reshape(-1), - padded_offsets.astype(jnp.int32), - alpha, - prob.astype(jnp.float32), - norm_const, - ) - - -@lru_cache(maxsize=None) -def _make_dswiglu_launcher( - expert_count: int, - sf_vec_size: int, - mma_tiler_m: int, - mma_tiler_n: int, -): - try: - import cutlass - from cutlass import cute - except ImportError as exc: - raise ImportError( - "NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION=1 requires nvidia-cutlass-dsl" - ) from exc - - kernel_cls = load_grouped_gemm_dswiglu_kernel() - use_2cta_instrs = mma_tiler_m == 256 - cluster_shape = (2, 1) if use_2cta_instrs else (1, 1) - kernel = kernel_cls( - sf_vec_size=sf_vec_size, - acc_dtype=cutlass.Float32, - use_2cta_instrs=use_2cta_instrs, - mma_tiler_mn=(mma_tiler_m, mma_tiler_n), - cluster_shape_mn=cluster_shape, - vectorized_f32=False, - discrete_col_sfd=True, - expert_cnt=expert_count, - use_mono_increase_expert_idx=True, - ) - max_active_clusters = cutlass.utils.HardwareInfo().get_max_active_clusters( - cluster_shape[0] * cluster_shape[1] - ) - - @cute.jit - def launch( - stream, - a, - b, - c, - sfa, - sfb, - padded_offsets, - alpha, - beta, - prob, - norm_const, - d, - d_col, - sfd_row, - sfd_col, - dprob, - ): - kernel( - a, - b, - c, - d, - d_col, - sfa, - sfb, - sfd_row, - sfd_col, - None, - norm_const, - padded_offsets, - alpha, - beta, - prob, - dprob, - max_active_clusters, - stream, - ) - - return launch - - -def grouped_gemm_dswiglu_mxfp8( - a: jax.Array, - b: jax.Array, - c: jax.Array, - sfa: jax.Array, - sfb: jax.Array, - padded_offsets: jax.Array, - prob: jax.Array, - *, - output_dtype: Any, -) -> tuple[jax.Array, jax.Array, jax.Array, jax.Array, jax.Array]: - """Call cuDNN frontend's grouped MXFP8 GEMM+dSwiGLU+quant kernel. - - Args: - a: Quantized upstream gradient payload with physical shape ``[M, K, 1]``. - b: Quantized FC2 weights with physical shape ``[N, K, E]``. - c: Forward packed gate/up activations with shape ``[M, 2N, 1]``. - sfa/sfb: Pre-swizzled E8M0 inverse-scale buffers. - padded_offsets: Exclusive, 256-aligned end offset for each expert. - prob: Per-row multiplier from early top-k weighting, or ones. - output_dtype: MXFP8 payload dtype for the quantized packed output. - - Returns: - Rowwise and colwise packed dSwiGLU payloads, rowwise and colwise - inverse scales, and the per-row probability gradient. - """ - try: - from cutlass.jax import TensorSpec, cutlass_call - except ImportError as exc: - raise ImportError( - "NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION=1 requires CUTLASS JAX bindings" - ) from exc - - if a.ndim != 3 or a.shape[-1] != 1: - raise ValueError(f"Expected A[M,K,1], got {a.shape}") - if b.ndim != 3: - raise ValueError(f"Expected B[N,K,E], got {b.shape}") - if c.ndim != 3 or c.shape[-1] != 1: - raise ValueError(f"Expected C[M,2N,1], got {c.shape}") - m, k_a, _ = a.shape - n, k_b, expert_count = b.shape - if k_a != k_b: - raise ValueError(f"A K={k_a} does not match B K={k_b}") - if c.shape != (m, 2 * n, 1): - raise ValueError(f"Expected C shape {(m, 2 * n, 1)}, got {c.shape}") - if m % 256: - raise ValueError(f"Padded activation rows M={m} must be divisible by 256") - - sf_vec_size = 32 - row_scale_size = 32 * 4 * _ceil_div(m, 128) * 4 * _ceil_div( - _ceil_div(2 * n, sf_vec_size), 4 - ) - col_scale_size = 32 * 4 * _ceil_div(2 * n, 128) * 4 * _ceil_div( - _ceil_div(m, sf_vec_size), 4 - ) - outputs = ( - jax.ShapeDtypeStruct((m, 2 * n, 1), output_dtype), - jax.ShapeDtypeStruct((m, 2 * n, 1), output_dtype), - jax.ShapeDtypeStruct((row_scale_size,), jnp.float8_e8m0fnu), - jax.ShapeDtypeStruct((col_scale_size,), jnp.float8_e8m0fnu), - jax.ShapeDtypeStruct((m, 1, 1), jnp.float32), - ) - launcher = _make_dswiglu_launcher(expert_count, sf_vec_size, 256, 256) - call = cutlass_call( - launcher, - output_shape_dtype=outputs, - input_spec=( - TensorSpec(layout=(1, 0, 2)), - TensorSpec(layout=(1, 0, 2)), - TensorSpec(layout=(1, 0, 2)), - TensorSpec(), - TensorSpec(), - TensorSpec(), - TensorSpec(), - TensorSpec(), - TensorSpec(), - TensorSpec(), - ), - output_spec=( - TensorSpec(layout=(1, 0, 2)), - TensorSpec(layout=(1, 0, 2)), - TensorSpec(), - TensorSpec(), - TensorSpec(layout=(1, 0, 2)), - ), - allow_cuda_graph=True, - ) - alpha = jnp.ones((expert_count,), dtype=jnp.float32) - beta = jnp.ones((expert_count,), dtype=jnp.float32) - norm_const = jnp.ones((1,), dtype=jnp.float32) - return call( - a, - b, - c, - sfa.reshape(-1), - sfb.reshape(-1), - padded_offsets.astype(jnp.int32), - alpha, - beta, - prob.astype(jnp.float32), - norm_const, - ) - - -__all__ = [ - "grouped_gemm_dswiglu_mxfp8", - "grouped_gemm_swiglu_mxfp8", - "load_grouped_gemm_dswiglu_kernel", - "load_grouped_gemm_swiglu_kernel", - "pack_swiglu_pair", - "unpack_swiglu_pair", -] diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index 5529033937e..f14e2ac1057 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -91,9 +91,6 @@ def _cudnn_cutedsl_fusion_rejection_reasons( """Return reasons this call cannot use the CuTeDSL fusion, or an empty list.""" from transformer_engine_jax import get_device_compute_capability - from .cutedsl_extensions.moe import ( - load_grouped_gemm_dswiglu_kernel, - ) from .quantize import GroupedQuantizer, ScalingMode errors = [] @@ -150,28 +147,10 @@ def _cudnn_cutedsl_fusion_rejection_reasons( if num_local_experts > 1024: errors.append(f"requires at most 1024 local experts, got {num_local_experts}") - try: - import cutlass.jax as cutlass_jax - except ImportError as exc: - errors.append(f"nvidia-cutlass-dsl with JAX bindings is required: {exc}") - else: - try: - if not cutlass_jax.is_available(): - errors.append( - "cutlass.jax.is_available() is false; check cute_dsl_runtime.so discovery" - ) - except (AttributeError, RuntimeError) as exc: - errors.append(f"CUTLASS JAX runtime check failed: {exc}") - dependencies_available, dependency_error = tex.grouped_gemm_swiglu_dependencies_available() if not dependencies_available: errors.append(f"could not load the TVM-FFI CuTeDSL compiler/bridge: {dependency_error}") - try: - load_grouped_gemm_dswiglu_kernel() - except (ImportError, ModuleNotFoundError, RuntimeError) as exc: - errors.append(f"could not load the cuDNN frontend dSwiGLU backward kernel: {exc}") - return errors @@ -445,9 +424,7 @@ def _ffn_fwd_per_shard( # blocks. The existing TE grouped-GEMM path consumes the two full # projections concatenated along N. if use_cudnn_cutedsl_fusion: - from .cutedsl_extensions.moe import pack_swiglu_pair - - wi_combined = pack_swiglu_pair(wi_0, wi_1) + wi_combined = tex.pack_swiglu_pair(wi_0, wi_1) else: # Concat along the trailing axis (NOT stack on a new axis). # grouped_gemm requires the 3D (G, K, N) weight layout with @@ -465,8 +442,6 @@ def _ffn_fwd_per_shard( casted_wi = tex.grouped_quantize(wi_combined, fc1_quantizer_set.kernel, flatten_axis=-1) casted_intermediate = None if use_cudnn_cutedsl_fusion: - from .cutedsl_extensions.moe import unpack_swiglu_pair - casted_sorted_x_lhs = casted_sorted_x.get_tensor(usage=TensorUsage.LHS) casted_wi_rhs = casted_wi.get_tensor(usage=TensorUsage.RHS) padded_offsets = jnp.cumsum(local_group_sizes, dtype=jnp.int32) @@ -497,7 +472,7 @@ def _ffn_fwd_per_shard( output_dtype=fc2_quantizer_set.x.q_dtype, ) combined_out = combined_out_3d.reshape(sorted_x.shape[0], wi_combined.shape[-1]) - gate_proj_out, up_proj_out = unpack_swiglu_pair(combined_out) + gate_proj_out, up_proj_out = tex.unpack_swiglu_pair(combined_out) intermediate_shape = (sorted_x.shape[0], gate_proj_out.shape[-1]) scaling_mode = fc2_quantizer_set.x.scaling_mode @@ -702,12 +677,7 @@ def _ffn_bwd_per_shard( d_wo_bias = tex.grouped_dbias(d_eo_2d, local_group_sizes) if has_bias else None if use_cudnn_cutedsl_fusion: - from .cutedsl_extensions.moe import ( - grouped_gemm_dswiglu_mxfp8, - pack_swiglu_pair, - ) - - packed_forward = pack_swiglu_pair(gate_proj_out, up_proj_out) + packed_forward = tex.pack_swiglu_pair(gate_proj_out, up_proj_out) padded_offsets = jnp.cumsum(local_group_sizes, dtype=jnp.int32) prob = ( recv_w_flat[:, None, None] @@ -720,11 +690,9 @@ def _ffn_bwd_per_shard( d_combined_scale_row, d_combined_scale_col, _dprob, - ) = grouped_gemm_dswiglu_mxfp8( + ) = tex.grouped_gemm_dswiglu( _casted_d_eo_lhs.data.reshape(recv_rows, hidden, 1), - casted_wo_rhs_trans.data.reshape( - num_local_experts, intermediate_size, hidden - ).transpose(1, 2, 0), + casted_wo_rhs_trans.data.reshape(num_local_experts, intermediate_size, hidden), packed_forward.reshape(recv_rows, 2 * intermediate_size, 1), _casted_d_eo_lhs.scale_inv, casted_wo_rhs_trans.scale_inv, @@ -772,7 +740,7 @@ def _ffn_bwd_per_shard( ) if apply_topk_weights_early: # The cuDNN frontend dSwiGLU kernel accumulates dprob with - # atomics and expects a zero-initialized output buffer. cutlass_call + # atomics and expects a zero-initialized output buffer. XLA FFI # allocates custom-call outputs uninitialized, so compute this # routing-weight cotangent explicitly until we can pass dprob as an # initialized input/output buffer. @@ -835,9 +803,7 @@ def _ffn_bwd_per_shard( ) d_wi_combined = jnp.where(wgrad_group_active, d_wi_combined, jnp.zeros_like(d_wi_combined)) if use_cudnn_cutedsl_fusion: - from .cutedsl_extensions.moe import unpack_swiglu_pair - - d_wi_0, d_wi_1 = unpack_swiglu_pair(d_wi_combined) + d_wi_0, d_wi_1 = tex.unpack_swiglu_pair(d_wi_combined) else: d_wi_0, d_wi_1 = jnp.split(d_wi_combined, 2, axis=-1) if has_bias: From d3199ae4c503bdf7ade3f98f01dd4ec6a0478e49 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Tue, 15 Sep 2026 13:20:35 -0700 Subject: [PATCH 05/38] Align external MoE bootstrap with cuDNN fusion --- transformer_engine/jax/moe.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index feca726e7ca..1322bd397f5 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -150,15 +150,22 @@ def get_moe_recv_capacity_per_rank( max_tokens_per_rank: int, ep_size: int, recv_capacity_factor: Optional[float] = None, - alignment: int = _ALIGN_SIZE, + alignment: Optional[int] = None, ) -> int: """Return the aligned receive capacity for one EP rank. ``recv_capacity_factor=None`` reserves the dropless worst case. A finite factor >= 1 scales the capacity needed by perfectly balanced routing and is capped at the worst case. The balanced baseline includes the independent - per-local-expert alignment required by NCCL EP. + per-local-expert alignment required by NCCL EP. When ``alignment`` is not + supplied, it follows the active MoE implementation: 256 for the cuDNN + grouped-SwiGLU fusion and 128 for the regular TE grouped GEMM. This keeps + eager bootstrap callers in sync with the later compiled ``moe()`` call. """ + if alignment is None: + alignment = ( + _CUDNN_JAX_ALIGN_SIZE if _use_cudnn_cutedsl_fusion_from_env() else _ALIGN_SIZE + ) if num_experts <= 0 or num_experts_per_tok <= 0 or max_tokens_per_rank <= 0: raise ValueError( "num_experts, num_experts_per_tok, and max_tokens_per_rank must be positive" From 4f278d89a6fc3f30d5cfe42d9b41e6ddf8051e9d Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Tue, 15 Sep 2026 14:02:46 -0700 Subject: [PATCH 06/38] Add temporary cuDNN grouped GEMM fusion flag --- tests/jax/run_te_ep_moe.sh | 8 ++++---- tests/jax/test_te_ep_moe.py | 5 ++++- transformer_engine/jax/moe.py | 4 ++-- 3 files changed, 10 insertions(+), 7 deletions(-) diff --git a/tests/jax/run_te_ep_moe.sh b/tests/jax/run_te_ep_moe.sh index ca625fc705b..2445cee359e 100755 --- a/tests/jax/run_te_ep_moe.sh +++ b/tests/jax/run_te_ep_moe.sh @@ -78,7 +78,7 @@ run_phase() { echo echo "============================================================" echo "Phase: $phase_name" - echo " NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION=$fusion_env" + echo " NVTE_JAX_TEMP_FLAG_FOR_ABHINAV_CUDNN_GROUPED_GEMM_FUSION=$fusion_env" echo " phase pytest args : ${phase_args[*]:-}" echo " logs : $phase_log_dir" echo "============================================================" @@ -96,10 +96,10 @@ run_phase() { ) if [ "$i" -eq 0 ]; then echo "=== Live output from process 0 ($phase_name) ===" - env NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION="$fusion_env" \ + env NVTE_JAX_TEMP_FLAG_FOR_ABHINAV_CUDNN_GROUPED_GEMM_FUSION="$fusion_env" \ "${pytest_cmd[@]}" 2>&1 | tee "$log_file" & else - env NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION="$fusion_env" \ + env NVTE_JAX_TEMP_FLAG_FOR_ABHINAV_CUDNN_GROUPED_GEMM_FUSION="$fusion_env" \ "${pytest_cmd[@]}" > "$log_file" 2>&1 & fi PIDS+=("$!") @@ -136,7 +136,7 @@ run_phase() { fi } -if [ "${NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION:-0}" = "1" ]; then +if [ "${NVTE_JAX_TEMP_FLAG_FOR_ABHINAV_CUDNN_GROUPED_GEMM_FUSION:-0}" = "1" ]; then # Keep ordinary CUDA C++ and cuDNN JAX coverage in separate Python process # groups. TE EP/NCCL caches layer alignment process-wide, so 128-token and # 256-token dispatch-alignment tests cannot safely share one interpreter. diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index 40a032bd542..d6078bf448f 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -704,7 +704,10 @@ class TestTeEpMoeCudnnCutedslFusion: @pytest.mark.parametrize("apply_topk_weights_early", [False, True]) def test_mxfp8_forward_and_backward(self, mesh, apply_topk_weights_early): if not _use_cudnn_cutedsl_fusion_from_env(): - pytest.skip("run separately with NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION=1") + pytest.skip( + "run separately with " + "NVTE_JAX_TEMP_FLAG_FOR_ABHINAV_CUDNN_GROUPED_GEMM_FUSION=1" + ) block = _make_block( apply_topk_weights_early=apply_topk_weights_early, quantization_recipe=MXFP8BlockScaling(), diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index 1322bd397f5..e8080639649 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -65,7 +65,7 @@ # same 128-token tile, so a single constant covers every supported path. _ALIGN_SIZE = 128 _CUDNN_JAX_ALIGN_SIZE = 256 -_CUDNN_JAX_ENV = "NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION" +_CUDNN_JAX_ENV = "NVTE_JAX_TEMP_FLAG_FOR_ABHINAV_CUDNN_GROUPED_GEMM_FUSION" def _use_cudnn_cutedsl_fusion_from_env() -> bool: @@ -1459,7 +1459,7 @@ def moe( at 128 tokens (``_ALIGN_SIZE``); see that constant's docstring for rationale and how to extend if a future recipe needs >128. - Set ``NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION=1`` to use cuDNN's + Set ``NVTE_JAX_TEMP_FLAG_FOR_ABHINAV_CUDNN_GROUPED_GEMM_FUSION=1`` to use cuDNN's dedicated JAX grouped MXFP8 GEMM + SwiGLU API for eligible SM100 calls. The fused path uses 256-token expert alignment. Ineligible calls warn and fall back to TE's regular grouped-GEMM implementation. From 4f71efd9c0caf9f5bb5d2583c81201feaade064a Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Tue, 15 Sep 2026 15:35:05 -0700 Subject: [PATCH 07/38] Add cuDNN grouped dSwiGLU backward fusion --- tests/jax/test_te_ep_moe.py | 38 ++ .../jax/cpp_extensions/grouped_gemm_swiglu.py | 71 +++- transformer_engine/jax/moe.py | 332 +++++++++++++----- 3 files changed, 340 insertions(+), 101 deletions(-) diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index d6078bf448f..e89e0733fd8 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -731,6 +731,44 @@ def test_mxfp8_forward_and_backward(self, mesh, apply_topk_weights_early): assert np.all(np.isfinite(grad_x_np)) assert np.any(grad_x_np != 0) + params_np = _params_global_numpy(variables, mesh) + x_np = np.asarray(jax.device_get(x)) + + def loss_fn(params, inputs): + out, _ = _pure_jax_moe_reference( + inputs, + params["gate_kernel"], + params["wi"], + params["wo"], + num_experts=NUM_EXPERTS, + num_experts_per_tok=TOPK, + score_function="softmax", + expert_bias=None, + ) + return jnp.mean(out.astype(jnp.float32) ** 2) + + grads_ref, grad_x_ref = jax.jit(jax.grad(loss_fn, argnums=(0, 1)))( + {k: jnp.asarray(v) for k, v in params_np.items() if k != "expert_bias"}, + jnp.asarray(x_np), + ) + for name in ("gate_kernel", "wi", "wo"): + np.testing.assert_allclose( + _to_global_numpy(_unwrap(grads["params"][name]), mesh).astype(np.float32), + np.asarray(jax.device_get(grads_ref[name])).astype(np.float32), + **( + GRAD_GATE_TOLERANCE["mxfp8"] + if name == "gate_kernel" + else GRAD_FFN_TOLERANCE["mxfp8"] + ), + err_msg=f"{name} fused MXFP8 gradient parity breach", + ) + np.testing.assert_allclose( + grad_x_np, + np.asarray(jax.device_get(grad_x_ref)).astype(np.float32), + **GRAD_FFN_TOLERANCE["mxfp8"], + err_msg="d_x fused MXFP8 gradient parity breach", + ) + class TestTeEpMoeAuxLoss: """Aux-loss path. Consolidated into: diff --git a/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py b/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py index e2e8fe58ef8..c67e6323b8a 100644 --- a/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py +++ b/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py @@ -9,6 +9,7 @@ import jax.numpy as jnp __all__ = [ + "grouped_gemm_dswiglu", "grouped_gemm_swiglu", "grouped_gemm_swiglu_dependencies_available", "pack_swiglu_pair", @@ -21,7 +22,9 @@ def pack_swiglu_pair(gate: jax.Array, up: jax.Array) -> jax.Array: if gate.shape != up.shape: raise ValueError(f"gate shape {gate.shape} must match up shape {up.shape}") if gate.shape[-1] % 32: - raise ValueError(f"SwiGLU intermediate dimension {gate.shape[-1]} must be divisible by 32") + raise ValueError( + f"SwiGLU intermediate dimension {gate.shape[-1]} must be divisible by 32" + ) blocks = gate.shape[-1] // 32 return jnp.stack( ( @@ -70,7 +73,10 @@ def grouped_gemm_swiglu_dependencies_available() -> tuple[bool, str]: """Check the public cuDNN JAX API without compiling a kernel.""" try: import cutlass.jax - from cudnn import grouped_gemm_swiglu_jax_sm100 # noqa: F401 + from cudnn import ( # noqa: F401 + grouped_gemm_dswiglu_jax_sm100, + grouped_gemm_swiglu_jax_sm100, + ) if not cutlass.jax.is_available(): return False, "CuTeDSL JAX support is unavailable" @@ -110,18 +116,16 @@ def grouped_gemm_swiglu( alpha = jnp.ones((experts,), dtype=jnp.float32) norm_const = jnp.ones((1,), dtype=jnp.float32) result = grouped_gemm_swiglu_jax_sm100( - a_tensor=a, + a_tensor=a.reshape(rows, hidden), b_tensor=b, sfa_tensor=_compact_sf(sfa, _sf_atom_shape(1, rows, hidden), "sfa"), sfb_tensor=_compact_sf(sfb, _sf_atom_shape(experts, combined, hidden), "sfb"), padded_offsets=padded_offsets.astype(jnp.int32), alpha_tensor=alpha, - prob_tensor=prob.astype(jnp.float32), + prob_tensor=prob.reshape(rows).astype(jnp.float32), norm_const_tensor=norm_const, c_dtype=jnp.dtype(compute_dtype), d_dtype=jnp.dtype(output_dtype), - sf_vec_size=32, - discrete_col_sfd=True, ) return ( result["c_tensor"], @@ -130,3 +134,58 @@ def grouped_gemm_swiglu( result["sfd_row_tensor"].reshape(-1), result["sfd_col_tensor"].reshape(-1), ) + + +def grouped_gemm_dswiglu( + a: jax.Array, + b: jax.Array, + c: jax.Array, + sfa: jax.Array, + sfb: jax.Array, + padded_offsets: jax.Array, + prob: jax.Array, + *, + output_dtype, +) -> tuple[jax.Array, jax.Array, jax.Array, jax.Array, jax.Array]: + """Run cuDNN's grouped MXFP8 GEMM + dSwiGLU + MXFP8 quantization API.""" + if a.ndim != 2: + raise ValueError(f"Expected A[M,K], got {a.shape}") + if b.ndim != 3: + raise ValueError(f"Expected physical B[E,N,K], got {b.shape}") + if c.ndim != 2: + raise ValueError(f"Expected C[M,2N], got {c.shape}") + + from cudnn import grouped_gemm_dswiglu_jax_sm100 + + rows, hidden = a.shape + experts, intermediate, b_hidden = b.shape + if hidden != b_hidden: + raise ValueError(f"A K={hidden} does not match B K={b_hidden}") + if c.shape != (rows, 2 * intermediate): + raise ValueError(f"Expected C shape {(rows, 2 * intermediate)}, got {c.shape}") + + alpha = jnp.ones((experts,), dtype=jnp.float32) + beta = jnp.ones((experts,), dtype=jnp.float32) + norm_const = jnp.ones((1,), dtype=jnp.float32) + result = grouped_gemm_dswiglu_jax_sm100( + a_tensor=a, + b_tensor=b, + c_tensor=c, + sfa_tensor=_compact_sf(sfa, _sf_atom_shape(1, rows, hidden), "sfa"), + sfb_tensor=_compact_sf( + sfb, _sf_atom_shape(experts, intermediate, hidden), "sfb" + ), + padded_offsets=padded_offsets.astype(jnp.int32), + alpha_tensor=alpha, + beta_tensor=beta, + prob_tensor=prob.reshape(rows).astype(jnp.float32), + norm_const_tensor=norm_const, + d_dtype=jnp.dtype(output_dtype), + ) + return ( + result["d_row_tensor"], + result["d_col_tensor"], + result["dprob_tensor"], + result["sfd_row_tensor"].reshape(-1), + result["sfd_col_tensor"].reshape(-1), + ) diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index e8080639649..804b1b28e4d 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -102,7 +102,9 @@ def _cudnn_jax_fusion_rejection_reasons( if wi_0_bias is not None or wi_1_bias is not None: errors.append("does not support FC1 gate/up bias") if wi.ndim != 3 or wi.shape[-1] % 64: - errors.append(f"requires rank-3 wi with a 64-aligned gated dimension, got {wi.shape}") + errors.append( + f"requires rank-3 wi with a 64-aligned gated dimension, got {wi.shape}" + ) if x.dtype not in (jnp.bfloat16, jnp.float16): errors.append(f"requires BF16 or FP16 activations, got {x.dtype}") @@ -125,21 +127,31 @@ def _cudnn_jax_fusion_rejection_reasons( if all(isinstance(q, GroupedQuantizer) for q in required_quantizers.values()): if fc1_quantizer_set.x.q_dtype != fc1_quantizer_set.kernel.q_dtype: - errors.append("requires identical FC1 activation and weight MXFP8 payload dtypes") + errors.append( + "requires identical FC1 activation and weight MXFP8 payload dtypes" + ) supported = (jnp.float8_e4m3fn, jnp.float8_e5m2) for name, quantizer in required_quantizers.items(): if quantizer.q_dtype not in supported: - errors.append(f"unsupported MXFP8 payload dtype {quantizer.q_dtype} for {name}") + errors.append( + f"unsupported MXFP8 payload dtype {quantizer.q_dtype} for {name}" + ) mesh = _get_mesh() if mesh is not None and not mesh.empty and ep_axis in mesh.shape: num_local_experts = num_experts // mesh.shape[ep_axis] if num_local_experts > 1024: - errors.append(f"requires at most 1024 local experts, got {num_local_experts}") + errors.append( + f"requires at most 1024 local experts, got {num_local_experts}" + ) - dependencies_available, dependency_error = tex.grouped_gemm_swiglu_dependencies_available() + dependencies_available, dependency_error = ( + tex.grouped_gemm_swiglu_dependencies_available() + ) if not dependencies_available: - errors.append(f"could not load cuDNN's grouped SwiGLU JAX API: {dependency_error}") + errors.append( + f"could not load cuDNN's grouped SwiGLU JAX API: {dependency_error}" + ) return errors @@ -164,14 +176,18 @@ def get_moe_recv_capacity_per_rank( """ if alignment is None: alignment = ( - _CUDNN_JAX_ALIGN_SIZE if _use_cudnn_cutedsl_fusion_from_env() else _ALIGN_SIZE + _CUDNN_JAX_ALIGN_SIZE + if _use_cudnn_cutedsl_fusion_from_env() + else _ALIGN_SIZE ) if num_experts <= 0 or num_experts_per_tok <= 0 or max_tokens_per_rank <= 0: raise ValueError( "num_experts, num_experts_per_tok, and max_tokens_per_rank must be positive" ) if ep_size <= 0 or num_experts % ep_size != 0: - raise ValueError(f"num_experts={num_experts} must be divisible by ep_size={ep_size}") + raise ValueError( + f"num_experts={num_experts} must be divisible by ep_size={ep_size}" + ) if alignment <= 0: raise ValueError(f"alignment must be positive, got {alignment}") if recv_capacity_factor is not None: @@ -184,12 +200,18 @@ def get_moe_recv_capacity_per_rank( num_local_experts = num_experts // ep_size tokens_per_ep_group = ep_size * max_tokens_per_rank - max_local_assignments = tokens_per_ep_group * min(num_experts_per_tok, num_local_experts) + max_local_assignments = tokens_per_ep_group * min( + num_experts_per_tok, num_local_experts + ) max_nonempty_experts = min(num_local_experts, max_local_assignments) padded_total_bound = max_local_assignments + (alignment - 1) * max_nonempty_experts - aligned_total_bound = ((padded_total_bound + alignment - 1) // alignment) * alignment + aligned_total_bound = ( + (padded_total_bound + alignment - 1) // alignment + ) * alignment per_expert_bound = ( - num_local_experts * ((tokens_per_ep_group + alignment - 1) // alignment) * alignment + num_local_experts + * ((tokens_per_ep_group + alignment - 1) // alignment) + * alignment ) worst_case = min(per_expert_bound, aligned_total_bound) if recv_capacity_factor is None: @@ -199,7 +221,9 @@ def get_moe_recv_capacity_per_rank( max_tokens_per_rank * num_experts_per_tok + num_local_experts - 1 ) // num_local_experts balanced_aligned = ( - num_local_experts * ((balanced_per_expert + alignment - 1) // alignment) * alignment + num_local_experts + * ((balanced_per_expert + alignment - 1) // alignment) + * alignment ) requested = math.ceil(balanced_aligned * recv_capacity_factor) requested = ((requested + alignment - 1) // alignment) * alignment @@ -234,10 +258,14 @@ def _constraint(y): return jax.lax.with_sharding_constraint(y, sharding) def _constraint_fwd(y): - return jax.lax.with_sharding_constraint(y, sharding), jnp.zeros((), dtype=y.dtype) + return jax.lax.with_sharding_constraint(y, sharding), jnp.zeros( + (), dtype=y.dtype + ) def _constraint_bwd(dtype_ref, grad): - return (jax.lax.with_sharding_constraint(grad.astype(dtype_ref.dtype), sharding),) + return ( + jax.lax.with_sharding_constraint(grad.astype(dtype_ref.dtype), sharding), + ) _constraint.defvjp(_constraint_fwd, _constraint_bwd) return _constraint(x) @@ -295,7 +323,9 @@ def _te_ep_assert_compatible_bootstrap( " transformer_engine.jax.moe.record_ep_bootstrap_signature_for_moe(...)" " with the same params, before invoking moe()." ) - b_num_experts, b_max_tpr, b_recv_pr, b_hidden, b_ep_size = _te_ep_bootstrap_signature + b_num_experts, b_max_tpr, b_recv_pr, b_hidden, b_ep_size = ( + _te_ep_bootstrap_signature + ) if ( num_experts != b_num_experts or hidden_dim != b_hidden @@ -374,7 +404,9 @@ def _validate_moe_quantizer_sets( supports only no-op quantizers and stateless MXFP8 grouped quantizers. """ if not isinstance(quantizer_sets, tuple) or len(quantizer_sets) != 2: - raise TypeError("MoE quantizer_sets must be a tuple of FC1 and FC2 QuantizerSet objects.") + raise TypeError( + "MoE quantizer_sets must be a tuple of FC1 and FC2 QuantizerSet objects." + ) expected_groups = { "x": num_token_groups, @@ -455,7 +487,9 @@ def _ffn_fwd_per_shard( else: wi_for_gemm = wi wi_combined_bias = ( - jnp.concatenate([wi_0_bias, wi_1_bias], axis=-1) if wi_0_bias is not None else None + jnp.concatenate([wi_0_bias, wi_1_bias], axis=-1) + if wi_0_bias is not None + else None ) fc1_quantizer_set, fc2_quantizer_set = quantizer_sets @@ -465,7 +499,9 @@ def _ffn_fwd_per_shard( group_sizes, flatten_axis=-1, ) - casted_wi = tex.grouped_quantize(wi_for_gemm, fc1_quantizer_set.kernel, flatten_axis=-1) + casted_wi = tex.grouped_quantize( + wi_for_gemm, fc1_quantizer_set.kernel, flatten_axis=-1 + ) casted_intermediate = None if use_cudnn_jax_fusion: casted_sorted_x_lhs = casted_sorted_x.get_tensor(usage=TensorUsage.LHS) @@ -484,9 +520,9 @@ def _ffn_fwd_per_shard( intermediate_scale_col, ) = tex.grouped_gemm_swiglu( casted_sorted_x_lhs.data.reshape(sorted_x.shape[0], hidden, 1), - casted_wi_rhs.data.reshape(num_local_experts, hidden, wi_for_gemm.shape[-1]).transpose( - 0, 2, 1 - ), + casted_wi_rhs.data.reshape( + num_local_experts, hidden, wi_for_gemm.shape[-1] + ).transpose(0, 2, 1), casted_sorted_x_lhs.scale_inv, casted_wi_rhs.scale_inv, padded_offsets, @@ -573,15 +609,25 @@ def _ffn_fwd_per_shard( contracting_dims=((1,), (1,)), bias=wo_bias, ) - expert_outputs_3d = expert_outputs.reshape(1, expert_outputs.shape[0], expert_outputs.shape[1]) + expert_outputs_3d = expert_outputs.reshape( + 1, expert_outputs.shape[0], expert_outputs.shape[1] + ) group_sizes_2d = group_sizes.reshape(1, num_local_experts) residuals = ( - casted_sorted_x.get_tensor(usage=TensorUsage.LHS_TRANS).checkpoint(fc1_quantizer_set.x), - casted_wi.get_tensor(usage=TensorUsage.RHS_TRANS).checkpoint(fc1_quantizer_set.kernel), - gate_proj_out, + casted_sorted_x.get_tensor(usage=TensorUsage.LHS_TRANS).checkpoint( + fc1_quantizer_set.x + ), + casted_wi.get_tensor(usage=TensorUsage.RHS_TRANS).checkpoint( + fc1_quantizer_set.kernel + ), + combined_out if use_cudnn_jax_fusion else gate_proj_out, up_proj_out, - casted_intermediate.get_tensor(usage=TensorUsage.LHS_TRANS).checkpoint(fc2_quantizer_set.x), - casted_wo.get_tensor(usage=TensorUsage.RHS_TRANS).checkpoint(fc2_quantizer_set.kernel), + casted_intermediate.get_tensor(usage=TensorUsage.LHS_TRANS).checkpoint( + fc2_quantizer_set.x + ), + casted_wo.get_tensor(usage=TensorUsage.RHS_TRANS).checkpoint( + fc2_quantizer_set.kernel + ), group_sizes_2d, ) return expert_outputs_3d, residuals @@ -619,11 +665,6 @@ def _ffn_bwd_per_shard( ) _casted_d_eo_lhs = casted_d_eo.get_tensor(usage=TensorUsage.LHS) _casted_d_eo_rhs = casted_d_eo.get_tensor(usage=TensorUsage.RHS) - d_intermediate = tex.grouped_gemm( - _casted_d_eo_lhs, - casted_wo_rhs_trans, - contracting_dims=((1,), (2,)), - ) d_wo = tex.grouped_gemm( casted_intermediate_lhs_trans, _casted_d_eo_rhs, @@ -631,51 +672,111 @@ def _ffn_bwd_per_shard( ) d_wo_bias = tex.grouped_dbias(d_eo_2d, group_sizes) if has_bias else None - act_fn = _convert_to_activation_function(activation_type) - if apply_topk_weights_early: - # intermediate' = intermediate * w. - # Masking is not required as: - # 1. Padding between groups is zero padded due to NCCL EP. - # 2. Overallocated padding past all groups is uninitialized, but subsequent GEMMs and EP are all group-size aware and will not read past the final group. - w_b = recv_w_flat[:, None].astype(d_intermediate.dtype) - gate_proj_for_bwd = gate_proj_out - up_proj_for_bwd = up_proj_out - intermediate_unweighted = act_fn(gate_proj_out) * up_proj_out - d_recv_w_from_intermediate = jnp.sum( - d_intermediate * intermediate_unweighted, - axis=-1, - ).astype(recv_w_flat.dtype) - d_intermediate = d_intermediate * w_b + if use_cudnn_jax_fusion: + # The forward residual uses this slot for cuDNN's interleaved pre-activation C. + combined_out = gate_proj_out + rows, combined = combined_out.shape + intermediate = combined // 2 + hidden = d_eo_2d.shape[-1] + num_local_experts = group_sizes.size + padded_offsets = jnp.cumsum(group_sizes, dtype=jnp.int32) + prob = ( + recv_w_flat + if apply_topk_weights_early + else jnp.ones((rows,), dtype=jnp.float32) + ) + ( + d_combined_row, + d_combined_col, + dprob, + d_combined_scale_row, + d_combined_scale_col, + ) = tex.grouped_gemm_dswiglu( + _casted_d_eo_lhs.data.reshape(rows, hidden), + casted_wo_rhs_trans.data.reshape(num_local_experts, intermediate, hidden), + combined_out, + _casted_d_eo_lhs.scale_inv, + casted_wo_rhs_trans.scale_inv, + padded_offsets, + prob, + output_dtype=fc1_quantizer_set.dgrad.q_dtype, + ) + d_recv_w_from_intermediate = ( + dprob.astype(recv_w_flat.dtype) + if apply_topk_weights_early + else jnp.zeros_like(recv_w_flat) + ) + + combined_shape = (rows, combined) + scaling_mode = fc1_quantizer_set.dgrad.scaling_mode + row_scale_size = scaling_mode.get_grouped_scale_shape( + combined_shape, + num_local_experts, + False, + is_padded=True, + flatten_axis=1, + )[0] + col_scale_size = scaling_mode.get_grouped_scale_shape( + combined_shape, + num_local_experts, + True, + is_padded=True, + flatten_axis=1, + )[0] + d_combined_scale_row = jnp.pad( + d_combined_scale_row, + (0, row_scale_size - d_combined_scale_row.size), + ) + d_combined_scale_col = jnp.pad( + d_combined_scale_col, + (0, col_scale_size - d_combined_scale_col.size), + ) + casted_d_combined = ScaledTensorFactory.create( + data=d_combined_row.reshape(-1), + scale_inv=d_combined_scale_row, + colwise_data=d_combined_col.reshape(-1), + colwise_scale_inv=d_combined_scale_col, + scaling_mode=scaling_mode, + dq_dtype=d_eo_2d.dtype, + data_layout=fc1_quantizer_set.dgrad.data_layout, + q_layout=fc1_quantizer_set.dgrad.q_layout, + flatten_axis=1, + first_dims=group_sizes, + original_shape=combined_shape, + pre_swizzled=True, + ) + d_combined_for_bias = None else: - gate_proj_for_bwd = gate_proj_out - up_proj_for_bwd = up_proj_out - d_recv_w_from_intermediate = jnp.zeros_like(recv_w_flat) - - # Activation bwd, symmetric with the fwd: silu' and the two - # elementwise products run in the GEMM dtype (no fp32 island), so - # the chain rule composes through at the same precision the wi/wo - # GEMMs consume. - act_gp, dact_pullback = jax.vjp(act_fn, gate_proj_for_bwd) - d_up_proj_out = d_intermediate * act_gp - (d_gate_proj_out,) = dact_pullback(d_intermediate * up_proj_for_bwd) - - # wi bwd (fused gate/up via concat). Mirror the fused fwd: pack the - # gate/up cotangents along the trailing axis, run a single - # grouped_quantize + two grouped_gemm pair (one dgrad, one wgrad) - # against the fused casted_wi_rhs_trans residual, then split the - # wgrad result remains in the contiguous gated-SwiGLU ``wi`` layout. - d_combined_for_bias = jnp.concatenate([d_gate_proj_out, d_up_proj_out], axis=-1) - d_combined = ( - tex.pack_swiglu_pair(d_gate_proj_out, d_up_proj_out) - if use_cudnn_jax_fusion - else d_combined_for_bias - ) - casted_d_combined = tex.grouped_quantize( - d_combined, - fc1_quantizer_set.dgrad, - group_sizes, - flatten_axis=-1, - ) + d_intermediate = tex.grouped_gemm( + _casted_d_eo_lhs, + casted_wo_rhs_trans, + contracting_dims=((1,), (2,)), + ) + act_fn = _convert_to_activation_function(activation_type) + if apply_topk_weights_early: + # intermediate' = intermediate * w. + # Masking is not required as grouped GEMMs consume only group-size rows. + w_b = recv_w_flat[:, None].astype(d_intermediate.dtype) + intermediate_unweighted = act_fn(gate_proj_out) * up_proj_out + d_recv_w_from_intermediate = jnp.sum( + d_intermediate * intermediate_unweighted, + axis=-1, + ).astype(recv_w_flat.dtype) + d_intermediate = d_intermediate * w_b + else: + d_recv_w_from_intermediate = jnp.zeros_like(recv_w_flat) + + # Activation bwd stays in the GEMM dtype, matching the forward path. + act_gp, dact_pullback = jax.vjp(act_fn, gate_proj_out) + d_up_proj_out = d_intermediate * act_gp + (d_gate_proj_out,) = dact_pullback(d_intermediate * up_proj_out) + d_combined_for_bias = jnp.concatenate([d_gate_proj_out, d_up_proj_out], axis=-1) + casted_d_combined = tex.grouped_quantize( + d_combined_for_bias, + fc1_quantizer_set.dgrad, + group_sizes, + flatten_axis=-1, + ) d_sorted_x = tex.grouped_gemm( casted_d_combined.get_tensor(usage=TensorUsage.LHS), casted_wi_rhs_trans, @@ -761,7 +862,9 @@ def _moe_fwd_rule( raise ValueError("moe(...) requires ep_axis to be set (TE EP backend).") num_ep = mesh.shape[ep_axis] if num_experts % num_ep != 0: - raise ValueError(f"num_experts={num_experts} must be divisible by EP size={num_ep}") + raise ValueError( + f"num_experts={num_experts} must be divisible by EP size={num_ep}" + ) num_local_experts = num_experts // num_ep dp_size = 1 @@ -835,7 +938,9 @@ def _moe_fwd_rule( # ---------------- Routing (global view) ---------------- # expert_bias is an empty (shape-(0,)) sentinel when the caller did # not enable it; the primitive treats that as "no bias". - eb_arg = expert_bias if expert_bias.shape != (0,) else jnp.zeros((0,), dtype=jnp.float32) + eb_arg = ( + expert_bias if expert_bias.shape != (0,) else jnp.zeros((0,), dtype=jnp.float32) + ) sparse_probs, routing_map, saved_scores = tex.fused_topk_with_score_function_fwd( logits_2d, topk=K, @@ -857,7 +962,9 @@ def _moe_fwd_rule( # single all-gather over (*dp, ep) and lives off the dispatch # critical path. if aux_loss_coeff > 0.0: - global_logits_2d = jax.lax.with_sharding_constraint(logits_2d, NamedSharding(mesh, P())) + global_logits_2d = jax.lax.with_sharding_constraint( + logits_2d, NamedSharding(mesh, P()) + ) _, global_routing_map, _ = tex.fused_topk_with_score_function_fwd( global_logits_2d, topk=K, @@ -908,8 +1015,12 @@ def _moe_fwd_rule( # each rank see B/ep rows (not B/num_procs) and overrun the bootstrap-sized # send buffer. Pin both routing tensors to the (outer, ep) leading sharding # so per-rank token counts match max_tokens_per_rank. - topk_idx_3d = jax.lax.with_sharding_constraint(topk_idx_3d, NamedSharding(mesh, ep3_spec)) - topk_w_3d = jax.lax.with_sharding_constraint(topk_w_3d, NamedSharding(mesh, ep3_spec)) + topk_idx_3d = jax.lax.with_sharding_constraint( + topk_idx_3d, NamedSharding(mesh, ep3_spec) + ) + topk_w_3d = jax.lax.with_sharding_constraint( + topk_w_3d, NamedSharding(mesh, ep3_spec) + ) # ---------------- TE EP dispatch (global view) ---------------- cfg = tex.EpLayerConfig( @@ -917,11 +1028,15 @@ def _moe_fwd_rule( dispatch_output_per_expert_alignment=dispatch_alignment, ) token_counts, total_recv_tokens, handle_mem = tex.ep_prepare(cfg, topk_idx_3d) - token_counts = jax.lax.with_sharding_constraint(token_counts, NamedSharding(mesh, ep2_spec)) + token_counts = jax.lax.with_sharding_constraint( + token_counts, NamedSharding(mesh, ep2_spec) + ) recv_tokens, recv_topk_weights = tex.ep_dispatch_fwd( cfg, handle_mem, topk_idx_3d, x, topk_w_3d, recv_pr ) - recv_tokens = jax.lax.with_sharding_constraint(recv_tokens, NamedSharding(mesh, ep3_spec)) + recv_tokens = jax.lax.with_sharding_constraint( + recv_tokens, NamedSharding(mesh, ep3_spec) + ) recv_topk_weights = jax.lax.with_sharding_constraint( recv_topk_weights, NamedSharding(mesh, ep2_spec) ) @@ -983,7 +1098,9 @@ def _ffn_fwd_body(*args): out_specs=(ep3_spec, residuals_spec), check_rep=False, )(*ffn_in_args) - expert_outputs = jax.lax.with_sharding_constraint(expert_outputs, NamedSharding(mesh, ep3_spec)) + expert_outputs = jax.lax.with_sharding_constraint( + expert_outputs, NamedSharding(mesh, ep3_spec) + ) # ---------------- TE EP combine (global view) ---------------- out_partition_spec = (batch_pspec_axis, None, None) @@ -1077,7 +1194,12 @@ def _moe_bwd_rule( cotangents, ): """Backward mirror of :func:`_moe_fwd_rule`.""" - del num_groups, group_topk, dtype, recv_capacity_per_rank # captured / unused in bwd + del ( + num_groups, + group_topk, + dtype, + recv_capacity_per_rank, + ) # captured / unused in bwd from jax.experimental.shard_map import shard_map # total_recv_tokens is a non-differentiable output; its cotangent is unused. @@ -1116,7 +1238,9 @@ def _moe_bwd_rule( w = ctx.recv_topk_weights[..., None].astype(grad_pre_combine.dtype) d_expert_outputs = grad_pre_combine * w d_recv_w_from_combine = (grad_pre_combine * ctx.expert_outputs).sum(axis=-1) - d_recv_w_from_combine = d_recv_w_from_combine.astype(ctx.recv_topk_weights.dtype) + d_recv_w_from_combine = d_recv_w_from_combine.astype( + ctx.recv_topk_weights.dtype + ) # ---------------- FFN bwd (per-shard via shard_map) ---------------- kernel_spec = P(ep_axis, None, None) @@ -1216,8 +1340,12 @@ def _ffn_bwd_body(*args): d_recv_w_total = d_recv_w_from_combine + d_recv_w_from_intermediate # ---------------- Dispatch bwd (global view) ---------------- - d_sorted_x = jax.lax.with_sharding_constraint(d_sorted_x, NamedSharding(mesh, ep3_spec)) - d_recv_w_total = jax.lax.with_sharding_constraint(d_recv_w_total, NamedSharding(mesh, ep2_spec)) + d_sorted_x = jax.lax.with_sharding_constraint( + d_sorted_x, NamedSharding(mesh, ep3_spec) + ) + d_recv_w_total = jax.lax.with_sharding_constraint( + d_recv_w_total, NamedSharding(mesh, ep2_spec) + ) d_x_from_dispatch, d_topk_w = tex.ep_dispatch_bwd( ctx.cfg, ctx.handle_mem, @@ -1264,7 +1392,9 @@ def _ffn_bwd_body(*args): ) # routing_map is ignored by the kernel when compute_aux_scores=True, # so pass a zero placeholder of the right shape/dtype. - zero_routing_map = jnp.zeros(ctx.aux_saved_scores.shape, dtype=ctx.routing_map.dtype) + zero_routing_map = jnp.zeros( + ctx.aux_saved_scores.shape, dtype=ctx.routing_map.dtype + ) d_logits_aux = tex.fused_topk_with_score_function_bwd( zero_routing_map, ctx.aux_saved_scores, @@ -1281,20 +1411,28 @@ def _ffn_bwd_body(*args): d_gate_logits = d_logits_2d.reshape(B, S, num_experts) gate_kernel_cast = ctx.gate_kernel.astype(ctx.x.dtype) d_x_from_gate = jnp.einsum("bse,he->bsh", d_gate_logits, gate_kernel_cast) - d_gate_kernel = jnp.einsum("bsh,bse->he", ctx.x, d_gate_logits).astype(ctx.gate_kernel.dtype) + d_gate_kernel = jnp.einsum("bsh,bse->he", ctx.x, d_gate_logits).astype( + ctx.gate_kernel.dtype + ) d_x = d_x_from_gate + d_x_from_dispatch # Pin output grads to the declared logical axes so downstream # optimizers see consistent shardings. d_x = with_sharding_constraint_by_logical_axes(d_x, input_axes) - d_gate_kernel = with_sharding_constraint_by_logical_axes(d_gate_kernel, gate_kernel_axes) + d_gate_kernel = with_sharding_constraint_by_logical_axes( + d_gate_kernel, gate_kernel_axes + ) d_wi = with_sharding_constraint_by_logical_axes(d_wi, wi_kernel_axes) d_wo = with_sharding_constraint_by_logical_axes(d_wo, wo_kernel_axes) if has_bias: wi_bias_axes = (wi_kernel_axes[0], *wi_kernel_axes[2:]) wo_bias_axes = (wo_kernel_axes[0], *wo_kernel_axes[2:]) - d_wi_0_bias = with_sharding_constraint_by_logical_axes(d_wi_0_bias, wi_bias_axes) - d_wi_1_bias = with_sharding_constraint_by_logical_axes(d_wi_1_bias, wi_bias_axes) + d_wi_0_bias = with_sharding_constraint_by_logical_axes( + d_wi_0_bias, wi_bias_axes + ) + d_wi_1_bias = with_sharding_constraint_by_logical_axes( + d_wi_1_bias, wi_bias_axes + ) d_wo_bias = with_sharding_constraint_by_logical_axes(d_wo_bias, wo_bias_axes) # expert_bias has no learnable bwd path through fused_topk: the @@ -1501,7 +1639,9 @@ def moe( mesh = _get_mesh() if mesh is None or mesh.empty: raise ValueError("moe(...) requires an active jax.sharding.Mesh.") - expected_leading: Any = (*data_parallelism_axes, ep_axis) if data_parallelism_axes else ep_axis + expected_leading: Any = ( + (*data_parallelism_axes, ep_axis) if data_parallelism_axes else ep_axis + ) expected_spec = P(expected_leading, None, None) actual_spec = getattr(getattr(x, "sharding", None), "spec", None) if actual_spec is not None and tuple(actual_spec) != tuple(expected_spec): @@ -1576,5 +1716,7 @@ def moe( ) if aux_loss_coeff <= 0.0: aux_loss = None - assert output.dtype == x.dtype, f"moe() output dtype {output.dtype} != input dtype {x.dtype}" + assert ( + output.dtype == x.dtype + ), f"moe() output dtype {output.dtype} != input dtype {x.dtype}" return output, aux_loss, total_recv_tokens From 98b88a229549aeba7ac374a74baefe210920dc9a Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Mon, 21 Sep 2026 17:44:39 -0700 Subject: [PATCH 08/38] Fix SM100 version check Signed-off-by: Jeremy Berchtold --- transformer_engine/jax/moe.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index 804b1b28e4d..34931126ddb 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -95,8 +95,8 @@ def _cudnn_jax_fusion_rejection_reasons( except RuntimeError as exc: errors.append(f"could not query GPU compute capability: {exc}") else: - if compute_capability != 100: - errors.append(f"requires an SM100 GPU, got SM{compute_capability}") + if compute_capability < 100: + errors.append(f"requires an SM100+ GPU, got SM{compute_capability}") if str(activation_type).lower() != "silu": errors.append("requires activation_type='silu'") if wi_0_bias is not None or wi_1_bias is not None: From a4e2c7ca1156196708e60e5cf6733d31a5cec686 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Tue, 22 Sep 2026 09:15:36 -0700 Subject: [PATCH 09/38] Update cuDNN grouped SwiGLU JAX API --- .../jax/cpp_extensions/grouped_gemm_swiglu.py | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py b/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py index c67e6323b8a..f683dd76f81 100644 --- a/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py +++ b/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py @@ -73,9 +73,9 @@ def grouped_gemm_swiglu_dependencies_available() -> tuple[bool, str]: """Check the public cuDNN JAX API without compiling a kernel.""" try: import cutlass.jax - from cudnn import ( # noqa: F401 - grouped_gemm_dswiglu_jax_sm100, - grouped_gemm_swiglu_jax_sm100, + from cudnn.jax import ( # noqa: F401 + grouped_gemm_dswiglu, + grouped_gemm_swiglu, ) if not cutlass.jax.is_available(): @@ -107,7 +107,7 @@ def grouped_gemm_swiglu( if b.ndim != 3: raise ValueError(f"Expected physical B[E,N,K], got {b.shape}") - from cudnn import grouped_gemm_swiglu_jax_sm100 + from cudnn.jax import grouped_gemm_swiglu as cudnn_grouped_gemm_swiglu rows, hidden, _ = a.shape experts, combined, b_hidden = b.shape @@ -115,7 +115,7 @@ def grouped_gemm_swiglu( raise ValueError(f"A K={hidden} does not match B K={b_hidden}") alpha = jnp.ones((experts,), dtype=jnp.float32) norm_const = jnp.ones((1,), dtype=jnp.float32) - result = grouped_gemm_swiglu_jax_sm100( + result = cudnn_grouped_gemm_swiglu( a_tensor=a.reshape(rows, hidden), b_tensor=b, sfa_tensor=_compact_sf(sfa, _sf_atom_shape(1, rows, hidden), "sfa"), @@ -155,7 +155,7 @@ def grouped_gemm_dswiglu( if c.ndim != 2: raise ValueError(f"Expected C[M,2N], got {c.shape}") - from cudnn import grouped_gemm_dswiglu_jax_sm100 + from cudnn.jax import grouped_gemm_dswiglu as cudnn_grouped_gemm_dswiglu rows, hidden = a.shape experts, intermediate, b_hidden = b.shape @@ -167,7 +167,7 @@ def grouped_gemm_dswiglu( alpha = jnp.ones((experts,), dtype=jnp.float32) beta = jnp.ones((experts,), dtype=jnp.float32) norm_const = jnp.ones((1,), dtype=jnp.float32) - result = grouped_gemm_dswiglu_jax_sm100( + result = cudnn_grouped_gemm_dswiglu( a_tensor=a, b_tensor=b, c_tensor=c, From abd0405ad8f702dd4ed1f173f3840f4298d1b5d5 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Mon, 28 Sep 2026 16:09:13 -0700 Subject: [PATCH 10/38] Route Rubin MoE grouped GLU through cuDNN JAX --- tests/jax/test_te_ep_moe.py | 53 ++++++++++++++- .../jax/cpp_extensions/grouped_gemm_swiglu.py | 66 +++++++++++++++++-- transformer_engine/jax/moe.py | 12 +++- 3 files changed, 122 insertions(+), 9 deletions(-) diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index e89e0733fd8..8e308cde5cd 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -699,21 +699,34 @@ def loss_fn(params, x): class TestTeEpMoeCudnnCutedslFusion: - """End-to-end MXFP8 coverage for cuDNN's dedicated grouped SwiGLU JAX API.""" + """End-to-end MXFP8 coverage for cuDNN's grouped GLU JAX APIs.""" @pytest.mark.parametrize("apply_topk_weights_early", [False, True]) - def test_mxfp8_forward_and_backward(self, mesh, apply_topk_weights_early): + def test_mxfp8_forward_and_backward(self, mesh, apply_topk_weights_early, monkeypatch): if not _use_cudnn_cutedsl_fusion_from_env(): pytest.skip( "run separately with " "NVTE_JAX_TEMP_FLAG_FOR_ABHINAV_CUDNN_GROUPED_GEMM_FUSION=1" ) + rubin_calls = [] + if get_device_compute_capability(0) == 107: + from transformer_engine.jax import cpp_extensions as tex + + original_glu = tex.grouped_gemm_glu + + def checked_glu(*args, **kwargs): + rubin_calls.append(True) + return original_glu(*args, **kwargs) + + monkeypatch.setattr(tex, "grouped_gemm_glu", checked_glu) block = _make_block( apply_topk_weights_early=apply_topk_weights_early, quantization_recipe=MXFP8BlockScaling(), ) x = _make_inputs(jax.random.PRNGKey(30)) variables, output, aux = _init_apply(block, mesh, x, jax.random.PRNGKey(31)) + if get_device_compute_capability(0) == 107: + assert rubin_calls, "Rubin MoE did not select the cuDNN GLU JAX API" grads, grad_x = _grad_step(block, variables, mesh, x) assert output.shape == x.shape @@ -731,6 +744,42 @@ def test_mxfp8_forward_and_backward(self, mesh, apply_topk_weights_early): assert np.all(np.isfinite(grad_x_np)) assert np.any(grad_x_np != 0) + if get_device_compute_capability(0) == 107: + # The dedicated SwiGLU path is the reference for this kernel + # substitution. Its MXFP8 gradients can differ from pure JAX by + # more than the strict unfused test threshold on Rubin. + monkeypatch.setattr(tex, "grouped_gemm_glu", tex.grouped_gemm_swiglu) + with _ctx(mesh): + x_sh = _shard_inputs(x, mesh) + baseline_output, _, _ = jax.jit(block.apply)(variables, x_sh) + baseline_output.block_until_ready() + baseline_grads, baseline_grad_x = _grad_step(block, variables, mesh, x) + np.testing.assert_allclose( + output_np, + _to_global_numpy(baseline_output, mesh).astype(np.float32), + **FWD_TOLERANCE["mxfp8"], + err_msg="Rubin GLU forward differs from dedicated SwiGLU", + ) + for name in ("gate_kernel", "wi", "wo"): + tolerance = ( + GRAD_GATE_TOLERANCE["mxfp8"] + if name == "gate_kernel" + else GRAD_FFN_TOLERANCE["mxfp8"] + ) + np.testing.assert_allclose( + _to_global_numpy(_unwrap(grads["params"][name]), mesh).astype(np.float32), + _to_global_numpy(_unwrap(baseline_grads["params"][name]), mesh).astype(np.float32), + **tolerance, + err_msg=f"Rubin GLU {name} gradient differs from dedicated SwiGLU", + ) + np.testing.assert_allclose( + grad_x_np, + _to_global_numpy(baseline_grad_x, mesh).astype(np.float32), + **GRAD_FFN_TOLERANCE["mxfp8"], + err_msg="Rubin GLU input gradient differs from dedicated SwiGLU", + ) + return + params_np = _params_global_numpy(variables, mesh) x_np = np.asarray(jax.device_get(x)) diff --git a/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py b/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py index f683dd76f81..15559420bed 100644 --- a/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py +++ b/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py @@ -10,6 +10,7 @@ __all__ = [ "grouped_gemm_dswiglu", + "grouped_gemm_glu", "grouped_gemm_swiglu", "grouped_gemm_swiglu_dependencies_available", "pack_swiglu_pair", @@ -69,7 +70,7 @@ def _compact_sf(scale: jax.Array, shape: tuple[int, ...], name: str) -> jax.Arra return scale.reshape(-1)[:size].reshape(shape) -def grouped_gemm_swiglu_dependencies_available() -> tuple[bool, str]: +def grouped_gemm_swiglu_dependencies_available(rubin: bool = False) -> tuple[bool, str]: """Check the public cuDNN JAX API without compiling a kernel.""" try: import cutlass.jax @@ -78,6 +79,9 @@ def grouped_gemm_swiglu_dependencies_available() -> tuple[bool, str]: grouped_gemm_swiglu, ) + if rubin: + from cudnn.jax import grouped_gemm_glu # noqa: F401 + if not cutlass.jax.is_available(): return False, "CuTeDSL JAX support is unavailable" except (ImportError, ModuleNotFoundError, RuntimeError, AttributeError) as exc: @@ -109,13 +113,65 @@ def grouped_gemm_swiglu( from cudnn.jax import grouped_gemm_swiglu as cudnn_grouped_gemm_swiglu + return _grouped_gemm_forward( + cudnn_grouped_gemm_swiglu, + a, + b, + sfa, + sfb, + padded_offsets, + prob, + compute_dtype=compute_dtype, + output_dtype=output_dtype, + ) + + +def grouped_gemm_glu( + a: jax.Array, + b: jax.Array, + sfa: jax.Array, + sfb: jax.Array, + padded_offsets: jax.Array, + prob: jax.Array, + *, + compute_dtype, + output_dtype, +) -> tuple[jax.Array, jax.Array, jax.Array, jax.Array, jax.Array]: + """Run the Rubin cuDNN grouped MXFP8 GEMM + SwiGLU kernel.""" + from cudnn.jax import grouped_gemm_glu as cudnn_grouped_gemm_glu + + return _grouped_gemm_forward( + cudnn_grouped_gemm_glu, + a, + b, + sfa, + sfb, + padded_offsets, + prob, + compute_dtype=compute_dtype, + output_dtype=output_dtype, + ) + + +def _grouped_gemm_forward( + cudnn_forward, + a, + b, + sfa, + sfb, + padded_offsets, + prob, + *, + compute_dtype, + output_dtype, +): rows, hidden, _ = a.shape experts, combined, b_hidden = b.shape if hidden != b_hidden: raise ValueError(f"A K={hidden} does not match B K={b_hidden}") alpha = jnp.ones((experts,), dtype=jnp.float32) norm_const = jnp.ones((1,), dtype=jnp.float32) - result = cudnn_grouped_gemm_swiglu( + result = cudnn_forward( a_tensor=a.reshape(rows, hidden), b_tensor=b, sfa_tensor=_compact_sf(sfa, _sf_atom_shape(1, rows, hidden), "sfa"), @@ -128,9 +184,9 @@ def grouped_gemm_swiglu( d_dtype=jnp.dtype(output_dtype), ) return ( - result["c_tensor"], - result["d_tensor"], - result["d_col_tensor"], + result["c_tensor"].reshape(rows, combined), + result["d_tensor"].reshape(rows, combined // 2), + result["d_col_tensor"].reshape(rows, combined // 2), result["sfd_row_tensor"].reshape(-1), result["sfd_col_tensor"].reshape(-1), ) diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index 34931126ddb..9affe7e2ce0 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -75,6 +75,13 @@ def _use_cudnn_cutedsl_fusion_from_env() -> bool: return value == "1" +def _is_rubin_device() -> bool: + """SM107 is reported as 107 by TransformerEngine's CUDA utility.""" + from transformer_engine_jax import get_device_compute_capability + + return get_device_compute_capability(0) == 107 + + def _cudnn_jax_fusion_rejection_reasons( x, wi, @@ -90,6 +97,7 @@ def _cudnn_jax_fusion_rejection_reasons( from transformer_engine_jax import get_device_compute_capability errors = [] + compute_capability = None try: compute_capability = get_device_compute_capability(0) except RuntimeError as exc: @@ -146,7 +154,7 @@ def _cudnn_jax_fusion_rejection_reasons( ) dependencies_available, dependency_error = ( - tex.grouped_gemm_swiglu_dependencies_available() + tex.grouped_gemm_swiglu_dependencies_available(rubin=compute_capability == 107) ) if not dependencies_available: errors.append( @@ -518,7 +526,7 @@ def _ffn_fwd_per_shard( intermediate_col, intermediate_scale_row, intermediate_scale_col, - ) = tex.grouped_gemm_swiglu( + ) = (tex.grouped_gemm_glu if _is_rubin_device() else tex.grouped_gemm_swiglu)( casted_sorted_x_lhs.data.reshape(sorted_x.shape[0], hidden, 1), casted_wi_rhs.data.reshape( num_local_experts, hidden, wi_for_gemm.shape[-1] From bffeab09e79ff2a4dc64ee70fa92288d2edbbd8b Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Tue, 29 Sep 2026 13:29:51 -0700 Subject: [PATCH 11/38] Add MoE activation checkpoint names Signed-off-by: Jeremy Berchtold --- transformer_engine/jax/moe.py | 52 +++++++++++++++++++++++++++++++++-- 1 file changed, 50 insertions(+), 2 deletions(-) diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index c672db2fec7..67f27222581 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -38,6 +38,7 @@ import flax.struct import jax import jax.numpy as jnp +from jax.ad_checkpoint import checkpoint_name from jax.sharding import NamedSharding, PartitionSpec as P from . import cpp_extensions as tex @@ -349,6 +350,9 @@ def _ffn_fwd_per_shard( num_local_experts: int, activation_type: str, apply_topk_weights_early: bool, + wi_0_checkpoint_name: Optional[str], + wi_1_checkpoint_name: Optional[str], + wo_checkpoint_name: Optional[str], ): """Run the grouped FFN on one shard's EP receive buffer.""" hidden = recv_tokens_local.shape[-1] @@ -380,6 +384,10 @@ def _ffn_fwd_per_shard( bias=wi_combined_bias, ) gate_proj_out, up_proj_out = jnp.split(combined_out, 2, axis=-1) + if wi_0_checkpoint_name is not None: + gate_proj_out = checkpoint_name(gate_proj_out, wi_0_checkpoint_name) + if wi_1_checkpoint_name is not None: + up_proj_out = checkpoint_name(up_proj_out, wi_1_checkpoint_name) # Activation inputs (gate_proj_out, up_proj_out) stay in the wi GEMM # output dtype; the activation output (`intermediate`) stays in the @@ -408,6 +416,8 @@ def _ffn_fwd_per_shard( contracting_dims=((1,), (1,)), bias=wo_bias, ) + if wo_checkpoint_name is not None: + expert_outputs = checkpoint_name(expert_outputs, wo_checkpoint_name) expert_outputs_3d = expert_outputs.reshape(1, expert_outputs.shape[0], expert_outputs.shape[1]) group_sizes_2d = group_sizes.reshape(1, num_local_experts) residuals = ( @@ -568,6 +578,9 @@ def _moe_fwd_rule( dtype, apply_topk_weights_early, recv_capacity_per_rank, + wi_0_checkpoint_name, + wi_1_checkpoint_name, + wo_checkpoint_name, ): """Forward: gate -> topk -> ep_dispatch -> FFN -> ep_combine. @@ -795,6 +808,9 @@ def _ffn_fwd_body(*args): num_local_experts=num_local_experts, activation_type=activation_type, apply_topk_weights_early=apply_topk_weights_early, + wi_0_checkpoint_name=wi_0_checkpoint_name, + wi_1_checkpoint_name=wi_1_checkpoint_name, + wo_checkpoint_name=wo_checkpoint_name, ) expert_outputs, ffn_residuals = shard_map( @@ -893,11 +909,22 @@ def _moe_bwd_rule( dtype, apply_topk_weights_early, recv_capacity_per_rank, + wi_0_checkpoint_name, + wi_1_checkpoint_name, + wo_checkpoint_name, residuals, cotangents, ): """Backward mirror of :func:`_moe_fwd_rule`.""" - del num_groups, group_topk, dtype, recv_capacity_per_rank # captured / unused in bwd + del ( + num_groups, + group_topk, + dtype, + recv_capacity_per_rank, + wi_0_checkpoint_name, + wi_1_checkpoint_name, + wo_checkpoint_name, + ) # captured / unused in bwd from jax.experimental.shard_map import shard_map # total_recv_tokens is a non-differentiable output; its cotangent is unused. @@ -1140,7 +1167,7 @@ def _ffn_bwd_body(*args): # ============================================================================= -@partial(jax.custom_vjp, nondiff_argnums=tuple(range(9, 27))) +@partial(jax.custom_vjp, nondiff_argnums=tuple(range(9, 30))) def _moe( x, gate_kernel, @@ -1169,6 +1196,9 @@ def _moe( dtype, apply_topk_weights_early, recv_capacity_per_rank, + wi_0_checkpoint_name, + wi_1_checkpoint_name, + wo_checkpoint_name, ): primal, _ = _moe_fwd_rule( x, @@ -1198,6 +1228,9 @@ def _moe( dtype, apply_topk_weights_early, recv_capacity_per_rank, + wi_0_checkpoint_name, + wi_1_checkpoint_name, + wo_checkpoint_name, ) return primal @@ -1237,6 +1270,9 @@ def moe( wo_kernel_axes: Tuple[Optional[str], ...] = ("exp", "mlp", "embed"), dtype: jnp.dtype = jnp.float32, recv_capacity_per_rank: Optional[int] = None, + wi_0_checkpoint_name: Optional[str] = None, + wi_1_checkpoint_name: Optional[str] = None, + wo_checkpoint_name: Optional[str] = None, ) -> Tuple[jnp.ndarray, Optional[jnp.ndarray], jnp.ndarray]: """Run a full MoE block under a single fused custom_vjp on the TE EP path. @@ -1271,6 +1307,15 @@ def moe( (default) reserves the dropless aligned worst case. The value must match the capacity used by ``ep_bootstrap``. Overflow is reported through ``total_recv_tokens`` when bootstrap used ``drop_on_overflow=True``. + wi_0_checkpoint_name : Optional[str] + JAX rematerialization checkpoint name for the gate projection output. + ``None`` leaves the value unnamed. + wi_1_checkpoint_name : Optional[str] + JAX rematerialization checkpoint name for the up projection output. + ``None`` leaves the value unnamed. + wo_checkpoint_name : Optional[str] + JAX rematerialization checkpoint name for the per-expert down projection output. + ``None`` leaves the value unnamed. Note that the per-expert dispatch-slot alignment is fixed internally at 128 tokens (``_ALIGN_SIZE``); see that constant's docstring for @@ -1362,6 +1407,9 @@ def moe( dtype, apply_topk_weights_early, recv_capacity_per_rank, + wi_0_checkpoint_name, + wi_1_checkpoint_name, + wo_checkpoint_name, ) if aux_loss_coeff <= 0.0: aux_loss = None From effe0a039b14a3fa6b948285c8f5a1f25302baa5 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Wed, 30 Sep 2026 08:42:14 -0700 Subject: [PATCH 12/38] [JAX] Add quantized MoE weight gather policy Signed-off-by: Jeremy Berchtold --- tests/jax/test_te_ep_moe.py | 48 ++++++- transformer_engine/jax/flax/moe.py | 8 +- transformer_engine/jax/moe.py | 214 ++++++++++++++++++++++++++++- 3 files changed, 262 insertions(+), 8 deletions(-) diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index 015c73343aa..e037e5514de 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -121,6 +121,7 @@ def _read_mp_options(): from transformer_engine.jax.flax import _MoEBlock as MoEBlock from transformer_engine.jax.moe import ( _ALIGN_SIZE, + WeightGather, get_moe_recv_capacity_per_rank, moe, record_ep_bootstrap_signature_for_moe, @@ -129,7 +130,6 @@ def _read_mp_options(): from transformer_engine.common.recipe import MXFP8BlockScaling from transformer_engine.jax.sharding import MeshResource, global_shard_guard - # ----------------------------------------------------------------------------- # Mesh / shape config # ----------------------------------------------------------------------------- @@ -359,6 +359,7 @@ def _make_block( expert_bias_init=None, input_axes=("batch", None, None), quantization_recipe=None, + weight_gather=WeightGather.full_precision(), ): kwargs = dict( num_experts=NUM_EXPERTS, @@ -372,6 +373,7 @@ def _make_block( dtype=DTYPE, input_axes=input_axes, quantization_recipe=quantization_recipe, + weight_gather=weight_gather, ) # Custom expert_bias_init lets tests inject a non-zero expert_bias without # poking variables['params'] post-init. @@ -563,6 +565,50 @@ def _quantization_recipe(quantization): _QUANTIZATION_CASES.append(pytest.param("mxfp8", id="mxfp8")) +def test_weight_gather_policy_axis_resolution(mesh): + """The policy resolves FSDP context only when no axis was supplied.""" + with pytest.raises(ValueError, match="active MeshResource"): + WeightGather.quantized() + with _ctx(mesh): + assert WeightGather.quantized().axis == FSDP_AXIS + with global_shard_guard(MeshResource()): + assert WeightGather.quantized(axis="custom").axis == "custom" + with pytest.raises(ValueError, match="fsdp_resource"): + WeightGather.quantized() + + +def test_quantized_weight_gather_matches_full_precision_gather(mesh): + """The FP8 weight gather retains forward and backward MoE semantics.""" + x = _make_inputs(jax.random.PRNGKey(41)) + baseline = _make_block(quantization_recipe=MXFP8BlockScaling()) + with _ctx(mesh): + gather_policy = WeightGather.quantized() + quantized_ag = _make_block(quantization_recipe=MXFP8BlockScaling(), weight_gather=gather_policy) + variables, baseline_out, _ = _init_apply(baseline, mesh, x, jax.random.PRNGKey(42)) + with _ctx(mesh): + x_sh = _shard_inputs(x, mesh) + quantized_out, _, _ = jax.jit(quantized_ag.apply)(variables, x_sh) + quantized_out.block_until_ready() + np.testing.assert_allclose( + _to_global_numpy(quantized_out, mesh).astype(np.float32), + _to_global_numpy(baseline_out, mesh).astype(np.float32), + **FWD_TOLERANCE["mxfp8"], + ) + baseline_grads, baseline_dx = _grad_step(baseline, variables, mesh, x) + quantized_grads, quantized_dx = _grad_step(quantized_ag, variables, mesh, x) + for name in ("wi", "wo"): + np.testing.assert_allclose( + _to_global_numpy(_unwrap(quantized_grads["params"][name]), mesh).astype(np.float32), + _to_global_numpy(_unwrap(baseline_grads["params"][name]), mesh).astype(np.float32), + **GRAD_FFN_TOLERANCE["mxfp8"], + ) + np.testing.assert_allclose( + _to_global_numpy(quantized_dx, mesh).astype(np.float32), + _to_global_numpy(baseline_dx, mesh).astype(np.float32), + **GRAD_FFN_TOLERANCE["mxfp8"], + ) + + def _reference_kwargs_from_config(config, params_np): """Pick out the reference-relevant pieces of a parametrize config.""" return dict( diff --git a/transformer_engine/jax/flax/moe.py b/transformer_engine/jax/flax/moe.py index 58573966edf..b9e37b14f52 100644 --- a/transformer_engine/jax/flax/moe.py +++ b/transformer_engine/jax/flax/moe.py @@ -34,7 +34,7 @@ from flax import linen as nn from transformer_engine.common.recipe import Recipe -from ..moe import moe +from ..moe import WeightGather, moe from ..quantize import QuantizerSet from ..router import ScoreFunction from ..sharding import _get_mesh, get_active_resource_axis @@ -105,6 +105,10 @@ class _MoEBlock(TransformerEngineBase): recv_capacity_per_rank : Optional[int] Exact aligned receive capacity per EP rank. ``None`` reserves the dropless worst case. + weight_gather : WeightGather + Expert-weight gather policy. Defaults to a full-precision gather. + For MXFP8 gather, use ``WeightGather.quantized()`` to select the + active ``MeshResource.fsdp_resource``, or pass ``axis`` explicitly. The per-expert dispatch-slot alignment is fixed internally at 128 tokens (see ``moe._ALIGN_SIZE``) -- the value required by NCCL EP @@ -151,6 +155,7 @@ class _MoEBlock(TransformerEngineBase): # MoE knobs forwarded to ``moe()`` apply_topk_weights_early: bool = False recv_capacity_per_rank: Optional[int] = None + weight_gather: WeightGather = WeightGather.full_precision() # Dtypes / init / misc dtype: DType = jnp.float32 @@ -304,6 +309,7 @@ def make_grouped_quantizer_set(postfix): apply_topk_weights_early=self.apply_topk_weights_early, quantizer_sets=quantizer_sets, recv_capacity_per_rank=self.recv_capacity_per_rank, + weight_gather=self.weight_gather, ep_axis=ep_axis, data_parallelism_axes=self.data_parallelism_axes, input_axes=self.input_axes, diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index c672db2fec7..5c5e0d925c8 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -32,8 +32,9 @@ import math import warnings +from dataclasses import dataclass from functools import partial -from typing import Any, Optional, Tuple, Union +from typing import Any, Literal, Optional, Tuple, Union import flax.struct import jax @@ -42,17 +43,68 @@ from . import cpp_extensions as tex from .quantize import ( + GroupedScaledTensor1x, GroupedQuantizer, QuantizerSet, + ScaledTensor2x, TensorUsage, noop_quantizer_set, with_sharding_constraint_by_logical_axes, ) +from .quantize.dequantizer import _unswizzle_mxfp8_grouped_scale +from .cpp_extensions.gemm import swizzled_scale from .flax.module import _convert_to_activation_function from .router import ScoreFunction, _validate_score_function -from .sharding import _get_mesh +from .sharding import _get_mesh, global_mesh_resource -__all__ = ["get_moe_recv_capacity_per_rank", "moe"] +__all__ = ["WeightGather", "get_moe_recv_capacity_per_rank", "moe"] + + +@dataclass(frozen=True) +class WeightGather: + """How MoE expert weights are gathered across a sharding mesh axis. + + The quantization recipe determines the wire format for ``quantized``; + this policy only chooses whether quantization precedes the gather. + """ + + mode: Literal["full_precision", "quantized"] = "full_precision" + axis: Optional[str] = None + + def __post_init__(self): + if self.mode not in ("full_precision", "quantized"): + raise ValueError(f"Unsupported weight gather mode: {self.mode!r}") + if self.mode == "quantized" and not self.axis: + raise ValueError("Quantized weight gather requires a mesh axis.") + if self.mode == "full_precision" and self.axis is not None: + raise ValueError("A weight gather axis is only used in quantized mode.") + + @classmethod + def full_precision(cls) -> "WeightGather": + """Gather weights before quantization (the default behavior).""" + return cls() + + @classmethod + def quantized(cls, *, axis: Optional[str] = None) -> "WeightGather": + """Quantize local shards before gathering data and scales. + + With no explicit axis, use the active ``MeshResource.fsdp_resource``. + Passing an axis bypasses the global resource lookup entirely. + """ + if axis is None: + try: + axis = global_mesh_resource().fsdp_resource + except AssertionError as exc: + raise ValueError( + "WeightGather.quantized() requires an active MeshResource " + "with fsdp_resource, or an explicit axis." + ) from exc + if not axis: + raise ValueError( + "WeightGather.quantized() requires MeshResource.fsdp_resource " + "or an explicit axis." + ) + return cls(mode="quantized", axis=axis) # Per-expert dispatch-slot alignment fed to ``tex.ep_prepare`` as @@ -335,6 +387,111 @@ def _validate_moe_quantizer_sets( ) +def _gather_quantized_weight(tensor, fsdp_axis: str, fsdp_size: int, sharded_axis: int): + """Gather an MXFP8 grouped weight without gathering its BF16 source. + + Grouped tensor data and scales are flat, with scales independently padded + and swizzled for each expert. Reassemble each expert in logical order, + then pad and swizzle its gathered scales for grouped GEMM. + """ + if isinstance(tensor, ScaledTensor2x): + return ScaledTensor2x( + _gather_quantized_weight(tensor.rowwise_tensor, fsdp_axis, fsdp_size, sharded_axis), + _gather_quantized_weight(tensor.colwise_tensor, fsdp_axis, fsdp_size, sharded_axis), + ) + if not isinstance(tensor, GroupedScaledTensor1x): + raise TypeError("Quantized FSDP gather requires grouped MXFP8 weight tensors.") + + local_shape = tensor.original_shape + num_experts = local_shape[0] + # The T layout swaps the two matrix dimensions within each expert. + data_axis = 3 - sharded_axis if tensor.data_layout == "T" else sharded_axis + global_shape = list(local_shape) + global_shape[data_axis] *= fsdp_size + global_shape = tuple(global_shape) + data = jax.lax.all_gather( + tensor.data.reshape(local_shape), fsdp_axis, axis=data_axis, tiled=True + ).reshape(-1) + + local_matrix = local_shape[1:] + global_matrix = global_shape[1:] + # Scales use a block-wise 2D view. The local shape can include padding + # that differs from the padding required after the gather. + local_scale_shape = tensor.scaling_mode.get_scale_shape( + local_matrix, + data_layout=tensor.data_layout, + is_colwise=tensor.is_colwise, + is_padded=True, + flatten_axis=tensor.flatten_axis - 1, + ) + local_unpadded_shape = tensor.scaling_mode.get_scale_shape( + local_matrix, + data_layout=tensor.data_layout, + is_colwise=tensor.is_colwise, + is_padded=False, + flatten_axis=tensor.flatten_axis - 1, + ) + global_unpadded_shape = tensor.scaling_mode.get_scale_shape( + global_matrix, + data_layout=tensor.data_layout, + is_colwise=tensor.is_colwise, + is_padded=False, + flatten_axis=tensor.flatten_axis - 1, + ) + global_padded_shape = tensor.scaling_mode.get_scale_shape( + global_matrix, + data_layout=tensor.data_layout, + is_colwise=tensor.is_colwise, + is_padded=True, + flatten_axis=tensor.flatten_axis - 1, + ) + scale_axis = data_axis - 1 + local_scale_size = math.prod(local_scale_shape) + gathered_scales = [] + for expert in range(num_experts): + local_swizzled = jax.lax.dynamic_slice_in_dim( + tensor.scale_inv, expert * local_scale_size, local_scale_size + ) + local_plain = _unswizzle_mxfp8_grouped_scale( + local_swizzled, local_scale_shape, tensor.is_colwise + ) + local_plain = local_plain[: local_unpadded_shape[0], : local_unpadded_shape[1]] + full_plain = jax.lax.all_gather(local_plain, fsdp_axis, axis=scale_axis, tiled=True) + assert full_plain.shape == global_unpadded_shape + full_padded = jnp.pad( + full_plain, + ( + (0, global_padded_shape[0] - global_unpadded_shape[0]), + (0, global_padded_shape[1] - global_unpadded_shape[1]), + ), + ) + gathered_scales.append(swizzled_scale(full_padded, 1, tensor.is_colwise).reshape(-1)) + scale_inv = jnp.concatenate(gathered_scales) + expected_scale_size = tensor.scaling_mode.get_grouped_scale_shape( + global_shape, + num_experts, + tensor.is_colwise, + is_padded=True, + flatten_axis=tensor.flatten_axis, + )[0] + scale_inv = jnp.pad(scale_inv, (0, expected_scale_size - scale_inv.size)) + return GroupedScaledTensor1x( + data=data, + scale_inv=scale_inv, + amax=tensor.amax, + first_dims=None, + last_dims=None, + scaling_mode=tensor.scaling_mode, + dq_dtype=tensor.dq_dtype, + _dq_func=tensor._dq_func, + is_colwise=tensor.is_colwise, + data_layout=tensor.data_layout, + flatten_axis=tensor.flatten_axis, + original_shape=global_shape, + pre_swizzled=True, + ) + + def _ffn_fwd_per_shard( recv_tokens_local: jnp.ndarray, recv_topk_weights_local: jnp.ndarray, @@ -349,6 +506,8 @@ def _ffn_fwd_per_shard( num_local_experts: int, activation_type: str, apply_topk_weights_early: bool, + weight_gather: WeightGather, + fsdp_size: int, ): """Run the grouped FFN on one shard's EP receive buffer.""" hidden = recv_tokens_local.shape[-1] @@ -373,6 +532,8 @@ def _ffn_fwd_per_shard( flatten_axis=-1, ) casted_wi = tex.grouped_quantize(wi, fc1_quantizer_set.kernel, flatten_axis=-1) + if weight_gather.mode == "quantized": + casted_wi = _gather_quantized_weight(casted_wi, weight_gather.axis, fsdp_size, 1) combined_out = tex.grouped_gemm( casted_sorted_x.get_tensor(usage=TensorUsage.LHS), casted_wi.get_tensor(usage=TensorUsage.RHS), @@ -402,6 +563,8 @@ def _ffn_fwd_per_shard( flatten_axis=-1, ) casted_wo = tex.grouped_quantize(wo, fc2_quantizer_set.kernel, flatten_axis=-1) + if weight_gather.mode == "quantized": + casted_wo = _gather_quantized_weight(casted_wo, weight_gather.axis, fsdp_size, 2) expert_outputs = tex.grouped_gemm( casted_intermediate.get_tensor(usage=TensorUsage.LHS), casted_wo.get_tensor(usage=TensorUsage.RHS), @@ -568,6 +731,7 @@ def _moe_fwd_rule( dtype, apply_topk_weights_early, recv_capacity_per_rank, + weight_gather, ): """Forward: gate -> topk -> ep_dispatch -> FFN -> ep_combine. @@ -598,6 +762,17 @@ def _moe_fwd_rule( num_token_groups=dp_size * num_experts, num_expert_groups=num_experts, ) + if weight_gather.mode == "quantized": + if weight_gather.axis not in data_parallelism_axes: + raise ValueError( + "Quantized weight all-gather requires its axis in data_parallelism_axes." + ) + if any(quantizer_set.kernel is None for quantizer_set in quantizer_sets): + raise ValueError("Quantized weight all-gather requires MXFP8 kernel quantizers.") + if wi.shape[1] % (mesh.shape[weight_gather.axis] * 32) or wo.shape[2] % ( + mesh.shape[weight_gather.axis] * 32 + ): + raise ValueError("FSDP weight shards must be divisible by the MXFP8 block size 32.") B, S, H = x.shape K = num_experts_per_tok @@ -751,8 +926,14 @@ def _moe_fwd_rule( # ---------------- FFN (per-shard via shard_map) ---------------- has_bias = wi_0_bias is not None kernel_spec = P(ep_axis, None, None) + wi_input_spec = ( + P(ep_axis, weight_gather.axis, None) if weight_gather.mode == "quantized" else kernel_spec + ) + wo_input_spec = ( + P(ep_axis, None, weight_gather.axis) if weight_gather.mode == "quantized" else kernel_spec + ) bias_spec = P(ep_axis, None) - ffn_in_specs = (ep3_spec, ep2_spec, ep2_spec, kernel_spec, kernel_spec) + ffn_in_specs = (ep3_spec, ep2_spec, ep2_spec, wi_input_spec, wo_input_spec) ffn_in_args = [recv_tokens, recv_topk_weights, token_counts, wi, wo] if has_bias: ffn_in_specs += (bias_spec, bias_spec, bias_spec) @@ -795,6 +976,8 @@ def _ffn_fwd_body(*args): num_local_experts=num_local_experts, activation_type=activation_type, apply_topk_weights_early=apply_topk_weights_early, + weight_gather=weight_gather, + fsdp_size=mesh.shape[weight_gather.axis] if weight_gather.axis is not None else 1, ) expert_outputs, ffn_residuals = shard_map( @@ -893,11 +1076,18 @@ def _moe_bwd_rule( dtype, apply_topk_weights_early, recv_capacity_per_rank, + weight_gather, residuals, cotangents, ): """Backward mirror of :func:`_moe_fwd_rule`.""" - del num_groups, group_topk, dtype, recv_capacity_per_rank # captured / unused in bwd + del ( + num_groups, + group_topk, + dtype, + recv_capacity_per_rank, + weight_gather, + ) from jax.experimental.shard_map import shard_map # total_recv_tokens is a non-differentiable output; its cotangent is unused. @@ -1140,7 +1330,7 @@ def _ffn_bwd_body(*args): # ============================================================================= -@partial(jax.custom_vjp, nondiff_argnums=tuple(range(9, 27))) +@partial(jax.custom_vjp, nondiff_argnums=tuple(range(9, 28))) def _moe( x, gate_kernel, @@ -1169,6 +1359,7 @@ def _moe( dtype, apply_topk_weights_early, recv_capacity_per_rank, + weight_gather, ): primal, _ = _moe_fwd_rule( x, @@ -1198,6 +1389,7 @@ def _moe( dtype, apply_topk_weights_early, recv_capacity_per_rank, + weight_gather, ) return primal @@ -1237,6 +1429,7 @@ def moe( wo_kernel_axes: Tuple[Optional[str], ...] = ("exp", "mlp", "embed"), dtype: jnp.dtype = jnp.float32, recv_capacity_per_rank: Optional[int] = None, + weight_gather: WeightGather = WeightGather.full_precision(), ) -> Tuple[jnp.ndarray, Optional[jnp.ndarray], jnp.ndarray]: """Run a full MoE block under a single fused custom_vjp on the TE EP path. @@ -1271,6 +1464,12 @@ def moe( (default) reserves the dropless aligned worst case. The value must match the capacity used by ``ep_bootstrap``. Overflow is reported through ``total_recv_tokens`` when bootstrap used ``drop_on_overflow=True``. + weight_gather : WeightGather + Expert-weight gather policy. ``WeightGather.full_precision()`` gathers + weights before quantization (the default). Use + ``WeightGather.quantized()`` to quantize local shards and gather their + MXFP8 data and scales along the active ``MeshResource.fsdp_resource``; + pass ``axis`` to select a physical mesh axis without consulting it. Note that the per-expert dispatch-slot alignment is fixed internally at 128 tokens (``_ALIGN_SIZE``); see that constant's docstring for @@ -1305,6 +1504,8 @@ def moe( surrounding design rationale. """ score_function = _validate_score_function(score_function) + if not isinstance(weight_gather, WeightGather): + raise TypeError("weight_gather must be a WeightGather policy.") # Enforce ((outer_dp..., ep), None, None) on inbound activations. The # EP comm groups consecutive global ranks (dp_color = rank // ep_size), @@ -1362,6 +1563,7 @@ def moe( dtype, apply_topk_weights_early, recv_capacity_per_rank, + weight_gather, ) if aux_loss_coeff <= 0.0: aux_loss = None From 358908e3bfc56d27009aa621920c553b34f8318b Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Thu, 1 Oct 2026 08:06:35 -0700 Subject: [PATCH 13/38] Checkpoint fused JAX MoE projections Signed-off-by: Jeremy Berchtold --- tests/jax/test_te_ep_moe.py | 45 +++++++++++++++++++++++++++++++++++ transformer_engine/jax/moe.py | 22 ++++++++++------- 2 files changed, 58 insertions(+), 9 deletions(-) diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index 8e308cde5cd..0b0955a4402 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -42,6 +42,7 @@ main+aux grads stay finite) in two consolidated tests. """ +import importlib import os os.environ.setdefault("XLA_PYTHON_CLIENT_PREALLOCATE", "false") @@ -818,6 +819,50 @@ def loss_fn(params, inputs): err_msg="d_x fused MXFP8 gradient parity breach", ) + def test_regular_swiglu_with_checkpoint_names(self, mesh, monkeypatch): + if not _use_cudnn_cutedsl_fusion_from_env(): + pytest.skip("cuDNN grouped GEMM fusion is disabled") + if get_device_compute_capability(0) != 107: + pytest.skip("Rubin test of the regular cuDNN grouped SwiGLU path") + + from transformer_engine.jax import cpp_extensions as tex + + moe_module = importlib.import_module("transformer_engine.jax.moe") + flax_moe_module = importlib.import_module("transformer_engine.jax.flax.moe") + monkeypatch.setattr(moe_module, "_is_rubin_device", lambda: False) + + regular_calls = [] + original_swiglu = tex.grouped_gemm_swiglu + + def checked_swiglu(*args, **kwargs): + regular_calls.append(True) + return original_swiglu(*args, **kwargs) + + monkeypatch.setattr(tex, "grouped_gemm_swiglu", checked_swiglu) + original_moe = flax_moe_module.moe + + def checkpointed_moe(*args, **kwargs): + kwargs.update( + wi_0_checkpoint_name="moe_mlpwi_0", + wi_1_checkpoint_name="moe_mlpwi_1", + wo_checkpoint_name="moe_mlpwo", + ) + return original_moe(*args, **kwargs) + + monkeypatch.setattr(flax_moe_module, "moe", checkpointed_moe) + block = _make_block(quantization_recipe=MXFP8BlockScaling()) + inputs = _make_inputs(jax.random.PRNGKey(32)) + variables, output, _ = _init_apply(block, mesh, inputs, jax.random.PRNGKey(33)) + grads, grad_inputs = _grad_step(block, variables, mesh, inputs) + + assert regular_calls, "regular cuDNN grouped SwiGLU was not selected" + assert np.all(np.isfinite(_to_global_numpy(output, mesh))) + assert np.all(np.isfinite(_to_global_numpy(grad_inputs, mesh))) + for name in ("gate_kernel", "wi", "wo"): + assert np.all( + np.isfinite(_to_global_numpy(_unwrap(grads["params"][name]), mesh)) + ) + class TestTeEpMoeAuxLoss: """Aux-loss path. Consolidated into: diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index d4c1bfee00e..a145afb94bd 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -543,10 +543,6 @@ def _ffn_fwd_per_shard( output_dtype=fc2_quantizer_set.x.q_dtype, ) combined_out = combined_out_3d.reshape(sorted_x.shape[0], wi_for_gemm.shape[-1]) - if wi_0_checkpoint_name is not None: - combined_out = checkpoint_name(combined_out, wi_0_checkpoint_name) - if wi_1_checkpoint_name is not None: - combined_out = checkpoint_name(combined_out, wi_1_checkpoint_name) gate_proj_out, up_proj_out = tex.unpack_swiglu_pair(combined_out) intermediate_shape = (sorted_x.shape[0], gate_proj_out.shape[-1]) @@ -595,10 +591,10 @@ def _ffn_fwd_per_shard( bias=wi_combined_bias, ) gate_proj_out, up_proj_out = jnp.split(combined_out, 2, axis=-1) - if wi_0_checkpoint_name is not None: - gate_proj_out = checkpoint_name(gate_proj_out, wi_0_checkpoint_name) - if wi_1_checkpoint_name is not None: - up_proj_out = checkpoint_name(up_proj_out, wi_1_checkpoint_name) + if wi_0_checkpoint_name is not None: + gate_proj_out = checkpoint_name(gate_proj_out, wi_0_checkpoint_name) + if wi_1_checkpoint_name is not None: + up_proj_out = checkpoint_name(up_proj_out, wi_1_checkpoint_name) # Activation inputs (gate_proj_out, up_proj_out) stay in the wi GEMM # output dtype; the activation output (`intermediate`) stays in the @@ -635,6 +631,14 @@ def _ffn_fwd_per_shard( 1, expert_outputs.shape[0], expert_outputs.shape[1] ) group_sizes_2d = group_sizes.reshape(1, num_local_experts) + if use_cudnn_jax_fusion: + ffn_activation_residual = ( + tex.pack_swiglu_pair(gate_proj_out, up_proj_out) + if wi_0_checkpoint_name is not None or wi_1_checkpoint_name is not None + else combined_out + ) + else: + ffn_activation_residual = gate_proj_out residuals = ( casted_sorted_x.get_tensor(usage=TensorUsage.LHS_TRANS).checkpoint( fc1_quantizer_set.x @@ -642,7 +646,7 @@ def _ffn_fwd_per_shard( casted_wi.get_tensor(usage=TensorUsage.RHS_TRANS).checkpoint( fc1_quantizer_set.kernel ), - combined_out if use_cudnn_jax_fusion else gate_proj_out, + ffn_activation_residual, up_proj_out, casted_intermediate.get_tensor(usage=TensorUsage.LHS_TRANS).checkpoint( fc2_quantizer_set.x From 00fe8b7eb9160305c6bd3d0a7cf0d742fa2ad2ea Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Thu, 1 Oct 2026 08:32:07 -0700 Subject: [PATCH 14/38] Checkpoint fused MoE combined projection without repacking --- transformer_engine/jax/moe.py | 22 ++++++++++------------ 1 file changed, 10 insertions(+), 12 deletions(-) diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index a145afb94bd..17f5a79f8ac 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -543,6 +543,10 @@ def _ffn_fwd_per_shard( output_dtype=fc2_quantizer_set.x.q_dtype, ) combined_out = combined_out_3d.reshape(sorted_x.shape[0], wi_for_gemm.shape[-1]) + if wi_0_checkpoint_name is not None: + combined_out = checkpoint_name(combined_out, wi_0_checkpoint_name) + if wi_1_checkpoint_name is not None: + combined_out = checkpoint_name(combined_out, wi_1_checkpoint_name) gate_proj_out, up_proj_out = tex.unpack_swiglu_pair(combined_out) intermediate_shape = (sorted_x.shape[0], gate_proj_out.shape[-1]) @@ -591,10 +595,11 @@ def _ffn_fwd_per_shard( bias=wi_combined_bias, ) gate_proj_out, up_proj_out = jnp.split(combined_out, 2, axis=-1) - if wi_0_checkpoint_name is not None: - gate_proj_out = checkpoint_name(gate_proj_out, wi_0_checkpoint_name) - if wi_1_checkpoint_name is not None: - up_proj_out = checkpoint_name(up_proj_out, wi_1_checkpoint_name) + if not use_cudnn_jax_fusion: + if wi_0_checkpoint_name is not None: + gate_proj_out = checkpoint_name(gate_proj_out, wi_0_checkpoint_name) + if wi_1_checkpoint_name is not None: + up_proj_out = checkpoint_name(up_proj_out, wi_1_checkpoint_name) # Activation inputs (gate_proj_out, up_proj_out) stay in the wi GEMM # output dtype; the activation output (`intermediate`) stays in the @@ -631,14 +636,7 @@ def _ffn_fwd_per_shard( 1, expert_outputs.shape[0], expert_outputs.shape[1] ) group_sizes_2d = group_sizes.reshape(1, num_local_experts) - if use_cudnn_jax_fusion: - ffn_activation_residual = ( - tex.pack_swiglu_pair(gate_proj_out, up_proj_out) - if wi_0_checkpoint_name is not None or wi_1_checkpoint_name is not None - else combined_out - ) - else: - ffn_activation_residual = gate_proj_out + ffn_activation_residual = combined_out if use_cudnn_jax_fusion else gate_proj_out residuals = ( casted_sorted_x.get_tensor(usage=TensorUsage.LHS_TRANS).checkpoint( fc1_quantizer_set.x From e254e043547da95e0930a4ae5ec946bbfaf78957 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Thu, 1 Oct 2026 09:39:19 -0700 Subject: [PATCH 15/38] Use expert-packed scales for fused JAX MoE GEMMs --- transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py b/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py index 15559420bed..95516e5df3d 100644 --- a/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py +++ b/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py @@ -182,6 +182,7 @@ def _grouped_gemm_forward( norm_const_tensor=norm_const, c_dtype=jnp.dtype(compute_dtype), d_dtype=jnp.dtype(output_dtype), + discrete_col_sfd=True, ) return ( result["c_tensor"].reshape(rows, combined), @@ -237,6 +238,7 @@ def grouped_gemm_dswiglu( prob_tensor=prob.reshape(rows).astype(jnp.float32), norm_const_tensor=norm_const, d_dtype=jnp.dtype(output_dtype), + discrete_col_sfd=True, ) return ( result["d_row_tensor"], From 89b19c4f56bf9bb6c535fa92add4d209478b2aaa Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Thu, 1 Oct 2026 10:41:26 -0700 Subject: [PATCH 16/38] Checkpoint fused cuDNN MoE quantized outputs Signed-off-by: Jeremy Berchtold --- tests/jax/test_te_ep_moe.py | 38 +++++++++++++++++++++++++---------- transformer_engine/jax/moe.py | 12 +++++++++++ 2 files changed, 39 insertions(+), 11 deletions(-) diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index 0b0955a4402..372bad2c6c1 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -819,26 +819,37 @@ def loss_fn(params, inputs): err_msg="d_x fused MXFP8 gradient parity breach", ) - def test_regular_swiglu_with_checkpoint_names(self, mesh, monkeypatch): + @pytest.mark.parametrize("use_regular_swiglu", [False, True]) + def test_cudnn_fused_with_checkpoint_names(self, mesh, monkeypatch, use_regular_swiglu): if not _use_cudnn_cutedsl_fusion_from_env(): pytest.skip("cuDNN grouped GEMM fusion is disabled") - if get_device_compute_capability(0) != 107: - pytest.skip("Rubin test of the regular cuDNN grouped SwiGLU path") + if not use_regular_swiglu and get_device_compute_capability(0) != 107: + pytest.skip("Rubin grouped GLU requires SM107") from transformer_engine.jax import cpp_extensions as tex moe_module = importlib.import_module("transformer_engine.jax.moe") flax_moe_module = importlib.import_module("transformer_engine.jax.flax.moe") - monkeypatch.setattr(moe_module, "_is_rubin_device", lambda: False) + if use_regular_swiglu: + monkeypatch.setattr(moe_module, "_is_rubin_device", lambda: False) - regular_calls = [] - original_swiglu = tex.grouped_gemm_swiglu + selected_calls = [] + fused_op_name = "grouped_gemm_swiglu" if use_regular_swiglu else "grouped_gemm_glu" + original_fused_op = getattr(tex, fused_op_name) - def checked_swiglu(*args, **kwargs): - regular_calls.append(True) - return original_swiglu(*args, **kwargs) + def checked_fused_op(*args, **kwargs): + selected_calls.append(True) + return original_fused_op(*args, **kwargs) - monkeypatch.setattr(tex, "grouped_gemm_swiglu", checked_swiglu) + monkeypatch.setattr(tex, fused_op_name, checked_fused_op) + named_values = {} + original_checkpoint_name = moe_module.checkpoint_name + + def recorded_checkpoint_name(value, name): + named_values.setdefault(name, []).append(value) + return original_checkpoint_name(value, name) + + monkeypatch.setattr(moe_module, "checkpoint_name", recorded_checkpoint_name) original_moe = flax_moe_module.moe def checkpointed_moe(*args, **kwargs): @@ -855,7 +866,12 @@ def checkpointed_moe(*args, **kwargs): variables, output, _ = _init_apply(block, mesh, inputs, jax.random.PRNGKey(33)) grads, grad_inputs = _grad_step(block, variables, mesh, inputs) - assert regular_calls, "regular cuDNN grouped SwiGLU was not selected" + assert selected_calls, f"{fused_op_name} was not selected" + for label in ("moe_mlpwi_0", "moe_mlpwi_1"): + assert len(named_values[label]) == 5 * len(selected_calls) + assert all(hasattr(value, "shape") for value in named_values[label]) + assert len(named_values["moe_mlpwo"]) == len(selected_calls) + assert all(hasattr(value, "shape") for value in named_values["moe_mlpwo"]) assert np.all(np.isfinite(_to_global_numpy(output, mesh))) assert np.all(np.isfinite(_to_global_numpy(grad_inputs, mesh))) for name in ("gate_kernel", "wi", "wo"): diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index 17f5a79f8ac..6c0884988dd 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -573,6 +573,16 @@ def _ffn_fwd_per_shard( intermediate_scale_col, (0, col_scale_size - intermediate_scale_col.size), ) + for checkpoint_label in (wi_0_checkpoint_name, wi_1_checkpoint_name): + if checkpoint_label is not None: + intermediate_row = checkpoint_name(intermediate_row, checkpoint_label) + intermediate_col = checkpoint_name(intermediate_col, checkpoint_label) + intermediate_scale_row = checkpoint_name( + intermediate_scale_row, checkpoint_label + ) + intermediate_scale_col = checkpoint_name( + intermediate_scale_col, checkpoint_label + ) casted_intermediate = ScaledTensorFactory.create( data=intermediate_row.reshape(-1), scale_inv=intermediate_scale_row, @@ -1639,9 +1649,11 @@ def moe( ``total_recv_tokens`` when bootstrap used ``drop_on_overflow=True``. wi_0_checkpoint_name : Optional[str] JAX rematerialization checkpoint name for the gate projection output. + With cuDNN fusion, also names the shared fused outputs, including its quantized output. ``None`` leaves the value unnamed. wi_1_checkpoint_name : Optional[str] JAX rematerialization checkpoint name for the up projection output. + With cuDNN fusion, also names the same fused outputs, including its quantized output. ``None`` leaves the value unnamed. wo_checkpoint_name : Optional[str] JAX rematerialization checkpoint name for the per-expert down projection output. From dfde5f32996688a61f388a7a8b7b8cce08d7f336 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Thu, 1 Oct 2026 13:38:14 -0700 Subject: [PATCH 17/38] Add EP dispatch and combine checkpoint names to JAX MoE Signed-off-by: Jeremy Berchtold --- tests/jax/test_te_ep_moe.py | 47 +++++++++++++++++++++++++++++- transformer_engine/jax/flax/moe.py | 7 +++++ transformer_engine/jax/moe.py | 27 ++++++++++++++++- 3 files changed, 79 insertions(+), 2 deletions(-) diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index 372bad2c6c1..77b5990ab32 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -366,6 +366,8 @@ def _make_block( expert_bias_init=None, input_axes=("batch", None, None), quantization_recipe=None, + dispatch_checkpoint_name=None, + combine_checkpoint_name=None, ): kwargs = dict( num_experts=NUM_EXPERTS, @@ -379,6 +381,8 @@ def _make_block( dtype=DTYPE, input_axes=input_axes, quantization_recipe=quantization_recipe, + dispatch_checkpoint_name=dispatch_checkpoint_name, + combine_checkpoint_name=combine_checkpoint_name, ) # Custom expert_bias_init lets tests inject a non-zero expert_bias without # poking variables['params'] post-init. @@ -699,6 +703,41 @@ def loss_fn(params, x): ) +def test_ep_checkpoint_names(mesh, monkeypatch): + if _use_cudnn_cutedsl_fusion_from_env(): + pytest.skip( + "BF16 fallback uses a different EP alignment than the cuDNN bootstrap" + ) + moe_module = importlib.import_module("transformer_engine.jax.moe") + named_values = {} + original_checkpoint_name = moe_module.checkpoint_name + + def recorded_checkpoint_name(value, name): + named_values.setdefault(name, []).append(value) + return original_checkpoint_name(value, name) + + monkeypatch.setattr(moe_module, "checkpoint_name", recorded_checkpoint_name) + block = _make_block( + dispatch_checkpoint_name="saved_dispatch", + combine_checkpoint_name="saved_combine", + ) + inputs = _make_inputs(jax.random.PRNGKey(34)) + variables, output, _ = _init_apply(block, mesh, inputs, jax.random.PRNGKey(35)) + grads, grad_inputs = _grad_step(block, variables, mesh, inputs) + + assert len(named_values["saved_dispatch"]) >= 2 + assert len(named_values["saved_combine"]) >= 1 + assert all( + hasattr(value, "shape") for values in named_values.values() for value in values + ) + assert np.all(np.isfinite(_to_global_numpy(output, mesh))) + assert np.all(np.isfinite(_to_global_numpy(grad_inputs, mesh))) + for name in ("gate_kernel", "wi", "wo"): + assert np.all( + np.isfinite(_to_global_numpy(_unwrap(grads["params"][name]), mesh)) + ) + + class TestTeEpMoeCudnnCutedslFusion: """End-to-end MXFP8 coverage for cuDNN's grouped GLU JAX APIs.""" @@ -861,7 +900,11 @@ def checkpointed_moe(*args, **kwargs): return original_moe(*args, **kwargs) monkeypatch.setattr(flax_moe_module, "moe", checkpointed_moe) - block = _make_block(quantization_recipe=MXFP8BlockScaling()) + block = _make_block( + quantization_recipe=MXFP8BlockScaling(), + dispatch_checkpoint_name="saved_dispatch", + combine_checkpoint_name="saved_combine", + ) inputs = _make_inputs(jax.random.PRNGKey(32)) variables, output, _ = _init_apply(block, mesh, inputs, jax.random.PRNGKey(33)) grads, grad_inputs = _grad_step(block, variables, mesh, inputs) @@ -872,6 +915,8 @@ def checkpointed_moe(*args, **kwargs): assert all(hasattr(value, "shape") for value in named_values[label]) assert len(named_values["moe_mlpwo"]) == len(selected_calls) assert all(hasattr(value, "shape") for value in named_values["moe_mlpwo"]) + assert len(named_values["saved_dispatch"]) >= 2 * len(selected_calls) + assert len(named_values["saved_combine"]) >= len(selected_calls) assert np.all(np.isfinite(_to_global_numpy(output, mesh))) assert np.all(np.isfinite(_to_global_numpy(grad_inputs, mesh))) for name in ("gate_kernel", "wi", "wo"): diff --git a/transformer_engine/jax/flax/moe.py b/transformer_engine/jax/flax/moe.py index 58573966edf..27a66aff87c 100644 --- a/transformer_engine/jax/flax/moe.py +++ b/transformer_engine/jax/flax/moe.py @@ -105,6 +105,9 @@ class _MoEBlock(TransformerEngineBase): recv_capacity_per_rank : Optional[int] Exact aligned receive capacity per EP rank. ``None`` reserves the dropless worst case. + dispatch_checkpoint_name, combine_checkpoint_name : Optional[str] + JAX rematerialization checkpoint names for the EP dispatch outputs + and combine output, respectively. ``None`` leaves them unnamed. The per-expert dispatch-slot alignment is fixed internally at 128 tokens (see ``moe._ALIGN_SIZE``) -- the value required by NCCL EP @@ -151,6 +154,8 @@ class _MoEBlock(TransformerEngineBase): # MoE knobs forwarded to ``moe()`` apply_topk_weights_early: bool = False recv_capacity_per_rank: Optional[int] = None + dispatch_checkpoint_name: Optional[str] = None + combine_checkpoint_name: Optional[str] = None # Dtypes / init / misc dtype: DType = jnp.float32 @@ -311,4 +316,6 @@ def make_grouped_quantizer_set(postfix): wi_kernel_axes=self.wi_kernel_axes, wo_kernel_axes=self.wo_kernel_axes, dtype=self.dtype, + dispatch_checkpoint_name=self.dispatch_checkpoint_name, + combine_checkpoint_name=self.combine_checkpoint_name, ) diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index 6c0884988dd..31730b7e2f8 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -881,6 +881,8 @@ def _moe_fwd_rule( wi_0_checkpoint_name, wi_1_checkpoint_name, wo_checkpoint_name, + dispatch_checkpoint_name, + combine_checkpoint_name, ): """Forward: gate -> topk -> ep_dispatch -> FFN -> ep_combine. @@ -1071,6 +1073,9 @@ def _moe_fwd_rule( recv_tokens, recv_topk_weights = tex.ep_dispatch_fwd( cfg, handle_mem, topk_idx_3d, x, topk_w_3d, recv_pr ) + if dispatch_checkpoint_name is not None: + recv_tokens = checkpoint_name(recv_tokens, dispatch_checkpoint_name) + recv_topk_weights = checkpoint_name(recv_topk_weights, dispatch_checkpoint_name) recv_tokens = jax.lax.with_sharding_constraint( recv_tokens, NamedSharding(mesh, ep3_spec) ) @@ -1165,6 +1170,8 @@ def _ffn_fwd_body(*args): num_local_tokens=(B, S), out_partition_spec=out_partition_spec, ) + if combine_checkpoint_name is not None: + output = checkpoint_name(output, combine_checkpoint_name) # output of MLP should be sharded the same way as the activation input output = with_sharding_constraint_by_logical_axes(output, input_axes) @@ -1233,6 +1240,8 @@ def _moe_bwd_rule( wi_0_checkpoint_name, wi_1_checkpoint_name, wo_checkpoint_name, + dispatch_checkpoint_name, + combine_checkpoint_name, residuals, cotangents, ): @@ -1245,6 +1254,8 @@ def _moe_bwd_rule( wi_0_checkpoint_name, wi_1_checkpoint_name, wo_checkpoint_name, + dispatch_checkpoint_name, + combine_checkpoint_name, ) # captured / unused in bwd from jax.experimental.shard_map import shard_map @@ -1505,7 +1516,7 @@ def _ffn_bwd_body(*args): # ============================================================================= -@partial(jax.custom_vjp, nondiff_argnums=tuple(range(9, 31))) +@partial(jax.custom_vjp, nondiff_argnums=tuple(range(9, 33))) def _moe( x, gate_kernel, @@ -1538,6 +1549,8 @@ def _moe( wi_0_checkpoint_name, wi_1_checkpoint_name, wo_checkpoint_name, + dispatch_checkpoint_name, + combine_checkpoint_name, ): primal, _ = _moe_fwd_rule( x, @@ -1571,6 +1584,8 @@ def _moe( wi_0_checkpoint_name, wi_1_checkpoint_name, wo_checkpoint_name, + dispatch_checkpoint_name, + combine_checkpoint_name, ) return primal @@ -1613,6 +1628,8 @@ def moe( wi_0_checkpoint_name: Optional[str] = None, wi_1_checkpoint_name: Optional[str] = None, wo_checkpoint_name: Optional[str] = None, + dispatch_checkpoint_name: Optional[str] = None, + combine_checkpoint_name: Optional[str] = None, ) -> Tuple[jnp.ndarray, Optional[jnp.ndarray], jnp.ndarray]: """Run a full MoE block under a single fused custom_vjp on the TE EP path. @@ -1658,6 +1675,12 @@ def moe( wo_checkpoint_name : Optional[str] JAX rematerialization checkpoint name for the per-expert down projection output. ``None`` leaves the value unnamed. + dispatch_checkpoint_name : Optional[str] + JAX rematerialization checkpoint name for the EP dispatch outputs. + ``None`` leaves these values unnamed; the opaque prepare handle is never named. + combine_checkpoint_name : Optional[str] + JAX rematerialization checkpoint name for the EP combine output. + ``None`` leaves the value unnamed. Note that the per-expert dispatch-slot alignment is fixed internally at 128 tokens (``_ALIGN_SIZE``); see that constant's docstring for @@ -1782,6 +1805,8 @@ def moe( wi_0_checkpoint_name, wi_1_checkpoint_name, wo_checkpoint_name, + dispatch_checkpoint_name, + combine_checkpoint_name, ) if aux_loss_coeff <= 0.0: aux_loss = None From 1410d5229709da3830eee552bb28575b6b4e5723 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Thu, 1 Oct 2026 14:04:01 -0700 Subject: [PATCH 18/38] Add native cuDNN grouped GEMM weight layout Signed-off-by: Jeremy Berchtold --- transformer_engine/jax/moe.py | 98 +++++++++++++++++++++++++++-------- 1 file changed, 76 insertions(+), 22 deletions(-) diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index 31730b7e2f8..89be852785d 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -110,9 +110,17 @@ def _cudnn_jax_fusion_rejection_reasons( errors.append("requires activation_type='silu'") if wi_0_bias is not None or wi_1_bias is not None: errors.append("does not support FC1 gate/up bias") - if wi.ndim != 3 or wi.shape[-1] % 64: + hidden = x.shape[-1] + standard_layout = ( + wi.ndim == 3 and wi.shape[-2] == hidden and wi.shape[-1] % 64 == 0 + ) + native_layout = ( + wi.ndim == 3 and wi.shape[-1] == hidden and wi.shape[-2] % 64 == 0 + ) + if not standard_layout and not native_layout: errors.append( - f"requires rank-3 wi with a 64-aligned gated dimension, got {wi.shape}" + "requires rank-3 wi in standard [E,K,2N] or cuDNN-native [E,2N,K] " + f"layout with a 64-aligned gated dimension and K={hidden}, got {wi.shape}" ) if x.dtype not in (jnp.bfloat16, jnp.float16): errors.append(f"requires BF16 or FP16 activations, got {x.dtype}") @@ -481,6 +489,7 @@ def _ffn_fwd_per_shard( wi_0_checkpoint_name: Optional[str], wi_1_checkpoint_name: Optional[str], wo_checkpoint_name: Optional[str], + cudnn_native_weight_layout: bool, ): """Run the grouped FFN on one shard's EP receive buffer.""" hidden = recv_tokens_local.shape[-1] @@ -491,13 +500,22 @@ def _ffn_fwd_per_shard( wi = wi.astype(sorted_x.dtype) wo = wo.astype(sorted_x.dtype) - # TE stores wi as [gate, up]. cuDNN consumes alternating 32-column - # gate/up blocks, while the regular grouped GEMM consumes TE's layout. + # The cuDNN-native parameter is persistent [E,2N,K] storage with alternating + # 32-column gate/up blocks. Standard TE storage is [E,K,2N] with contiguous + # gate/up halves. Keep conversion only as a compatibility fallback. if use_cudnn_jax_fusion: - wi_gate, wi_up = jnp.split(wi, 2, axis=-1) - wi_for_gemm = tex.pack_swiglu_pair(wi_gate, wi_up) + if cudnn_native_weight_layout: + wi_for_gemm = wi + else: + wi_gate, wi_up = jnp.split(wi, 2, axis=-1) + wi_for_gemm = tex.pack_swiglu_pair(wi_gate, wi_up) else: - wi_for_gemm = wi + if cudnn_native_weight_layout: + wi_interleaved = wi.transpose(0, 2, 1) + wi_gate, wi_up = tex.unpack_swiglu_pair(wi_interleaved) + wi_for_gemm = jnp.concatenate((wi_gate, wi_up), axis=-1) + else: + wi_for_gemm = wi wi_combined_bias = ( jnp.concatenate([wi_0_bias, wi_1_bias], axis=-1) if wi_0_bias is not None @@ -517,7 +535,14 @@ def _ffn_fwd_per_shard( casted_intermediate = None if use_cudnn_jax_fusion: casted_sorted_x_lhs = casted_sorted_x.get_tensor(usage=TensorUsage.LHS) - casted_wi_rhs = casted_wi.get_tensor(usage=TensorUsage.RHS) + casted_wi_rhs = casted_wi.get_tensor( + usage=TensorUsage.LHS if cudnn_native_weight_layout else TensorUsage.RHS + ) + combined = ( + wi_for_gemm.shape[-2] + if cudnn_native_weight_layout + else wi_for_gemm.shape[-1] + ) padded_offsets = jnp.cumsum(group_sizes, dtype=jnp.int32) prob = ( recv_w_flat[:, None, None] @@ -532,9 +557,13 @@ def _ffn_fwd_per_shard( intermediate_scale_col, ) = (tex.grouped_gemm_glu if _is_rubin_device() else tex.grouped_gemm_swiglu)( casted_sorted_x_lhs.data.reshape(sorted_x.shape[0], hidden, 1), - casted_wi_rhs.data.reshape( - num_local_experts, hidden, wi_for_gemm.shape[-1] - ).transpose(0, 2, 1), + ( + casted_wi_rhs.data.reshape(num_local_experts, combined, hidden) + if cudnn_native_weight_layout + else casted_wi_rhs.data.reshape( + num_local_experts, hidden, combined + ).transpose(0, 2, 1) + ), casted_sorted_x_lhs.scale_inv, casted_wi_rhs.scale_inv, padded_offsets, @@ -542,14 +571,16 @@ def _ffn_fwd_per_shard( compute_dtype=sorted_x.dtype, output_dtype=fc2_quantizer_set.x.q_dtype, ) - combined_out = combined_out_3d.reshape(sorted_x.shape[0], wi_for_gemm.shape[-1]) + combined_out = combined_out_3d.reshape(sorted_x.shape[0], combined) if wi_0_checkpoint_name is not None: combined_out = checkpoint_name(combined_out, wi_0_checkpoint_name) if wi_1_checkpoint_name is not None: combined_out = checkpoint_name(combined_out, wi_1_checkpoint_name) - gate_proj_out, up_proj_out = tex.unpack_swiglu_pair(combined_out) + # The fused backward consumes interleaved C directly. Keep both residual + # slots as aliases instead of materializing an otherwise-unused unpack. + gate_proj_out = up_proj_out = combined_out - intermediate_shape = (sorted_x.shape[0], gate_proj_out.shape[-1]) + intermediate_shape = (sorted_x.shape[0], combined // 2) scaling_mode = fc2_quantizer_set.x.scaling_mode row_scale_size = scaling_mode.get_grouped_scale_shape( intermediate_shape, @@ -651,7 +682,13 @@ def _ffn_fwd_per_shard( casted_sorted_x.get_tensor(usage=TensorUsage.LHS_TRANS).checkpoint( fc1_quantizer_set.x ), - casted_wi.get_tensor(usage=TensorUsage.RHS_TRANS).checkpoint( + casted_wi.get_tensor( + usage=( + TensorUsage.LHS + if use_cudnn_jax_fusion and cudnn_native_weight_layout + else TensorUsage.RHS_TRANS + ) + ).checkpoint( fc1_quantizer_set.kernel ), ffn_activation_residual, @@ -683,6 +720,7 @@ def _ffn_bwd_per_shard( apply_topk_weights_early: bool, has_bias: bool, use_cudnn_jax_fusion: bool, + cudnn_native_weight_layout: bool, ): """Backward mirror of :func:`_ffn_fwd_per_shard`.""" group_sizes = local_group_sizes.reshape(-1).astype(jnp.int32) @@ -814,16 +852,28 @@ def _ffn_bwd_per_shard( d_sorted_x = tex.grouped_gemm( casted_d_combined.get_tensor(usage=TensorUsage.LHS), casted_wi_rhs_trans, - contracting_dims=((1,), (2,)), + contracting_dims=((1,), (1 if cudnn_native_weight_layout else 2,)), ) - d_wi_combined = tex.grouped_gemm( - casted_sorted_x_lhs_trans, - casted_d_combined.get_tensor(usage=TensorUsage.RHS), - contracting_dims=((0,), (0,)), - ) - if use_cudnn_jax_fusion: + if cudnn_native_weight_layout: + # dY^T @ X directly produces [E,2N,K], matching the persistent native + # parameter, without a post-GEMM transpose or de-interleave/repack. + d_wi_combined = tex.grouped_gemm( + casted_d_combined.get_tensor(usage=TensorUsage.LHS_TRANS), + casted_sorted_x_lhs_trans, + contracting_dims=((0,), (0,)), + ) + else: + d_wi_combined = tex.grouped_gemm( + casted_sorted_x_lhs_trans, + casted_d_combined.get_tensor(usage=TensorUsage.RHS), + contracting_dims=((0,), (0,)), + ) + if use_cudnn_jax_fusion and not cudnn_native_weight_layout: d_wi_gate, d_wi_up = tex.unpack_swiglu_pair(d_wi_combined) d_wi_combined = jnp.concatenate([d_wi_gate, d_wi_up], axis=-1) + elif not use_cudnn_jax_fusion and cudnn_native_weight_layout: + d_wi_gate, d_wi_up = jnp.split(d_wi_combined, 2, axis=-1) + d_wi_combined = tex.pack_swiglu_pair(d_wi_gate, d_wi_up).transpose(0, 2, 1) if has_bias: d_wi_combined_bias = tex.grouped_dbias(d_combined_for_bias, group_sizes) d_wi_0_bias, d_wi_1_bias = jnp.split(d_wi_combined_bias, 2, axis=-1) @@ -918,6 +968,7 @@ def _moe_fwd_rule( B, S, H = x.shape K = num_experts_per_tok + cudnn_native_weight_layout = wi.ndim == 3 and wi.shape[-1] == H if B % num_procs != 0: raise ValueError(f"batch={B} not divisible by ep*dp={num_procs}") @@ -1134,6 +1185,7 @@ def _ffn_fwd_body(*args): wi_0_checkpoint_name=wi_0_checkpoint_name, wi_1_checkpoint_name=wi_1_checkpoint_name, wo_checkpoint_name=wo_checkpoint_name, + cudnn_native_weight_layout=cudnn_native_weight_layout, ) expert_outputs, ffn_residuals = shard_map( @@ -1212,6 +1264,7 @@ def _ffn_fwd_body(*args): "has_bias": has_bias, "x_shape": x.shape, "recv_pr": recv_pr, + "cudnn_native_weight_layout": cudnn_native_weight_layout, } # total_recv_tokens is a non-differentiable overflow signal (see moe()). return (output, aux_loss, total_recv_tokens), (ctx, static) @@ -1335,6 +1388,7 @@ def _ffn_bwd_body(*args): apply_topk_weights_early=apply_topk_weights_early, has_bias=has_bias, use_cudnn_jax_fusion=use_cudnn_jax_fusion, + cudnn_native_weight_layout=static["cudnn_native_weight_layout"], ) ( d_sorted_x_local, From 68d6ef0af4f6865e057b461fbc73286c488822fb Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Thu, 1 Oct 2026 14:24:54 -0700 Subject: [PATCH 19/38] Reject cuDNN-native MoE weights without fusion Signed-off-by: Jeremy Berchtold --- transformer_engine/jax/moe.py | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index 89be852785d..c077a7074ea 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -1806,6 +1806,7 @@ def moe( expert_bias_arg = expert_bias.astype(jnp.float32) use_cudnn_jax_fusion = False + cudnn_native_weight_layout = wi.ndim == 3 and wi.shape[-1] == x.shape[-1] if _use_cudnn_cutedsl_fusion_from_env(): rejection_reasons = _cudnn_jax_fusion_rejection_reasons( x, @@ -1818,6 +1819,12 @@ def moe( ep_axis=ep_axis, ) if rejection_reasons: + if cudnn_native_weight_layout: + raise ValueError( + "cuDNN-native MoE weight layout requires the fused cuDNN grouped-GEMM " + "path, which is unsupported for this moe() call: " + + "; ".join(rejection_reasons) + ) warnings.warn( f"{_CUDNN_JAX_ENV}=1 is unsupported for this moe() call; falling back to " "the regular TE grouped-GEMM path: " + "; ".join(rejection_reasons), @@ -1826,6 +1833,11 @@ def moe( ) else: use_cudnn_jax_fusion = True + elif cudnn_native_weight_layout: + raise ValueError( + "cuDNN-native MoE weight layout requires the fused cuDNN grouped-GEMM path; " + f"set {_CUDNN_JAX_ENV}=1." + ) output, aux_loss, total_recv_tokens = _moe( x, From 0d2326bd41d323dd21ce4a4686e30d7a7e8eb052 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Thu, 1 Oct 2026 14:39:00 -0700 Subject: [PATCH 20/38] Remove redundant MoE combine checkpoint name Signed-off-by: Jeremy Berchtold --- tests/jax/test_te_ep_moe.py | 6 ------ transformer_engine/jax/flax/moe.py | 8 +++----- transformer_engine/jax/moe.py | 15 +-------------- 3 files changed, 4 insertions(+), 25 deletions(-) diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index 77b5990ab32..953f87950f9 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -367,7 +367,6 @@ def _make_block( input_axes=("batch", None, None), quantization_recipe=None, dispatch_checkpoint_name=None, - combine_checkpoint_name=None, ): kwargs = dict( num_experts=NUM_EXPERTS, @@ -382,7 +381,6 @@ def _make_block( input_axes=input_axes, quantization_recipe=quantization_recipe, dispatch_checkpoint_name=dispatch_checkpoint_name, - combine_checkpoint_name=combine_checkpoint_name, ) # Custom expert_bias_init lets tests inject a non-zero expert_bias without # poking variables['params'] post-init. @@ -719,14 +717,12 @@ def recorded_checkpoint_name(value, name): monkeypatch.setattr(moe_module, "checkpoint_name", recorded_checkpoint_name) block = _make_block( dispatch_checkpoint_name="saved_dispatch", - combine_checkpoint_name="saved_combine", ) inputs = _make_inputs(jax.random.PRNGKey(34)) variables, output, _ = _init_apply(block, mesh, inputs, jax.random.PRNGKey(35)) grads, grad_inputs = _grad_step(block, variables, mesh, inputs) assert len(named_values["saved_dispatch"]) >= 2 - assert len(named_values["saved_combine"]) >= 1 assert all( hasattr(value, "shape") for values in named_values.values() for value in values ) @@ -903,7 +899,6 @@ def checkpointed_moe(*args, **kwargs): block = _make_block( quantization_recipe=MXFP8BlockScaling(), dispatch_checkpoint_name="saved_dispatch", - combine_checkpoint_name="saved_combine", ) inputs = _make_inputs(jax.random.PRNGKey(32)) variables, output, _ = _init_apply(block, mesh, inputs, jax.random.PRNGKey(33)) @@ -916,7 +911,6 @@ def checkpointed_moe(*args, **kwargs): assert len(named_values["moe_mlpwo"]) == len(selected_calls) assert all(hasattr(value, "shape") for value in named_values["moe_mlpwo"]) assert len(named_values["saved_dispatch"]) >= 2 * len(selected_calls) - assert len(named_values["saved_combine"]) >= len(selected_calls) assert np.all(np.isfinite(_to_global_numpy(output, mesh))) assert np.all(np.isfinite(_to_global_numpy(grad_inputs, mesh))) for name in ("gate_kernel", "wi", "wo"): diff --git a/transformer_engine/jax/flax/moe.py b/transformer_engine/jax/flax/moe.py index 27a66aff87c..5d729f552f7 100644 --- a/transformer_engine/jax/flax/moe.py +++ b/transformer_engine/jax/flax/moe.py @@ -105,9 +105,9 @@ class _MoEBlock(TransformerEngineBase): recv_capacity_per_rank : Optional[int] Exact aligned receive capacity per EP rank. ``None`` reserves the dropless worst case. - dispatch_checkpoint_name, combine_checkpoint_name : Optional[str] - JAX rematerialization checkpoint names for the EP dispatch outputs - and combine output, respectively. ``None`` leaves them unnamed. + dispatch_checkpoint_name : Optional[str] + JAX rematerialization checkpoint name for the EP dispatch outputs. + ``None`` leaves them unnamed. The per-expert dispatch-slot alignment is fixed internally at 128 tokens (see ``moe._ALIGN_SIZE``) -- the value required by NCCL EP @@ -155,7 +155,6 @@ class _MoEBlock(TransformerEngineBase): apply_topk_weights_early: bool = False recv_capacity_per_rank: Optional[int] = None dispatch_checkpoint_name: Optional[str] = None - combine_checkpoint_name: Optional[str] = None # Dtypes / init / misc dtype: DType = jnp.float32 @@ -317,5 +316,4 @@ def make_grouped_quantizer_set(postfix): wo_kernel_axes=self.wo_kernel_axes, dtype=self.dtype, dispatch_checkpoint_name=self.dispatch_checkpoint_name, - combine_checkpoint_name=self.combine_checkpoint_name, ) diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index c077a7074ea..92e51226a81 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -932,7 +932,6 @@ def _moe_fwd_rule( wi_1_checkpoint_name, wo_checkpoint_name, dispatch_checkpoint_name, - combine_checkpoint_name, ): """Forward: gate -> topk -> ep_dispatch -> FFN -> ep_combine. @@ -1222,8 +1221,6 @@ def _ffn_fwd_body(*args): num_local_tokens=(B, S), out_partition_spec=out_partition_spec, ) - if combine_checkpoint_name is not None: - output = checkpoint_name(output, combine_checkpoint_name) # output of MLP should be sharded the same way as the activation input output = with_sharding_constraint_by_logical_axes(output, input_axes) @@ -1294,7 +1291,6 @@ def _moe_bwd_rule( wi_1_checkpoint_name, wo_checkpoint_name, dispatch_checkpoint_name, - combine_checkpoint_name, residuals, cotangents, ): @@ -1308,7 +1304,6 @@ def _moe_bwd_rule( wi_1_checkpoint_name, wo_checkpoint_name, dispatch_checkpoint_name, - combine_checkpoint_name, ) # captured / unused in bwd from jax.experimental.shard_map import shard_map @@ -1570,7 +1565,7 @@ def _ffn_bwd_body(*args): # ============================================================================= -@partial(jax.custom_vjp, nondiff_argnums=tuple(range(9, 33))) +@partial(jax.custom_vjp, nondiff_argnums=tuple(range(9, 32))) def _moe( x, gate_kernel, @@ -1604,7 +1599,6 @@ def _moe( wi_1_checkpoint_name, wo_checkpoint_name, dispatch_checkpoint_name, - combine_checkpoint_name, ): primal, _ = _moe_fwd_rule( x, @@ -1639,7 +1633,6 @@ def _moe( wi_1_checkpoint_name, wo_checkpoint_name, dispatch_checkpoint_name, - combine_checkpoint_name, ) return primal @@ -1683,7 +1676,6 @@ def moe( wi_1_checkpoint_name: Optional[str] = None, wo_checkpoint_name: Optional[str] = None, dispatch_checkpoint_name: Optional[str] = None, - combine_checkpoint_name: Optional[str] = None, ) -> Tuple[jnp.ndarray, Optional[jnp.ndarray], jnp.ndarray]: """Run a full MoE block under a single fused custom_vjp on the TE EP path. @@ -1732,10 +1724,6 @@ def moe( dispatch_checkpoint_name : Optional[str] JAX rematerialization checkpoint name for the EP dispatch outputs. ``None`` leaves these values unnamed; the opaque prepare handle is never named. - combine_checkpoint_name : Optional[str] - JAX rematerialization checkpoint name for the EP combine output. - ``None`` leaves the value unnamed. - Note that the per-expert dispatch-slot alignment is fixed internally at 128 tokens (``_ALIGN_SIZE``); see that constant's docstring for rationale and how to extend if a future recipe needs >128. @@ -1872,7 +1860,6 @@ def moe( wi_1_checkpoint_name, wo_checkpoint_name, dispatch_checkpoint_name, - combine_checkpoint_name, ) if aux_loss_coeff <= 0.0: aux_loss = None From 791e8ae7bf47fb13896478596502a9d6f0c2c0a7 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Mon, 5 Oct 2026 11:22:56 -0700 Subject: [PATCH 21/38] Use MeshResource for JAX MoE parallelism and quantized weight gathering --- docs/api/jax.rst | 36 + tests/jax/test_moe_resource_api.py | 170 ++ tests/jax/test_te_ep_moe.py | 77 +- transformer_engine/jax/cpp_extensions/ep.py | 177 ++- transformer_engine/jax/flax/moe.py | 151 +- transformer_engine/jax/moe.py | 1551 ++++++++++--------- 6 files changed, 1327 insertions(+), 835 deletions(-) create mode 100644 tests/jax/test_moe_resource_api.py diff --git a/docs/api/jax.rst b/docs/api/jax.rst index 7a31c9d3797..947d22a2ea7 100644 --- a/docs/api/jax.rst +++ b/docs/api/jax.rst @@ -59,3 +59,39 @@ Modules :members: __call__ .. autoapifunction:: transformer_engine.jax.flax.extend_logical_axis_rules + +MoE parallelism resources +------------------------- + +The experimental ``transformer_engine.jax.moe.moe`` function and Flax +``_MoEBlock`` accept ``mesh_resource=None`` and ``quant_before_fsdp_ag=False``. +Pass a ``MeshResource`` explicitly, or provide one through +``global_shard_guard``. An explicit resource takes precedence; a missing +resource raises ``ValueError``. The resource must name an EP axis. DP and +FSDP are outer batch-sharding axes, in that order, with EP innermost. +Unset outer axes are omitted, and a shared DP/FSDP axis is included once. + +For example, with an active physical mesh containing ``dp``, ``fsdp`` and +``ep`` axes:: + + from transformer_engine.jax.flax import _MoEBlock + from transformer_engine.jax.sharding import MeshResource + + block = _MoEBlock( + mesh_resource=MeshResource( + dp_resource="dp", fsdp_resource="fsdp", ep_resource="ep" + ), + quant_before_fsdp_ag=True, + # Supply an MXFP8 quantization recipe and the model dimensions. + ) + +``quant_before_fsdp_ag=True`` quantizes local expert-weight shards before +all-gathering their MXFP8 data and scales on the resource's FSDP axis. +It requires an FSDP resource and MXFP8 kernel quantizers. The default +gathers full-precision weights before quantization. + +The old ``ep_axis``, ``data_parallelism_axes`` and ``weight_gather`` arguments +remain accepted with a ``DeprecationWarning``. They are translated into a +resource and the boolean before calling the new API. Conflicting old and +new arguments raise ``ValueError``. ``WeightGather`` remains available only +for this compatibility path; new callers should use the boolean. diff --git a/tests/jax/test_moe_resource_api.py b/tests/jax/test_moe_resource_api.py new file mode 100644 index 00000000000..4265319d532 --- /dev/null +++ b/tests/jax/test_moe_resource_api.py @@ -0,0 +1,170 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""MoE resource resolution and deprecated API compatibility without EP kernels.""" + +import importlib +import inspect +import warnings + +import jax +import jax.numpy as jnp +import numpy as np +import pytest +from jax.sharding import Mesh + +from transformer_engine.jax.flax import _MoEBlock +from transformer_engine.jax.moe import WeightGather, _moe_mesh_axes, _resolve_moe_mesh_resource +from transformer_engine.jax.sharding import MeshResource, global_mesh_resource, global_shard_guard + + +@pytest.fixture(autouse=True) +def no_global_resource(): + with global_shard_guard(None): + yield + + +def test_resource_required_for_function_and_block(): + module = importlib.import_module("transformer_engine.jax.moe") + with pytest.raises(ValueError, match="active global_shard_guard"): + module.moe(None, None, None, None, num_experts=2, num_experts_per_tok=1) + with pytest.raises(ValueError, match="active global_shard_guard"): + _MoEBlock().init(jax.random.PRNGKey(0), jnp.ones((1, 1, 4))) + + +def test_explicit_resource_overrides_global_and_is_snapshotted(): + explicit = MeshResource(dp_resource="dp", fsdp_resource="fsdp", ep_resource="ep") + enclosing = MeshResource(ep_resource="other") + with global_shard_guard(enclosing): + resolved, quantize = _resolve_moe_mesh_resource(explicit, True) + assert global_mesh_resource() is enclosing + explicit.ep_resource = "changed" + assert _moe_mesh_axes(resolved) == ("ep", ("dp", "fsdp")) + assert quantize is True + + +def test_global_default_and_duplicate_outer_axis(): + with global_shard_guard( + MeshResource(dp_resource="fsdp", fsdp_resource="fsdp", ep_resource="ep") + ): + resource, quantize = _resolve_moe_mesh_resource() + assert _moe_mesh_axes(resource) == ("ep", ("fsdp",)) + assert quantize is False + + +def test_ep_partitioning_retains_bind_time_resources(): + from transformer_engine.jax.cpp_extensions import ep + from jax.sharding import PartitionSpec as P + + with global_shard_guard(MeshResource(dp_resource="dp", fsdp_resource="fsdp", ep_resource="ep")): + captured = ep._capture_ep_resource_axes() + spec = P(("dp", "fsdp", "ep"), None, None) + assert ep._leading_axis_ok(spec, captured) == (True, "ep", ("dp", "fsdp")) + assert ep._ep_output_spec(None, None, resource_axes=captured) == spec + assert ep._ep_spec_ok(spec, 2, resource_axes=captured) + assert not ep._ep_spec_ok(P(("wrong", "ep"), None, None), 2, resource_axes=captured) + + +@pytest.mark.parametrize( + "resource,quantize,error", + [ + (MeshResource(), False, "ep_resource"), + (MeshResource(ep_resource="ep", dp_resource="ep"), False, "distinct"), + (MeshResource(ep_resource="ep"), True, "fsdp_resource"), + ], +) +def test_invalid_resource(resource, quantize, error): + with pytest.raises(ValueError, match=error): + _resolve_moe_mesh_resource(resource, quantize) + + +def test_bool_and_resource_types(): + with pytest.raises(TypeError, match="must be a bool"): + _resolve_moe_mesh_resource(MeshResource(ep_resource="ep"), "true") + with pytest.raises(TypeError, match="must be a MeshResource"): + _resolve_moe_mesh_resource("ep") + + +def test_legacy_preserves_arbitrary_outer_axis_order(): + with pytest.warns(DeprecationWarning): + resource, quantize = _resolve_moe_mesh_resource( + ep_axis="ep", + data_parallelism_axes=("outer", "fsdp", "replica"), + weight_gather=WeightGather.quantized(axis="fsdp"), + ) + assert _moe_mesh_axes(resource) == ("ep", ("outer", "fsdp", "replica")) + assert resource.fsdp_resource == "fsdp" + assert quantize is True + + +def test_legacy_defaults_to_no_outer_axes(): + with global_shard_guard(MeshResource(dp_resource="dp", fsdp_resource="fsdp", ep_resource="ep")): + with pytest.warns(DeprecationWarning): + resource, _ = _resolve_moe_mesh_resource(ep_axis="ep") + assert _moe_mesh_axes(resource) == ("ep", ()) + + +@pytest.mark.parametrize( + "kwargs,error", + [ + ({"ep_axis": "other"}, "ep_axis conflicts"), + ({"data_parallelism_axes": ("other",)}, "data_parallelism_axes conflicts"), + ({"weight_gather": WeightGather.quantized(axis="other")}, "weight_gather axis conflicts"), + ], +) +def test_conflicting_old_and_new_args(kwargs, error): + with pytest.warns(DeprecationWarning): + with pytest.raises(ValueError, match=error): + _resolve_moe_mesh_resource( + MeshResource(fsdp_resource="fsdp", ep_resource="ep"), **kwargs + ) + + +@pytest.mark.parametrize("legacy", [False, True]) +def test_public_api_delegates_with_selected_resource(monkeypatch, legacy): + module = importlib.import_module("transformer_engine.jax.moe") + original_moe = module.moe + signature = inspect.signature(module._moe) + captured = {} + + def fake_vjp(*args): + captured.update(signature.bind(*args).arguments) + assert global_mesh_resource().ep_resource == "ep" + return args[0], None, jnp.zeros((1,), jnp.int32) + + def delegated_moe(*args, **kwargs): + assert "mesh_resource" in kwargs + assert not {"ep_axis", "data_parallelism_axes", "weight_gather"}.intersection(kwargs) + return original_moe(*args, **kwargs) + + monkeypatch.setattr(module, "_moe", fake_vjp) + monkeypatch.setattr(module, "moe", delegated_moe) + mesh = Mesh(np.asarray(jax.devices()[:1]).reshape(1, 1, 1), ("dp", "fsdp", "ep")) + kwargs = dict(num_experts=2, num_experts_per_tok=1) + if legacy: + kwargs.update( + ep_axis="ep", + data_parallelism_axes=("dp", "fsdp"), + weight_gather=WeightGather.quantized(axis="fsdp"), + ) + else: + kwargs.update( + mesh_resource=MeshResource(dp_resource="dp", fsdp_resource="fsdp", ep_resource="ep"), + quant_before_fsdp_ag=True, + ) + with jax.set_mesh(mesh), warnings.catch_warnings(record=True) as recorded: + warnings.simplefilter("always", DeprecationWarning) + original_moe( + jnp.ones((1, 1, 4)), + jnp.ones((4, 2)), + jnp.ones((2, 4, 8)), + jnp.ones((2, 4, 4)), + **kwargs, + ) + assert any("deprecated for TE MoE" in str(w.message) for w in recorded) == legacy + assert _moe_mesh_axes(captured["mesh_resource"]) == ("ep", ("dp", "fsdp")) + assert captured["quant_before_fsdp_ag"] is True + assert "ep_axis" not in captured + with pytest.raises(AssertionError, match="Global mesh resource is not set"): + global_mesh_resource() diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index a578e62b1c9..8c949b5644f 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -365,13 +365,14 @@ def _make_block( input_axes=("batch", None, None), quantization_recipe=None, dispatch_checkpoint_name=None, - weight_gather=WeightGather.full_precision(), + quant_before_fsdp_ag=False, + mesh_resource=None, ): kwargs = dict( num_experts=NUM_EXPERTS, num_experts_per_tok=TOPK, intermediate_size=INTER, - data_parallelism_axes=(FSDP_AXIS,), + mesh_resource=mesh_resource, apply_topk_weights_early=apply_topk_weights_early, aux_loss_coeff=aux_loss_coeff, use_expert_routing_bias=use_expert_routing_bias, @@ -380,7 +381,7 @@ def _make_block( input_axes=input_axes, quantization_recipe=quantization_recipe, dispatch_checkpoint_name=dispatch_checkpoint_name, - weight_gather=weight_gather, + quant_before_fsdp_ag=quant_before_fsdp_ag, ) # Custom expert_bias_init lets tests inject a non-zero expert_bias without # poking variables['params'] post-init. @@ -607,9 +608,7 @@ def native_moe(*args, **kwargs): monkeypatch.setattr(flax_moe_module, "moe", native_moe) x = _make_inputs(jax.random.PRNGKey(41)) baseline = _make_block(quantization_recipe=MXFP8BlockScaling()) - with _ctx(mesh): - gather_policy = WeightGather.quantized() - quantized_ag = _make_block(quantization_recipe=MXFP8BlockScaling(), weight_gather=gather_policy) + quantized_ag = _make_block(quantization_recipe=MXFP8BlockScaling(), quant_before_fsdp_ag=True) variables, baseline_out, _ = _init_apply(baseline, mesh, x, jax.random.PRNGKey(42)) with _ctx(mesh): x_sh = _shard_inputs(x, mesh) @@ -635,6 +634,70 @@ def native_moe(*args, **kwargs): ) +@pytest.mark.parametrize("quant_before_fsdp_ag", [False, True]) +@pytest.mark.parametrize("api", ["explicit", "legacy"]) +def test_mesh_resource_api_forward_and_backward(mesh, quant_before_fsdp_ag, api): + """Explicit resources need no global context; legacy calls retain numerical semantics.""" + baseline = _make_block( + quantization_recipe=MXFP8BlockScaling(), quant_before_fsdp_ag=quant_before_fsdp_ag + ) + if api == "explicit": + candidate = baseline.clone( + mesh_resource=MeshResource(ep_resource=EP_AXIS, fsdp_resource=FSDP_AXIS) + ) + else: + candidate = baseline.clone( + data_parallelism_axes=(FSDP_AXIS,), + quant_before_fsdp_ag=False, + weight_gather=( + WeightGather.quantized(axis=FSDP_AXIS) + if quant_before_fsdp_ag + else WeightGather.full_precision() + ), + ) + x = _make_inputs(jax.random.PRNGKey(51)) + variables, baseline_out, _ = _init_apply(baseline, mesh, x, jax.random.PRNGKey(52)) + baseline_grads, baseline_dx = _grad_step(baseline, variables, mesh, x) + with jax.set_mesh(mesh), nn_partitioning.axis_rules(LOGICAL_AXIS_RULES): + resource = None if api == "explicit" else MeshResource(ep_resource=EP_AXIS) + with global_shard_guard(resource): + x_sh = _shard_inputs(x, mesh) + if api == "legacy": + with pytest.warns(DeprecationWarning, match="deprecated for TE MoE"): + candidate_out, _, _ = jax.jit(candidate.apply)(variables, x_sh) + else: + candidate_out, _, _ = jax.jit(candidate.apply)(variables, x_sh) + + def loss_fn(variables, inputs): + output, _, _ = candidate.apply(variables, inputs) + return jnp.mean(output.astype(jnp.float32) ** 2) + + candidate_grads, candidate_dx = jax.jit(jax.grad(loss_fn, argnums=(0, 1)))( + variables, x_sh + ) + jax.block_until_ready((candidate_out, candidate_grads, candidate_dx)) + np.testing.assert_allclose( + _to_global_numpy(candidate_out, mesh).astype(np.float32), + _to_global_numpy(baseline_out, mesh).astype(np.float32), + **FWD_TOLERANCE["mxfp8"], + ) + for name in ("gate_kernel", "wi", "wo"): + np.testing.assert_allclose( + _to_global_numpy(_unwrap(candidate_grads["params"][name]), mesh).astype(np.float32), + _to_global_numpy(_unwrap(baseline_grads["params"][name]), mesh).astype(np.float32), + **( + GRAD_GATE_TOLERANCE["mxfp8"] + if name == "gate_kernel" + else GRAD_FFN_TOLERANCE["mxfp8"] + ), + ) + np.testing.assert_allclose( + _to_global_numpy(candidate_dx, mesh).astype(np.float32), + _to_global_numpy(baseline_dx, mesh).astype(np.float32), + **GRAD_FFN_TOLERANCE["mxfp8"], + ) + + def _reference_kwargs_from_config(config, params_np): """Pick out the reference-relevant pieces of a parametrize config.""" return dict( @@ -798,7 +861,7 @@ class TestTeEpMoeCudnnCutedslFusion: def test_mxfp8_forward_and_backward(self, mesh, apply_topk_weights_early, monkeypatch): if not _use_cudnn_cutedsl_fusion_from_env(): pytest.skip( - "run separately with " "NVTE_JAX_TEMP_FLAG_FOR_ABHINAV_CUDNN_GROUPED_GEMM_FUSION=1" + "run separately with NVTE_JAX_TEMP_FLAG_FOR_ABHINAV_CUDNN_GROUPED_GEMM_FUSION=1" ) rubin_calls = [] if get_device_compute_capability(0) == 107: diff --git a/transformer_engine/jax/cpp_extensions/ep.py b/transformer_engine/jax/cpp_extensions/ep.py index 9e2cf2b67c5..5556d0b665f 100644 --- a/transformer_engine/jax/cpp_extensions/ep.py +++ b/transformer_engine/jax/cpp_extensions/ep.py @@ -196,15 +196,24 @@ 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 _capture_ep_resource_axes(): + """Capture physical axes at bind time for asynchronous SPMD partitioning.""" + resource = global_mesh_resource() + outer_axes = getattr(resource, "_legacy_data_parallelism_axes", None) + if outer_axes is None: + outer_axes = tuple( + dict.fromkeys( + axis for axis in (resource.dp_resource, resource.fsdp_resource) if axis is not None + ) + ) + return resource.ep_resource, outer_axes - Leading dim is ``ep`` or a tuple ending in ``ep`` (outer dp/fsdp axes - first); all other dims 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) + +def _leading_axis_ok(spec, resource_axes=None): + """Validate replicated trailing dimensions and EP plus optional outer axes.""" + ep_axis, outer_axes = ( + resource_axes if resource_axes is not None else _capture_ep_resource_axes() + ) 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:]): @@ -218,20 +227,13 @@ def _leading_axis_ok(spec): def _ep_outer_axis(): - """The single dp/fsdp axis (if any) sitting outside ep on EP-output tensors. - - When set, EP-output globals carry an extra leading ``dp_size`` dim so SPMD - sees each DP color's slab as distinct (rather than replicated across DP). - - A dp/fsdp axis that is sized 1 in the active mesh is treated as absent so - we don't pin EP-output specs to a degenerate axis that JAX may collapse. - """ - gsr = global_mesh_resource() - if gsr.dp_resource is not None and get_mesh_axis_size(gsr.dp_resource) > 1: - return gsr.dp_resource - if gsr.fsdp_resource is not None and get_mesh_axis_size(gsr.fsdp_resource) > 1: - return gsr.fsdp_resource - return gsr.dp_resource or gsr.fsdp_resource + """Legacy single outer axis used by the standalone EP custom VJP wrapper.""" + resource = global_mesh_resource() + if resource.dp_resource is not None and get_mesh_axis_size(resource.dp_resource) > 1: + return resource.dp_resource + if resource.fsdp_resource is not None and get_mesh_axis_size(resource.fsdp_resource) > 1: + return resource.fsdp_resource + return resource.dp_resource or resource.fsdp_resource def _ep_leading_dims(is_outer): @@ -243,24 +245,20 @@ def _ep_leading_dims(is_outer): return (cfg.num_ep_groups * cfg.ep_size,) -def _ep_output_spec(*trailing): - """PartitionSpec for an EP-output tensor: ``(("dp","ep"), *trailing)`` when - DP is set (compound leading axis on a single dim), else ``("ep",*trailing)``.""" - gsr = global_mesh_resource() - outer = _ep_outer_axis() - if outer is None: - return PartitionSpec(gsr.ep_resource, *trailing) - return PartitionSpec((outer, gsr.ep_resource), *trailing) +def _ep_output_spec(*trailing, resource_axes=None): + """Output sharding uses every DP/FSDP outer axis, with EP innermost.""" + ep_axis, outer_axes = ( + resource_axes if resource_axes is not None else _capture_ep_resource_axes() + ) + leading = (*outer_axes, ep_axis) if outer_axes else ep_axis + return PartitionSpec(leading, *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. - """ - gsr = global_mesh_resource() - ep_axis = gsr.ep_resource - outer = _ep_outer_axis() +def _ep_spec_ok(spec, trailing_count, resource_axes=None): + """Validate EP-output axes, allowing JAX to drop size-one mesh axes.""" + ep_axis, outer_axes = ( + resource_axes if resource_axes is not None else _capture_ep_resource_axes() + ) if len(spec) != 1 + trailing_count: return False if any(ax is not None for ax in spec[1:]): @@ -268,8 +266,7 @@ def _ep_spec_ok(spec, trailing_count): leading = spec[0] elts = leading if isinstance(leading, tuple) else (leading,) actual = frozenset(a for a in elts if a is not None) - expected = {ep_axis} if outer is None else {ep_axis, outer} - return actual <= expected + return actual <= set((*outer_axes, ep_axis)) # ── ep_prepare ────────────────────────────────────────────────────────────── @@ -280,14 +277,17 @@ class EpPreparePrimitive(BasePrimitive): name = "te_ep_prepare_ffi" multiple_results = True - impl_static_args = (1, 2, 3) # top_k, dispatch_output_per_expert_alignment, is_outer + impl_static_args = (1, 2, 3, 4) # top_k, dispatch_output_per_expert_alignment, is_outer inner_primitive = None outer_primitive = None @staticmethod - def abstract(topk_idx_aval, *, top_k, dispatch_output_per_expert_alignment, is_outer): + def abstract( + topk_idx_aval, *, top_k, dispatch_output_per_expert_alignment, is_outer, resource_axes + ): # is_outer=True: global leading dim = (dp*ep,) (or (ep,) with no DP); # False: per-shard = (1,). + del resource_axes cfg = get_ep_config() num_local_experts = cfg.num_local_experts assert ( @@ -311,7 +311,10 @@ def outer_abstract(*args, **kwargs): return EpPreparePrimitive.abstract(*args, **kwargs) # pylint: disable=missing-kwoa @staticmethod - def lowering(ctx, topk_idx, *, top_k, dispatch_output_per_expert_alignment, is_outer): + def lowering( + ctx, topk_idx, *, top_k, dispatch_output_per_expert_alignment, is_outer, resource_axes + ): + del resource_axes del is_outer return ffi.ffi_lowering(EpPreparePrimitive.name)( ctx, @@ -321,27 +324,42 @@ def lowering(ctx, topk_idx, *, top_k, dispatch_output_per_expert_alignment, is_o ) @staticmethod - def impl(topk_idx, top_k, dispatch_output_per_expert_alignment, is_outer): + def impl(topk_idx, top_k, dispatch_output_per_expert_alignment, is_outer, resource_axes): assert EpPreparePrimitive.inner_primitive is not None token_counts, total_recv_tokens, handle_mem = EpPreparePrimitive.inner_primitive.bind( topk_idx, top_k=top_k, dispatch_output_per_expert_alignment=dispatch_output_per_expert_alignment, is_outer=is_outer, + resource_axes=resource_axes, ) return token_counts, total_recv_tokens, handle_mem @staticmethod - def batcher(batched_args, batch_dims, *, top_k, dispatch_output_per_expert_alignment, is_outer): + def batcher( + batched_args, + batch_dims, + *, + top_k, + dispatch_output_per_expert_alignment, + is_outer, + resource_axes, + ): raise NotImplementedError("EpPreparePrimitive does not support vmap") @staticmethod def partition( - top_k, dispatch_output_per_expert_alignment, is_outer, mesh, arg_infos, result_infos + top_k, + dispatch_output_per_expert_alignment, + is_outer, + resource_axes, + mesh, + arg_infos, + result_infos, ): del is_outer, result_infos idx_spec = arg_infos[0].sharding.spec - ok, ep_axis, outer_axes = _leading_axis_ok(idx_spec) + ok, ep_axis, outer_axes = _leading_axis_ok(idx_spec, resource_axes) if not ok: raise NotImplementedError( "EpPrepare: topk_idx leading dim must include ep_resource" @@ -358,7 +376,11 @@ def partition( def sharded_impl(topk_idx): return EpPreparePrimitive.impl( - topk_idx, top_k, dispatch_output_per_expert_alignment, False + topk_idx, + top_k, + dispatch_output_per_expert_alignment, + False, + resource_axes=resource_axes, ) return mesh, sharded_impl, (tc_sharding, trt_sharding, hm_sharding), arg_shardings @@ -384,7 +406,7 @@ class EpDispatchPrimitive(BasePrimitive): name = "te_ep_dispatch_ffi" multiple_results = True - impl_static_args = (4, 5, 6, 7) # top_k, dispatch_output_per_expert_alignment, + impl_static_args = (4, 5, 6, 7, 8) # top_k, dispatch_output_per_expert_alignment, # recv_capacity_per_rank, is_outer inner_primitive = None outer_primitive = None @@ -400,9 +422,11 @@ def abstract( dispatch_output_per_expert_alignment, recv_capacity_per_rank, is_outer, + resource_axes, ): # is_outer=True: global leading dim = (dp*ep,) (or (ep,) with no DP); # False: per-shard = (1,). + del resource_axes del topk_idx_aval, topk_weights_aval, top_k, dispatch_output_per_expert_alignment del handle_mem_aval assert ( @@ -433,7 +457,9 @@ def lowering( dispatch_output_per_expert_alignment, recv_capacity_per_rank, is_outer, + resource_axes, ): + del resource_axes del recv_capacity_per_rank, is_outer return ffi.ffi_lowering(EpDispatchPrimitive.name)( ctx, @@ -455,6 +481,7 @@ def impl( dispatch_output_per_expert_alignment, recv_capacity_per_rank, is_outer, + resource_axes, ): assert EpDispatchPrimitive.inner_primitive is not None recv_tokens, recv_topk_weights = EpDispatchPrimitive.inner_primitive.bind( @@ -466,6 +493,7 @@ def impl( dispatch_output_per_expert_alignment=dispatch_output_per_expert_alignment, recv_capacity_per_rank=recv_capacity_per_rank, is_outer=is_outer, + resource_axes=resource_axes, ) return recv_tokens, recv_topk_weights @@ -478,6 +506,7 @@ def batcher( dispatch_output_per_expert_alignment, recv_capacity_per_rank, is_outer, + resource_axes, ): raise NotImplementedError("EpDispatchPrimitive does not support vmap") @@ -487,13 +516,14 @@ def partition( dispatch_output_per_expert_alignment, recv_capacity_per_rank, is_outer, + resource_axes, mesh, arg_infos, result_infos, ): 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, resource_axes) if not ok: raise NotImplementedError( "EpDispatch: tokens leading dim must include ep_resource" @@ -525,6 +555,7 @@ def sharded_impl(handle_mem, topk_idx, tokens, topk_weights): dispatch_output_per_expert_alignment, recv_capacity_per_rank, False, + resource_axes=resource_axes, ) return mesh, sharded_impl, out_shardings, arg_shardings @@ -580,7 +611,7 @@ class EpCombinePrimitive(BasePrimitive): name = "te_ep_combine_ffi" multiple_results = False - impl_static_args = (2, 3, 4, 5) # top_k, dispatch_output_per_expert_alignment, + impl_static_args = (2, 3, 4, 5, 6) # top_k, dispatch_output_per_expert_alignment, # out_leading_shape, out_partition_spec inner_primitive = None outer_primitive = None @@ -594,7 +625,9 @@ def abstract( dispatch_output_per_expert_alignment, out_leading_shape, out_partition_spec, + resource_axes, ): + del resource_axes del top_k, dispatch_output_per_expert_alignment, out_partition_spec, handle_mem_aval assert ( len(expert_out_aval.shape) == 3 @@ -614,7 +647,9 @@ def lowering( dispatch_output_per_expert_alignment, out_leading_shape, out_partition_spec, + resource_axes, ): + del resource_axes del out_leading_shape, out_partition_spec return ffi.ffi_lowering(EpCombinePrimitive.name)( ctx, @@ -632,6 +667,7 @@ def impl( dispatch_output_per_expert_alignment, out_leading_shape, out_partition_spec, + resource_axes, ): assert EpCombinePrimitive.inner_primitive is not None return EpCombinePrimitive.inner_primitive.bind( @@ -641,6 +677,7 @@ def impl( dispatch_output_per_expert_alignment=dispatch_output_per_expert_alignment, out_leading_shape=out_leading_shape, out_partition_spec=out_partition_spec, + resource_axes=resource_axes, ) @staticmethod @@ -652,6 +689,7 @@ def batcher( dispatch_output_per_expert_alignment, out_leading_shape, out_partition_spec, + resource_axes, ): raise NotImplementedError("EpCombinePrimitive does not support vmap") @@ -661,13 +699,14 @@ def partition( dispatch_output_per_expert_alignment, out_leading_shape, out_partition_spec, + resource_axes, mesh, arg_infos, result_infos, ): 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, resource_axes=resource_axes): raise NotImplementedError( "EpCombine: expert_out must be sharded as PartitionSpec(ep_resource," " None, None) (or ((dp, ep), None, None) when dp/fsdp is set)" @@ -689,6 +728,7 @@ def sharded_impl(handle_mem, expert_out): dispatch_output_per_expert_alignment, per_shard_leading, out_partition_spec, + resource_axes=resource_axes, ) return mesh, sharded_impl, out_sharding, arg_shardings @@ -714,7 +754,7 @@ class EpDispatchBwdPrimitive(BasePrimitive): name = "te_ep_dispatch_bwd_ffi" multiple_results = True - impl_static_args = (3, 4, 5, 6) # top_k, dispatch_output_per_expert_alignment, + impl_static_args = (3, 4, 5, 6, 7) # top_k, dispatch_output_per_expert_alignment, # out_leading_shape, out_partition_spec inner_primitive = None outer_primitive = None @@ -729,7 +769,9 @@ def abstract( dispatch_output_per_expert_alignment, out_leading_shape, out_partition_spec, + resource_axes, ): + del resource_axes del dispatch_output_per_expert_alignment del g_recv_topk_weights_aval, out_partition_spec, handle_mem_aval assert ( @@ -754,7 +796,9 @@ def lowering( dispatch_output_per_expert_alignment, out_leading_shape, out_partition_spec, + resource_axes, ): + del resource_axes del out_leading_shape, out_partition_spec return ffi.ffi_lowering(EpDispatchBwdPrimitive.name)( ctx, @@ -774,6 +818,7 @@ def impl( dispatch_output_per_expert_alignment, out_leading_shape, out_partition_spec, + resource_axes, ): assert EpDispatchBwdPrimitive.inner_primitive is not None return EpDispatchBwdPrimitive.inner_primitive.bind( @@ -784,6 +829,7 @@ def impl( dispatch_output_per_expert_alignment=dispatch_output_per_expert_alignment, out_leading_shape=out_leading_shape, out_partition_spec=out_partition_spec, + resource_axes=resource_axes, ) @staticmethod @@ -795,6 +841,7 @@ def batcher( dispatch_output_per_expert_alignment, out_leading_shape, out_partition_spec, + resource_axes, ): raise NotImplementedError("EpDispatchBwdPrimitive does not support vmap") @@ -804,20 +851,21 @@ def partition( dispatch_output_per_expert_alignment, out_leading_shape, out_partition_spec, + resource_axes, mesh, arg_infos, result_infos, ): 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, resource_axes=resource_axes): raise NotImplementedError( "EpDispatchBwd: grad must be sharded as PartitionSpec(ep_resource," " None, None) (or ((dp, ep), None, None) when dp/fsdp is set)" f" over [num_procs, recv_pr, H]; got spec={g_spec}." ) gw_spec = arg_infos[2].sharding.spec - if not _ep_spec_ok(gw_spec, trailing_count=1): + if not _ep_spec_ok(gw_spec, trailing_count=1, resource_axes=resource_axes): raise NotImplementedError( "EpDispatchBwd: g_recv_topk_weights must be sharded as" " PartitionSpec(ep_resource, None) (or ((dp, ep), None) when dp/fsdp is set)" @@ -846,6 +894,7 @@ def sharded_impl(handle_mem, grad, g_recv_topk_weights): dispatch_output_per_expert_alignment, per_shard_leading, out_partition_spec, + resource_axes=resource_axes, ) return mesh, sharded_impl, out_shardings, arg_shardings @@ -871,7 +920,7 @@ class EpCombineBwdPrimitive(BasePrimitive): name = "te_ep_combine_bwd_ffi" multiple_results = False - impl_static_args = (2, 3, 4, 5) # top_k, dispatch_output_per_expert_alignment, + impl_static_args = (2, 3, 4, 5, 6) # top_k, dispatch_output_per_expert_alignment, # recv_capacity_per_rank, is_outer inner_primitive = None outer_primitive = None @@ -885,9 +934,11 @@ def abstract( dispatch_output_per_expert_alignment, recv_capacity_per_rank, is_outer, + resource_axes, ): # is_outer=True: global leading dim = (dp*ep,) (or (ep,) with no DP); # False: per-shard = (1,). + del resource_axes del top_k, dispatch_output_per_expert_alignment, handle_mem_aval assert ( len(grad_aval.shape) >= 2 @@ -912,7 +963,9 @@ def lowering( dispatch_output_per_expert_alignment, recv_capacity_per_rank, is_outer, + resource_axes, ): + del resource_axes del recv_capacity_per_rank, is_outer return ffi.ffi_lowering(EpCombineBwdPrimitive.name)( ctx, @@ -930,6 +983,7 @@ def impl( dispatch_output_per_expert_alignment, recv_capacity_per_rank, is_outer, + resource_axes, ): assert EpCombineBwdPrimitive.inner_primitive is not None return EpCombineBwdPrimitive.inner_primitive.bind( @@ -939,6 +993,7 @@ def impl( dispatch_output_per_expert_alignment=dispatch_output_per_expert_alignment, recv_capacity_per_rank=recv_capacity_per_rank, is_outer=is_outer, + resource_axes=resource_axes, ) @staticmethod @@ -950,6 +1005,7 @@ def batcher( dispatch_output_per_expert_alignment, recv_capacity_per_rank, is_outer, + resource_axes, ): raise NotImplementedError("EpCombineBwdPrimitive does not support vmap") @@ -959,6 +1015,7 @@ def partition( dispatch_output_per_expert_alignment, recv_capacity_per_rank, is_outer, + resource_axes, mesh, arg_infos, result_infos, @@ -966,7 +1023,7 @@ def partition( del is_outer, result_infos 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()) + out_sharding = NamedSharding(mesh, _ep_output_spec(resource_axes=resource_axes)) def sharded_impl(handle_mem, grad): return EpCombineBwdPrimitive.impl( @@ -976,6 +1033,7 @@ def sharded_impl(handle_mem, grad): dispatch_output_per_expert_alignment, recv_capacity_per_rank, False, + resource_axes=resource_axes, ) return mesh, sharded_impl, out_sharding, arg_shardings @@ -1005,6 +1063,7 @@ def ep_prepare(cfg: EpLayerConfig, topk_idx): top_k=int(cfg.top_k), dispatch_output_per_expert_alignment=int(cfg.dispatch_output_per_expert_alignment), is_outer=True, + resource_axes=_capture_ep_resource_axes(), ) @@ -1022,6 +1081,7 @@ def ep_dispatch_fwd( dispatch_output_per_expert_alignment=int(cfg.dispatch_output_per_expert_alignment), recv_capacity_per_rank=recv_capacity_per_rank, is_outer=True, + resource_axes=_capture_ep_resource_axes(), ) @@ -1038,6 +1098,7 @@ def ep_combine_fwd( dispatch_output_per_expert_alignment=int(cfg.dispatch_output_per_expert_alignment), out_leading_shape=out_leading, out_partition_spec=out_partition_spec, + resource_axes=_capture_ep_resource_axes(), ) @@ -1060,6 +1121,7 @@ def ep_dispatch_bwd( dispatch_output_per_expert_alignment=int(cfg.dispatch_output_per_expert_alignment), out_leading_shape=out_leading, out_partition_spec=out_partition_spec, + resource_axes=_capture_ep_resource_axes(), ) @@ -1073,4 +1135,5 @@ def ep_combine_bwd(cfg: EpLayerConfig, handle_mem, grad, recv_capacity_per_rank) dispatch_output_per_expert_alignment=int(cfg.dispatch_output_per_expert_alignment), recv_capacity_per_rank=recv_capacity_per_rank, is_outer=True, + resource_axes=_capture_ep_resource_axes(), ) diff --git a/transformer_engine/jax/flax/moe.py b/transformer_engine/jax/flax/moe.py index 1ddeb47108d..429cce08e24 100644 --- a/transformer_engine/jax/flax/moe.py +++ b/transformer_engine/jax/flax/moe.py @@ -34,10 +34,10 @@ from flax import linen as nn from transformer_engine.common.recipe import Recipe -from ..moe import WeightGather, moe +from ..moe import WeightGather, _moe_mesh_axes, _resolve_moe_mesh_resource, moe from ..quantize import QuantizerSet from ..router import ScoreFunction -from ..sharding import _get_mesh, get_active_resource_axis +from ..sharding import MeshResource, _get_mesh, global_shard_guard from .module import TransformerEngineBase PRNGKey = Any @@ -92,12 +92,17 @@ class _MoEBlock(TransformerEngineBase): Logical sharding axis tuples (consumed by Flax's :func:`with_logical_partitioning` and our internal :func:`with_sharding_constraint_by_logical_axes`). - data_parallelism_axes : tuple[str, ...] - FSDP axes over which the input *batch* dim is sharded IN - ADDITION to the EP axis. Empty (default) means activations are - replicated across non-EP axes within an EP group; set e.g. - ``("fsdp",)`` for true FSDP-of-batch where each device owns a - unique slice of the batch. + mesh_resource : Optional[MeshResource] + Physical DP, FSDP and EP mesh axes. ``None`` resolves the active global + MeshResource context. An explicit resource takes precedence; one of + these is required. Batch sharding uses DP and FSDP as outer axes and + EP innermost. + quant_before_fsdp_ag : bool + Quantize MXFP8 weight shards before their FSDP all-gather. Defaults to + ``False``; ``True`` requires ``mesh_resource.fsdp_resource``. + ep_axis, data_parallelism_axes, weight_gather : deprecated + Compatibility arguments converted into MeshResource and the boolean + with a DeprecationWarning. apply_topk_weights_early : bool If ``True``, multiply expert outputs by their top-k weights *inside* each shard before ``ep_combine`` (saves one global @@ -108,10 +113,6 @@ class _MoEBlock(TransformerEngineBase): dispatch_checkpoint_name : Optional[str] JAX rematerialization checkpoint name for the EP dispatch outputs. ``None`` leaves them unnamed. - weight_gather : WeightGather - Expert-weight gather policy. Defaults to a full-precision gather. - For MXFP8 gather, use ``WeightGather.quantized()`` to select the - active ``MeshResource.fsdp_resource``, or pass ``axis`` explicitly. The per-expert dispatch-slot alignment is fixed internally at 128 tokens (see ``moe._ALIGN_SIZE``) -- the value required by NCCL EP @@ -153,13 +154,17 @@ class _MoEBlock(TransformerEngineBase): input_axes: Tuple[Optional[str], ...] = () # Parallelism - data_parallelism_axes: Tuple[str, ...] = () + mesh_resource: Optional[MeshResource] = None + quant_before_fsdp_ag: bool = False + # Deprecated compatibility arguments. + ep_axis: Optional[str] = None + data_parallelism_axes: Optional[Tuple[str, ...]] = None # MoE knobs forwarded to ``moe()`` apply_topk_weights_early: bool = False recv_capacity_per_rank: Optional[int] = None dispatch_checkpoint_name: Optional[str] = None - weight_gather: WeightGather = WeightGather.full_precision() + weight_gather: Optional[WeightGather] = None # Dtypes / init / misc dtype: DType = jnp.float32 @@ -200,6 +205,14 @@ def __call__(self, inputs: Array) -> Tuple[Array, Optional[Array], Array]: Non-differentiable per-rank pre-drop recv-slot total; flags overflow when ``drop_on_overflow`` is set at ep_bootstrap. """ + mesh_resource, quant_before_fsdp_ag = _resolve_moe_mesh_resource( + self.mesh_resource, + self.quant_before_fsdp_ag, + self.ep_axis, + self.data_parallelism_axes, + self.weight_gather, + ) + _, data_parallelism_axes = _moe_mesh_axes(mesh_resource) assert ( inputs.ndim == 3 ), f"_MoEBlock expects [batch, sequence, hidden] input, got shape {inputs.shape}" @@ -262,64 +275,64 @@ def __call__(self, inputs: Array) -> Tuple[Array, Optional[Array], Array]: jnp.float32, ) - ep_axis = get_active_resource_axis("ep_resource") mesh = _get_mesh() data_parallel_size = 1 - for axis in self.data_parallelism_axes: + for axis in data_parallelism_axes: data_parallel_size *= mesh.shape[axis] - def make_grouped_quantizer_set(postfix): - # Dispatched token groups span every data-parallel replica, - # whereas expert kernels have one group per global expert. - token_set = self.generate_quantizer_set( - f"{postfix}_token", - fp8_recipe=self.quantization_recipe, - n_groups=data_parallel_size * self.num_experts, - ) - expert_set = self.generate_quantizer_set( - f"{postfix}_expert", - fp8_recipe=self.quantization_recipe, - n_groups=self.num_experts, - ) - return QuantizerSet( - x=token_set.x, - kernel=expert_set.kernel, - dgrad=token_set.dgrad, + with global_shard_guard(mesh_resource): + + def make_grouped_quantizer_set(postfix): + # Dispatched token groups span every data-parallel replica, + # whereas expert kernels have one group per global expert. + token_set = self.generate_quantizer_set( + f"{postfix}_token", + fp8_recipe=self.quantization_recipe, + n_groups=data_parallel_size * self.num_experts, + ) + expert_set = self.generate_quantizer_set( + f"{postfix}_expert", + fp8_recipe=self.quantization_recipe, + n_groups=self.num_experts, + ) + return QuantizerSet( + x=token_set.x, + kernel=expert_set.kernel, + dgrad=token_set.dgrad, + ) + + quantizer_sets = ( + make_grouped_quantizer_set("_fc1"), + make_grouped_quantizer_set("_fc2"), ) - quantizer_sets = ( - make_grouped_quantizer_set("_fc1"), - make_grouped_quantizer_set("_fc2"), - ) - - return moe( - inputs, - gate_kernel, - wi, - wo, - wi_0_bias, - wi_1_bias, - wo_bias, - expert_bias, - num_experts=self.num_experts, - num_experts_per_tok=self.num_experts_per_tok, - activation_type=self.activation_type, - score_function=self.score_function, - use_pre_softmax=self.use_pre_softmax, - num_groups=self.num_groups, - group_topk=self.group_topk, - scaling_factor=self.scaling_factor, - aux_loss_coeff=self.aux_loss_coeff, - apply_topk_weights_early=self.apply_topk_weights_early, - quantizer_sets=quantizer_sets, - recv_capacity_per_rank=self.recv_capacity_per_rank, - weight_gather=self.weight_gather, - ep_axis=ep_axis, - data_parallelism_axes=self.data_parallelism_axes, - input_axes=self.input_axes, - gate_kernel_axes=self.gate_kernel_axes, - wi_kernel_axes=self.wi_kernel_axes, - wo_kernel_axes=self.wo_kernel_axes, - dtype=self.dtype, - dispatch_checkpoint_name=self.dispatch_checkpoint_name, - ) + return moe( + inputs, + gate_kernel, + wi, + wo, + wi_0_bias, + wi_1_bias, + wo_bias, + expert_bias, + num_experts=self.num_experts, + num_experts_per_tok=self.num_experts_per_tok, + activation_type=self.activation_type, + score_function=self.score_function, + use_pre_softmax=self.use_pre_softmax, + num_groups=self.num_groups, + group_topk=self.group_topk, + scaling_factor=self.scaling_factor, + aux_loss_coeff=self.aux_loss_coeff, + apply_topk_weights_early=self.apply_topk_weights_early, + quantizer_sets=quantizer_sets, + recv_capacity_per_rank=self.recv_capacity_per_rank, + quant_before_fsdp_ag=quant_before_fsdp_ag, + mesh_resource=mesh_resource, + input_axes=self.input_axes, + gate_kernel_axes=self.gate_kernel_axes, + wi_kernel_axes=self.wi_kernel_axes, + wo_kernel_axes=self.wo_kernel_axes, + dtype=self.dtype, + dispatch_checkpoint_name=self.dispatch_checkpoint_name, + ) diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index e5d982c896f..cc0470ca85c 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -33,7 +33,7 @@ import math import os import warnings -from dataclasses import dataclass +from dataclasses import dataclass, fields, replace from functools import partial from typing import Any, Literal, Optional, Tuple, Union @@ -59,14 +59,16 @@ from .cpp_extensions.gemm import swizzled_scale from .flax.module import _convert_to_activation_function from .router import ScoreFunction, _validate_score_function -from .sharding import _get_mesh, global_mesh_resource +from .sharding import MeshResource, _get_mesh, global_mesh_resource, global_shard_guard __all__ = ["WeightGather", "get_moe_recv_capacity_per_rank", "moe"] @dataclass(frozen=True) class WeightGather: - """How MoE expert weights are gathered across a sharding mesh axis. + """Deprecated compatibility policy; use ``quant_before_fsdp_ag`` instead. + + How MoE expert weights are gathered across a sharding mesh axis. The quantization recipe determines the wire format for ``quantized``; this policy only chooses whether quantization precedes the gather. @@ -111,6 +113,129 @@ def quantized(cls, *, axis: Optional[str] = None) -> "WeightGather": return cls(mode="quantized", axis=axis) +@dataclass +class _LegacyMoEMeshResource(MeshResource): + """Preserve arbitrary ordered outer axes accepted by the deprecated API.""" + + _legacy_data_parallelism_axes: Tuple[str, ...] = () + + +def _moe_mesh_axes(resource: MeshResource): + """Resolve physical MoE axes, keeping EP innermost in the batch shard.""" + if not isinstance(resource.ep_resource, str) or not resource.ep_resource: + raise ValueError("TE MoE requires MeshResource.ep_resource to name a physical mesh axis.") + if isinstance(resource, _LegacyMoEMeshResource): + outer_axes = resource._legacy_data_parallelism_axes + else: + outer_axes = tuple( + dict.fromkeys( + axis for axis in (resource.dp_resource, resource.fsdp_resource) if axis is not None + ) + ) + if any(not isinstance(axis, str) or not axis for axis in outer_axes): + raise ValueError("TE MoE DP and FSDP resources must name physical mesh axes.") + if resource.ep_resource in outer_axes or len(set(outer_axes)) != len(outer_axes): + raise ValueError("TE MoE EP and outer data-parallel axes must be distinct.") + return resource.ep_resource, outer_axes + + +def _resolve_moe_mesh_resource( + mesh_resource=None, + quant_before_fsdp_ag=False, + ep_axis=None, + data_parallelism_axes=None, + weight_gather=None, +): + """Resolve the canonical API and adapt deprecated axes/gather arguments.""" + if not isinstance(quant_before_fsdp_ag, bool): + raise TypeError("quant_before_fsdp_ag must be a bool.") + if mesh_resource is not None and not isinstance(mesh_resource, MeshResource): + raise TypeError("mesh_resource must be a MeshResource or None.") + explicit_resource = mesh_resource is not None + if mesh_resource is None: + try: + mesh_resource = global_mesh_resource() + except AssertionError: + mesh_resource = None + + legacy = ep_axis is not None or data_parallelism_axes is not None or weight_gather is not None + if legacy: + warnings.warn( + "ep_axis, data_parallelism_axes, and weight_gather are deprecated for TE MoE; " + "pass mesh_resource=MeshResource(...) and quant_before_fsdp_ag instead.", + DeprecationWarning, + stacklevel=3, + ) + if weight_gather is not None and not isinstance(weight_gather, WeightGather): + raise TypeError("weight_gather must be a WeightGather policy.") + if weight_gather is not None: + if quant_before_fsdp_ag and weight_gather.mode != "quantized": + raise ValueError("quant_before_fsdp_ag conflicts with weight_gather.") + quant_before_fsdp_ag = weight_gather.mode == "quantized" + if explicit_resource: + if ep_axis is not None and ep_axis != mesh_resource.ep_resource: + raise ValueError("ep_axis conflicts with mesh_resource.ep_resource.") + if ( + data_parallelism_axes is not None + and tuple(data_parallelism_axes) != _moe_mesh_axes(mesh_resource)[1] + ): + raise ValueError("data_parallelism_axes conflicts with mesh_resource.") + if ( + weight_gather is not None + and weight_gather.axis is not None + and weight_gather.axis != mesh_resource.fsdp_resource + ): + raise ValueError("weight_gather axis conflicts with mesh_resource.fsdp_resource.") + elif ep_axis is not None or data_parallelism_axes is not None: + # The old functional API defaulted to no outer axes even in a global context. + axes = tuple(data_parallelism_axes or ()) + fsdp_axis = ( + weight_gather.axis + if weight_gather is not None and weight_gather.axis is not None + else getattr(mesh_resource, "fsdp_resource", None) + ) + if ( + weight_gather is not None + and weight_gather.axis is not None + and weight_gather.axis not in axes + ): + raise ValueError( + "Quantized weight all-gather requires its FSDP axis among the outer batch axes." + ) + if fsdp_axis not in axes: + fsdp_axis = axes[-1] if axes else None + dp_axes = tuple(axis for axis in axes if axis != fsdp_axis) + resources = { + field.name: getattr(mesh_resource, field.name, None) + for field in fields(MeshResource) + } + resources.update( + ep_resource=( + ep_axis if ep_axis is not None else getattr(mesh_resource, "ep_resource", None) + ), + dp_resource=dp_axes[0] if dp_axes else None, + fsdp_resource=fsdp_axis, + ) + mesh_resource = _LegacyMoEMeshResource( + **resources, + _legacy_data_parallelism_axes=axes, + ) + elif weight_gather.axis is not None and mesh_resource is not None: + mesh_resource = replace(mesh_resource, fsdp_resource=weight_gather.axis) + + if mesh_resource is None: + raise ValueError( + "TE MoE requires mesh_resource=MeshResource(...) or an active global_shard_guard" + " context." + ) + # Snapshot the mutable resource so the custom VJP retains its forward-time axes. + mesh_resource = replace(mesh_resource) + _moe_mesh_axes(mesh_resource) + if quant_before_fsdp_ag and not mesh_resource.fsdp_resource: + raise ValueError("quant_before_fsdp_ag=True requires MeshResource.fsdp_resource.") + return mesh_resource, quant_before_fsdp_ag + + # Per-expert dispatch-slot alignment fed to ``tex.ep_prepare`` as # ``dispatch_output_per_expert_alignment``. NCCL EP HT requires the # per-expert recv block to be at least 128-token aligned, and all current @@ -613,7 +738,8 @@ def _ffn_fwd_per_shard( wi_1_checkpoint_name: Optional[str], wo_checkpoint_name: Optional[str], cudnn_native_weight_layout: bool, - weight_gather: WeightGather, + quant_before_fsdp_ag: bool, + fsdp_axis: Optional[str], fsdp_size: int, ): """Run the grouped FFN on one shard's EP receive buffer.""" @@ -653,10 +779,10 @@ def _ffn_fwd_per_shard( flatten_axis=-1, ) casted_wi = tex.grouped_quantize(wi_for_gemm, fc1_quantizer_set.kernel, flatten_axis=-1) - if weight_gather.mode == "quantized": + if quant_before_fsdp_ag: casted_wi = _gather_quantized_weight( casted_wi, - weight_gather.axis, + fsdp_axis, fsdp_size, 2 if cudnn_native_weight_layout else 1, ) @@ -785,8 +911,8 @@ def _ffn_fwd_per_shard( flatten_axis=-1, ) casted_wo = tex.grouped_quantize(wo, fc2_quantizer_set.kernel, flatten_axis=-1) - if weight_gather.mode == "quantized": - casted_wo = _gather_quantized_weight(casted_wo, weight_gather.axis, fsdp_size, 2) + if quant_before_fsdp_ag: + casted_wo = _gather_quantized_weight(casted_wo, fsdp_axis, fsdp_size, 2) expert_outputs = tex.grouped_gemm( casted_intermediate.get_tensor(usage=TensorUsage.LHS), casted_wo.get_tensor(usage=TensorUsage.RHS), @@ -1026,8 +1152,7 @@ def _moe_fwd_rule( group_topk, scaling_factor, aux_loss_coeff, - ep_axis, - data_parallelism_axes, + mesh_resource, input_axes, gate_kernel_axes, wi_kernel_axes, @@ -1040,137 +1165,117 @@ def _moe_fwd_rule( wi_1_checkpoint_name, wo_checkpoint_name, dispatch_checkpoint_name, - weight_gather, + quant_before_fsdp_ag, ): """Forward: gate -> topk -> ep_dispatch -> FFN -> ep_combine. Returns ``(output, aux_loss)``. ``aux_loss`` is a zero scalar when ``aux_loss_coeff == 0``. """ - del gate_kernel_axes, wi_kernel_axes, wo_kernel_axes # used in bwd only - from jax.experimental.shard_map import shard_map - - x = with_sharding_constraint_by_logical_axes(x, input_axes) - - mesh = _get_mesh() - if mesh is None or mesh.empty: - raise ValueError("moe(...) requires an active jax.sharding.Mesh.") - if ep_axis is None: - raise ValueError("moe(...) requires ep_axis to be set (TE EP backend).") - num_ep = mesh.shape[ep_axis] - if num_experts % num_ep != 0: - raise ValueError(f"num_experts={num_experts} must be divisible by EP size={num_ep}") - num_local_experts = num_experts // num_ep - - dp_size = 1 - for ax in data_parallelism_axes: - dp_size *= mesh.shape[ax] - num_procs = num_ep * dp_size - _validate_moe_quantizer_sets( - quantizer_sets, - num_token_groups=dp_size * num_experts, - num_expert_groups=num_experts, - ) - B, S, H = x.shape - K = num_experts_per_tok - cudnn_native_weight_layout = wi.ndim == 3 and wi.shape[-1] == H - wi_hidden_axis = 2 if cudnn_native_weight_layout else 1 - if weight_gather.mode == "quantized": - if weight_gather.axis not in data_parallelism_axes: - raise ValueError( - "Quantized weight all-gather requires its axis in data_parallelism_axes." - ) - if any(quantizer_set.kernel is None for quantizer_set in quantizer_sets): - raise ValueError("Quantized weight all-gather requires MXFP8 kernel quantizers.") - if wi.shape[wi_hidden_axis] % (mesh.shape[weight_gather.axis] * 32) or wo.shape[2] % ( - mesh.shape[weight_gather.axis] * 32 - ): - raise ValueError("FSDP weight shards must be divisible by the MXFP8 block size 32.") - - if B % num_procs != 0: - raise ValueError(f"batch={B} not divisible by ep*dp={num_procs}") - - # Per-rank send capacity: B/num_procs rows x S tokens per rank. - max_tokens_per_rank = (B // num_procs) * S - dispatch_alignment = _CUDNN_JAX_ALIGN_SIZE if use_cudnn_jax_fusion else _ALIGN_SIZE - worst_case_recv_pr = get_moe_recv_capacity_per_rank( - num_experts=num_experts, - num_experts_per_tok=K, - max_tokens_per_rank=max_tokens_per_rank, - ep_size=num_ep, - alignment=dispatch_alignment, - ) - if recv_capacity_per_rank is None: - recv_pr = worst_case_recv_pr - else: - recv_pr = int(recv_capacity_per_rank) - if recv_pr <= 0 or recv_pr % dispatch_alignment != 0: - raise ValueError( - "recv_capacity_per_rank must be a positive multiple of " - f"{dispatch_alignment}, got" - f" {recv_pr}" - ) + with global_shard_guard(mesh_resource): + ep_axis, data_parallelism_axes = _moe_mesh_axes(mesh_resource) + del gate_kernel_axes, wi_kernel_axes, wo_kernel_axes # used in bwd only + from jax.experimental.shard_map import shard_map + + x = with_sharding_constraint_by_logical_axes(x, input_axes) + + mesh = _get_mesh() + if mesh is None or mesh.empty: + raise ValueError("moe(...) requires an active jax.sharding.Mesh.") + if ep_axis is None: + raise ValueError("moe(...) requires ep_axis to be set (TE EP backend).") + num_ep = mesh.shape[ep_axis] + if num_experts % num_ep != 0: + raise ValueError(f"num_experts={num_experts} must be divisible by EP size={num_ep}") + num_local_experts = num_experts // num_ep + + dp_size = 1 + for ax in data_parallelism_axes: + dp_size *= mesh.shape[ax] + num_procs = num_ep * dp_size + _validate_moe_quantizer_sets( + quantizer_sets, + num_token_groups=dp_size * num_experts, + num_expert_groups=num_experts, + ) + B, S, H = x.shape + K = num_experts_per_tok + cudnn_native_weight_layout = wi.ndim == 3 and wi.shape[-1] == H + wi_hidden_axis = 2 if cudnn_native_weight_layout else 1 + if quant_before_fsdp_ag: + if mesh_resource.fsdp_resource not in data_parallelism_axes: + raise ValueError( + "Quantized weight all-gather requires its FSDP axis among the outer batch axes." + ) + if any(quantizer_set.kernel is None for quantizer_set in quantizer_sets): + raise ValueError("Quantized weight all-gather requires MXFP8 kernel quantizers.") + if wi.shape[wi_hidden_axis] % ( + mesh.shape[mesh_resource.fsdp_resource] * 32 + ) or wo.shape[2] % (mesh.shape[mesh_resource.fsdp_resource] * 32): + raise ValueError("FSDP weight shards must be divisible by the MXFP8 block size 32.") + + if B % num_procs != 0: + raise ValueError(f"batch={B} not divisible by ep*dp={num_procs}") + + # Per-rank send capacity: B/num_procs rows x S tokens per rank. + max_tokens_per_rank = (B // num_procs) * S + dispatch_alignment = _CUDNN_JAX_ALIGN_SIZE if use_cudnn_jax_fusion else _ALIGN_SIZE + worst_case_recv_pr = get_moe_recv_capacity_per_rank( + num_experts=num_experts, + num_experts_per_tok=K, + max_tokens_per_rank=max_tokens_per_rank, + ep_size=num_ep, + alignment=dispatch_alignment, + ) + if recv_capacity_per_rank is None: + recv_pr = worst_case_recv_pr + else: + recv_pr = int(recv_capacity_per_rank) + if recv_pr <= 0 or recv_pr % dispatch_alignment != 0: + raise ValueError( + "recv_capacity_per_rank must be a positive multiple of " + f"{dispatch_alignment}, got" + f" {recv_pr}" + ) - _te_ep_assert_compatible_bootstrap( - num_experts=num_experts, - max_tokens_per_rank=max_tokens_per_rank, - recv_capacity_per_rank=recv_pr, - hidden_dim=H, - ep_size=num_ep, - ) + _te_ep_assert_compatible_bootstrap( + num_experts=num_experts, + max_tokens_per_rank=max_tokens_per_rank, + recv_capacity_per_rank=recv_pr, + hidden_dim=H, + ep_size=num_ep, + ) - if not data_parallelism_axes: - batch_pspec_axis: Any = ep_axis - else: - # ep must be innermost: ep_bootstrap forms NCCL EP comms from - # consecutive global ranks (dp_color = rank // ep_size), so the - # comm only stays within one model replica under (outer_dp, ep). - batch_pspec_axis = (*data_parallelism_axes, ep_axis) - ep3_spec = P(batch_pspec_axis, None, None) - ep2_spec = P(batch_pspec_axis, None) - x = jax.lax.with_sharding_constraint(x, NamedSharding(mesh, ep3_spec)) - - # ---------------- Gate (global view) ---------------- - # tex.fused_topk_with_score_function is only validated against its - # pytorch reference at fp32 (see tests/pytorch/test_fused_router.py: - # parametrize gates dtype on torch.float32 only; the tolerance helper - # raises NotImplementedError for any other dtype). Keeping logits in - # the activation dtype (e.g. bf16) lets sigmoid / softmax / topk - # accumulate at low precision and silently produce NaNs on tokens - # whose normalised weights underflow. Cast to fp32 here to stay in - # the validated regime. - gate_kernel_cast = gate_kernel.astype(x.dtype) - gate_logits = jnp.einsum("bsh,he->bse", x, gate_kernel_cast) - logits_2d = gate_logits.reshape(-1, num_experts).astype(jnp.float32) - - # ---------------- Routing (global view) ---------------- - # expert_bias is an empty (shape-(0,)) sentinel when the caller did - # not enable it; the primitive treats that as "no bias". - eb_arg = expert_bias if expert_bias.shape != (0,) else jnp.zeros((0,), dtype=jnp.float32) - sparse_probs, routing_map, saved_scores = tex.fused_topk_with_score_function_fwd( - logits_2d, - topk=K, - use_pre_softmax=use_pre_softmax, - num_groups=-1 if num_groups is None else num_groups, - group_topk=-1 if group_topk is None else group_topk, - scaling_factor=scaling_factor, - score_function=score_function, - expert_bias=eb_arg, - compute_aux_scores=False, - ) - sparse_probs = sparse_probs.astype(dtype) - - # ---------------- Aux loss (global view, replicated) ---------------- - # ``fused_moe_aux_loss_fwd`` sums probs and tokens_per_expert across - # all tokens, which is wrong when T is sharded. Force-replicate the - # gate logits and recompute the routing map at global view so the - # kernel sees a complete [T_global, E] tensor. The replication is a - # single all-gather over (*dp, ep) and lives off the dispatch - # critical path. - if aux_loss_coeff > 0.0: - global_logits_2d = jax.lax.with_sharding_constraint(logits_2d, NamedSharding(mesh, P())) - _, global_routing_map, _ = tex.fused_topk_with_score_function_fwd( - global_logits_2d, + if not data_parallelism_axes: + batch_pspec_axis: Any = ep_axis + else: + # ep must be innermost: ep_bootstrap forms NCCL EP comms from + # consecutive global ranks (dp_color = rank // ep_size), so the + # comm only stays within one model replica under (outer_dp, ep). + batch_pspec_axis = (*data_parallelism_axes, ep_axis) + ep3_spec = P(batch_pspec_axis, None, None) + ep2_spec = P(batch_pspec_axis, None) + x = jax.lax.with_sharding_constraint(x, NamedSharding(mesh, ep3_spec)) + + # ---------------- Gate (global view) ---------------- + # tex.fused_topk_with_score_function is only validated against its + # pytorch reference at fp32 (see tests/pytorch/test_fused_router.py: + # parametrize gates dtype on torch.float32 only; the tolerance helper + # raises NotImplementedError for any other dtype). Keeping logits in + # the activation dtype (e.g. bf16) lets sigmoid / softmax / topk + # accumulate at low precision and silently produce NaNs on tokens + # whose normalised weights underflow. Cast to fp32 here to stay in + # the validated regime. + gate_kernel_cast = gate_kernel.astype(x.dtype) + gate_logits = jnp.einsum("bsh,he->bse", x, gate_kernel_cast) + logits_2d = gate_logits.reshape(-1, num_experts).astype(jnp.float32) + + # ---------------- Routing (global view) ---------------- + # expert_bias is an empty (shape-(0,)) sentinel when the caller did + # not enable it; the primitive treats that as "no bias". + eb_arg = expert_bias if expert_bias.shape != (0,) else jnp.zeros((0,), dtype=jnp.float32) + sparse_probs, routing_map, saved_scores = tex.fused_topk_with_score_function_fwd( + logits_2d, topk=K, use_pre_softmax=use_pre_softmax, num_groups=-1 if num_groups is None else num_groups, @@ -1180,210 +1285,239 @@ def _moe_fwd_rule( expert_bias=eb_arg, compute_aux_scores=False, ) - aux_tokens_per_expert = jnp.sum(global_routing_map.astype(jnp.int32), axis=0) - # compute_aux_scores=True takes a separate kernel path: clean - # per-expert softmax, no grouping / bias / scaling. - aux_probs, _aux_rm, aux_saved_scores = tex.fused_topk_with_score_function_fwd( - global_logits_2d.astype(jnp.float32), - topk=K, - use_pre_softmax=False, - num_groups=-1, - group_topk=-1, - scaling_factor=1.0, - score_function=score_function, - expert_bias=jnp.zeros((0,), dtype=jnp.float32), - compute_aux_scores=True, + sparse_probs = sparse_probs.astype(dtype) + + # ---------------- Aux loss (global view, replicated) ---------------- + # ``fused_moe_aux_loss_fwd`` sums probs and tokens_per_expert across + # all tokens, which is wrong when T is sharded. Force-replicate the + # gate logits and recompute the routing map at global view so the + # kernel sees a complete [T_global, E] tensor. The replication is a + # single all-gather over (*dp, ep) and lives off the dispatch + # critical path. + if aux_loss_coeff > 0.0: + global_logits_2d = jax.lax.with_sharding_constraint(logits_2d, NamedSharding(mesh, P())) + _, global_routing_map, _ = tex.fused_topk_with_score_function_fwd( + global_logits_2d, + topk=K, + use_pre_softmax=use_pre_softmax, + num_groups=-1 if num_groups is None else num_groups, + group_topk=-1 if group_topk is None else group_topk, + scaling_factor=scaling_factor, + score_function=score_function, + expert_bias=eb_arg, + compute_aux_scores=False, + ) + aux_tokens_per_expert = jnp.sum(global_routing_map.astype(jnp.int32), axis=0) + # compute_aux_scores=True takes a separate kernel path: clean + # per-expert softmax, no grouping / bias / scaling. + aux_probs, _aux_rm, aux_saved_scores = tex.fused_topk_with_score_function_fwd( + global_logits_2d.astype(jnp.float32), + topk=K, + use_pre_softmax=False, + num_groups=-1, + group_topk=-1, + scaling_factor=1.0, + score_function=score_function, + expert_bias=jnp.zeros((0,), dtype=jnp.float32), + compute_aux_scores=True, + ) + aux_loss, aux_const_buf = tex.fused_moe_aux_loss_fwd( + aux_probs.astype(jnp.float32), + aux_tokens_per_expert.astype(jnp.int32), + topk=K, + coeff=aux_loss_coeff, + ) + aux_loss = aux_loss.astype(dtype) + else: + aux_loss = jnp.zeros((), dtype=dtype) + aux_const_buf = None + aux_tokens_per_expert = None + aux_saved_scores = None + + # ---------------- Routing -> (topk_idx, topk_w) at 3D ---------------- + # argsort on a bool tensor places True last (False=0 < True=1), so the + # last K indices are the selected expert IDs. + selected_experts = jnp.argsort(routing_map, axis=-1)[..., -K:] + routing_weights = jnp.take_along_axis(sparse_probs, selected_experts, axis=-1) + topk_idx_3d = selected_experts.reshape(B, S, K).astype(jnp.int32) + topk_w_3d = routing_weights.reshape(B, S, K).astype(jnp.float32) + # tex.ep_prepare/dispatch's partition only folds ep_axis into a replicated + # leading dim, not the outer dp/fsdp axes, so a replicated topk_idx makes + # each rank see B/ep rows (not B/num_procs) and overrun the bootstrap-sized + # send buffer. Pin both routing tensors to the (outer, ep) leading sharding + # so per-rank token counts match max_tokens_per_rank. + topk_idx_3d = jax.lax.with_sharding_constraint(topk_idx_3d, NamedSharding(mesh, ep3_spec)) + topk_w_3d = jax.lax.with_sharding_constraint(topk_w_3d, NamedSharding(mesh, ep3_spec)) + + # ---------------- TE EP dispatch (global view) ---------------- + cfg = tex.EpLayerConfig( + top_k=K, + dispatch_output_per_expert_alignment=dispatch_alignment, ) - aux_loss, aux_const_buf = tex.fused_moe_aux_loss_fwd( - aux_probs.astype(jnp.float32), - aux_tokens_per_expert.astype(jnp.int32), - topk=K, - coeff=aux_loss_coeff, + token_counts, total_recv_tokens, handle_mem = tex.ep_prepare(cfg, topk_idx_3d) + token_counts = jax.lax.with_sharding_constraint(token_counts, NamedSharding(mesh, ep2_spec)) + recv_tokens, recv_topk_weights = tex.ep_dispatch_fwd( + cfg, handle_mem, topk_idx_3d, x, topk_w_3d, recv_pr ) - aux_loss = aux_loss.astype(dtype) - else: - aux_loss = jnp.zeros((), dtype=dtype) - aux_const_buf = None - aux_tokens_per_expert = None - aux_saved_scores = None - - # ---------------- Routing -> (topk_idx, topk_w) at 3D ---------------- - # argsort on a bool tensor places True last (False=0 < True=1), so the - # last K indices are the selected expert IDs. - selected_experts = jnp.argsort(routing_map, axis=-1)[..., -K:] - routing_weights = jnp.take_along_axis(sparse_probs, selected_experts, axis=-1) - topk_idx_3d = selected_experts.reshape(B, S, K).astype(jnp.int32) - topk_w_3d = routing_weights.reshape(B, S, K).astype(jnp.float32) - # tex.ep_prepare/dispatch's partition only folds ep_axis into a replicated - # leading dim, not the outer dp/fsdp axes, so a replicated topk_idx makes - # each rank see B/ep rows (not B/num_procs) and overrun the bootstrap-sized - # send buffer. Pin both routing tensors to the (outer, ep) leading sharding - # so per-rank token counts match max_tokens_per_rank. - topk_idx_3d = jax.lax.with_sharding_constraint(topk_idx_3d, NamedSharding(mesh, ep3_spec)) - topk_w_3d = jax.lax.with_sharding_constraint(topk_w_3d, NamedSharding(mesh, ep3_spec)) - - # ---------------- TE EP dispatch (global view) ---------------- - cfg = tex.EpLayerConfig( - top_k=K, - dispatch_output_per_expert_alignment=dispatch_alignment, - ) - token_counts, total_recv_tokens, handle_mem = tex.ep_prepare(cfg, topk_idx_3d) - token_counts = jax.lax.with_sharding_constraint(token_counts, NamedSharding(mesh, ep2_spec)) - recv_tokens, recv_topk_weights = tex.ep_dispatch_fwd( - cfg, handle_mem, topk_idx_3d, x, topk_w_3d, recv_pr - ) - if dispatch_checkpoint_name is not None: - recv_tokens = checkpoint_name(recv_tokens, dispatch_checkpoint_name) - recv_topk_weights = checkpoint_name(recv_topk_weights, dispatch_checkpoint_name) - recv_tokens = jax.lax.with_sharding_constraint(recv_tokens, NamedSharding(mesh, ep3_spec)) - recv_topk_weights = jax.lax.with_sharding_constraint( - recv_topk_weights, NamedSharding(mesh, ep2_spec) - ) - - # ---------------- FFN (per-shard via shard_map) ---------------- - has_bias = wi_0_bias is not None - kernel_spec = P(ep_axis, None, None) - wi_input_spec = ( - ( - P(ep_axis, None, weight_gather.axis) - if cudnn_native_weight_layout - else P(ep_axis, weight_gather.axis, None) + if dispatch_checkpoint_name is not None: + recv_tokens = checkpoint_name(recv_tokens, dispatch_checkpoint_name) + recv_topk_weights = checkpoint_name(recv_topk_weights, dispatch_checkpoint_name) + recv_tokens = jax.lax.with_sharding_constraint(recv_tokens, NamedSharding(mesh, ep3_spec)) + recv_topk_weights = jax.lax.with_sharding_constraint( + recv_topk_weights, NamedSharding(mesh, ep2_spec) ) - if weight_gather.mode == "quantized" - else kernel_spec - ) - wo_input_spec = ( - P(ep_axis, None, weight_gather.axis) if weight_gather.mode == "quantized" else kernel_spec - ) - bias_spec = P(ep_axis, None) - ffn_in_specs = (ep3_spec, ep2_spec, ep2_spec, wi_input_spec, wo_input_spec) - ffn_in_args = [recv_tokens, recv_topk_weights, token_counts, wi, wo] - if has_bias: - ffn_in_specs += (bias_spec, bias_spec, bias_spec) - ffn_in_args.extend([wi_0_bias, wi_1_bias, wo_bias]) - - # Quantized grouped tensors store their data, scales, and group metadata - # as physical buffers rather than in the source tensor's logical shape. - # A PartitionSpec used as a pytree prefix applies the same ownership to - # every array leaf of the grouped tensor: dispatched-token buffers belong - # to the compound batch shard, while expert-weight buffers belong to EP. - token_buffer_spec = P(batch_pspec_axis) - token_matrix_spec = P(batch_pspec_axis, None) - expert_buffer_spec = P(ep_axis) - residuals_spec = ( - token_buffer_spec, - expert_buffer_spec, - token_matrix_spec, - token_matrix_spec, - token_buffer_spec, - expert_buffer_spec, - ep2_spec, - ) - def _ffn_fwd_body(*args): + # ---------------- FFN (per-shard via shard_map) ---------------- + has_bias = wi_0_bias is not None + kernel_spec = P(ep_axis, None, None) + wi_input_spec = ( + ( + P(ep_axis, None, mesh_resource.fsdp_resource) + if cudnn_native_weight_layout + else P(ep_axis, mesh_resource.fsdp_resource, None) + ) + if quant_before_fsdp_ag + else kernel_spec + ) + wo_input_spec = ( + P(ep_axis, None, mesh_resource.fsdp_resource) if quant_before_fsdp_ag else kernel_spec + ) + bias_spec = P(ep_axis, None) + ffn_in_specs = (ep3_spec, ep2_spec, ep2_spec, wi_input_spec, wo_input_spec) + ffn_in_args = [recv_tokens, recv_topk_weights, token_counts, wi, wo] if has_bias: - r_tok, r_w, tc, local_wi, local_wo, w0b, w1b, wob = args - else: - r_tok, r_w, tc, local_wi, local_wo = args - w0b = w1b = wob = None - return _ffn_fwd_per_shard( - r_tok, - r_w, - tc, - local_wi, - local_wo, - w0b, - w1b, - wob, - quantizer_sets, - num_local_experts=num_local_experts, - activation_type=activation_type, - apply_topk_weights_early=apply_topk_weights_early, - use_cudnn_jax_fusion=use_cudnn_jax_fusion, - wi_0_checkpoint_name=wi_0_checkpoint_name, - wi_1_checkpoint_name=wi_1_checkpoint_name, - wo_checkpoint_name=wo_checkpoint_name, - cudnn_native_weight_layout=cudnn_native_weight_layout, - weight_gather=weight_gather, - fsdp_size=mesh.shape[weight_gather.axis] if weight_gather.axis is not None else 1, + ffn_in_specs += (bias_spec, bias_spec, bias_spec) + ffn_in_args.extend([wi_0_bias, wi_1_bias, wo_bias]) + + # Quantized grouped tensors store their data, scales, and group metadata + # as physical buffers rather than in the source tensor's logical shape. + # A PartitionSpec used as a pytree prefix applies the same ownership to + # every array leaf of the grouped tensor: dispatched-token buffers belong + # to the compound batch shard, while expert-weight buffers belong to EP. + token_buffer_spec = P(batch_pspec_axis) + token_matrix_spec = P(batch_pspec_axis, None) + expert_buffer_spec = P(ep_axis) + residuals_spec = ( + token_buffer_spec, + expert_buffer_spec, + token_matrix_spec, + token_matrix_spec, + token_buffer_spec, + expert_buffer_spec, + ep2_spec, ) - expert_outputs, ffn_residuals = shard_map( - _ffn_fwd_body, - mesh=mesh, - in_specs=ffn_in_specs, - out_specs=(ep3_spec, residuals_spec), - check_rep=False, - )(*ffn_in_args) - expert_outputs = jax.lax.with_sharding_constraint(expert_outputs, NamedSharding(mesh, ep3_spec)) - - # ---------------- TE EP combine (global view) ---------------- - out_partition_spec = (batch_pspec_axis, None, None) - if apply_topk_weights_early: - # expert_outputs is already weighted upstream. - output = tex.ep_combine_fwd( - cfg, - handle_mem, - expert_outputs, - num_local_tokens=(B, S), - out_partition_spec=out_partition_spec, - ) - else: - # HT combine is unweighted; apply routing weights before calling it. - # Padded recv slots are ignored by combine via handle_mem metadata. - w = recv_topk_weights[..., None].astype(expert_outputs.dtype) - weighted = expert_outputs * w - output = tex.ep_combine_fwd( - cfg, - handle_mem, - weighted, - num_local_tokens=(B, S), - out_partition_spec=out_partition_spec, + def _ffn_fwd_body(*args): + if has_bias: + r_tok, r_w, tc, local_wi, local_wo, w0b, w1b, wob = args + else: + r_tok, r_w, tc, local_wi, local_wo = args + w0b = w1b = wob = None + return _ffn_fwd_per_shard( + r_tok, + r_w, + tc, + local_wi, + local_wo, + w0b, + w1b, + wob, + quantizer_sets, + num_local_experts=num_local_experts, + activation_type=activation_type, + apply_topk_weights_early=apply_topk_weights_early, + use_cudnn_jax_fusion=use_cudnn_jax_fusion, + wi_0_checkpoint_name=wi_0_checkpoint_name, + wi_1_checkpoint_name=wi_1_checkpoint_name, + wo_checkpoint_name=wo_checkpoint_name, + cudnn_native_weight_layout=cudnn_native_weight_layout, + quant_before_fsdp_ag=quant_before_fsdp_ag, + fsdp_axis=mesh_resource.fsdp_resource, + fsdp_size=( + mesh.shape[mesh_resource.fsdp_resource] + if mesh_resource.fsdp_resource is not None + else 1 + ), + ) + + expert_outputs, ffn_residuals = shard_map( + _ffn_fwd_body, + mesh=mesh, + in_specs=ffn_in_specs, + out_specs=(ep3_spec, residuals_spec), + check_rep=False, + )(*ffn_in_args) + expert_outputs = jax.lax.with_sharding_constraint( + expert_outputs, NamedSharding(mesh, ep3_spec) ) - # output of MLP should be sharded the same way as the activation input - output = with_sharding_constraint_by_logical_axes(output, input_axes) - ( - casted_sorted_x_lhs_trans, - casted_wi_rhs_trans, - gate_proj_out, - up_proj_out, - casted_intermediate_lhs_trans, - casted_wo_rhs_trans, - local_group_sizes, - ) = ffn_residuals - - ctx = _Ctx( - x=x, - gate_kernel=gate_kernel, - expert_bias=expert_bias, - logits_2d=logits_2d, - saved_scores=saved_scores, - routing_map=routing_map, - cfg=cfg, - handle_mem=handle_mem, - recv_topk_weights=recv_topk_weights, - casted_sorted_x_lhs_trans=casted_sorted_x_lhs_trans, - casted_wi_rhs_trans=casted_wi_rhs_trans, - gate_proj_out=gate_proj_out, - up_proj_out=up_proj_out, - casted_intermediate_lhs_trans=casted_intermediate_lhs_trans, - casted_wo_rhs_trans=casted_wo_rhs_trans, - expert_outputs=expert_outputs, - local_group_sizes=local_group_sizes, - quantizer_sets=quantizer_sets, - aux_const_buf=aux_const_buf, - aux_tokens_per_expert=aux_tokens_per_expert, - aux_saved_scores=aux_saved_scores, - ) - static = { - "has_bias": has_bias, - "x_shape": x.shape, - "recv_pr": recv_pr, - "cudnn_native_weight_layout": cudnn_native_weight_layout, - } - # total_recv_tokens is a non-differentiable overflow signal (see moe()). - return (output, aux_loss, total_recv_tokens), (ctx, static) + # ---------------- TE EP combine (global view) ---------------- + out_partition_spec = (batch_pspec_axis, None, None) + if apply_topk_weights_early: + # expert_outputs is already weighted upstream. + output = tex.ep_combine_fwd( + cfg, + handle_mem, + expert_outputs, + num_local_tokens=(B, S), + out_partition_spec=out_partition_spec, + ) + else: + # HT combine is unweighted; apply routing weights before calling it. + # Padded recv slots are ignored by combine via handle_mem metadata. + w = recv_topk_weights[..., None].astype(expert_outputs.dtype) + weighted = expert_outputs * w + output = tex.ep_combine_fwd( + cfg, + handle_mem, + weighted, + num_local_tokens=(B, S), + out_partition_spec=out_partition_spec, + ) + # output of MLP should be sharded the same way as the activation input + output = with_sharding_constraint_by_logical_axes(output, input_axes) + + ( + casted_sorted_x_lhs_trans, + casted_wi_rhs_trans, + gate_proj_out, + up_proj_out, + casted_intermediate_lhs_trans, + casted_wo_rhs_trans, + local_group_sizes, + ) = ffn_residuals + + ctx = _Ctx( + x=x, + gate_kernel=gate_kernel, + expert_bias=expert_bias, + logits_2d=logits_2d, + saved_scores=saved_scores, + routing_map=routing_map, + cfg=cfg, + handle_mem=handle_mem, + recv_topk_weights=recv_topk_weights, + casted_sorted_x_lhs_trans=casted_sorted_x_lhs_trans, + casted_wi_rhs_trans=casted_wi_rhs_trans, + gate_proj_out=gate_proj_out, + up_proj_out=up_proj_out, + casted_intermediate_lhs_trans=casted_intermediate_lhs_trans, + casted_wo_rhs_trans=casted_wo_rhs_trans, + expert_outputs=expert_outputs, + local_group_sizes=local_group_sizes, + quantizer_sets=quantizer_sets, + aux_const_buf=aux_const_buf, + aux_tokens_per_expert=aux_tokens_per_expert, + aux_saved_scores=aux_saved_scores, + ) + static = { + "has_bias": has_bias, + "x_shape": x.shape, + "recv_pr": recv_pr, + "cudnn_native_weight_layout": cudnn_native_weight_layout, + } + # total_recv_tokens is a non-differentiable overflow signal (see moe()). + return (output, aux_loss, total_recv_tokens), (ctx, static) def _moe_bwd_rule( @@ -1396,8 +1530,7 @@ def _moe_bwd_rule( group_topk, scaling_factor, aux_loss_coeff, - ep_axis, - data_parallelism_axes, + mesh_resource, input_axes, gate_kernel_axes, wi_kernel_axes, @@ -1410,259 +1543,265 @@ def _moe_bwd_rule( wi_1_checkpoint_name, wo_checkpoint_name, dispatch_checkpoint_name, - weight_gather, + quant_before_fsdp_ag, residuals, cotangents, ): """Backward mirror of :func:`_moe_fwd_rule`.""" - del ( - num_groups, - group_topk, - dtype, - recv_capacity_per_rank, - wi_0_checkpoint_name, - wi_1_checkpoint_name, - wo_checkpoint_name, - dispatch_checkpoint_name, - weight_gather, - ) # captured / unused in bwd - from jax.experimental.shard_map import shard_map - - # total_recv_tokens is a non-differentiable output; its cotangent is unused. - d_output, d_aux_loss, _d_total_recv_tokens = cotangents - - ctx, static = residuals - has_bias = static["has_bias"] - x_shape = static["x_shape"] - recv_pr = static["recv_pr"] - - mesh = _get_mesh() - if mesh is None or mesh.empty: - raise ValueError("moe(...) requires an active jax.sharding.Mesh.") - B, S, _ = x_shape - K = num_experts_per_tok - if not data_parallelism_axes: - batch_pspec_axis: Any = ep_axis - else: - batch_pspec_axis = (*data_parallelism_axes, ep_axis) - ep3_spec = P(batch_pspec_axis, None, None) - ep2_spec = P(batch_pspec_axis, None) - out_partition_spec = (batch_pspec_axis, None, None) - - # ---------------- Combine bwd (global view) ---------------- - d_output = jax.lax.with_sharding_constraint(d_output, NamedSharding(mesh, ep3_spec)) - grad_pre_combine = tex.ep_combine_bwd(ctx.cfg, ctx.handle_mem, d_output, recv_pr) - grad_pre_combine = jax.lax.with_sharding_constraint( - grad_pre_combine, NamedSharding(mesh, ep3_spec) - ) - if apply_topk_weights_early: - # combine_fwd consumed already-weighted expert_outputs; the recv_w - # cotangent flows through the early-weighting step inside the FFN bwd. - d_expert_outputs = grad_pre_combine - d_recv_w_from_combine = jnp.zeros_like(ctx.recv_topk_weights) - else: - w = ctx.recv_topk_weights[..., None].astype(grad_pre_combine.dtype) - d_expert_outputs = grad_pre_combine * w - d_recv_w_from_combine = (grad_pre_combine * ctx.expert_outputs).sum(axis=-1) - d_recv_w_from_combine = d_recv_w_from_combine.astype(ctx.recv_topk_weights.dtype) - - # ---------------- FFN bwd (per-shard via shard_map) ---------------- - kernel_spec = P(ep_axis, None, None) - bias_spec = P(ep_axis, None) - token_buffer_spec = P(batch_pspec_axis) - token_matrix_spec = P(batch_pspec_axis, None) - expert_buffer_spec = P(ep_axis) - residuals_specs = ( - token_buffer_spec, - expert_buffer_spec, - token_matrix_spec, - token_matrix_spec, - token_buffer_spec, - expert_buffer_spec, - ep2_spec, - ) - bwd_in_specs = (ep3_spec, *residuals_specs, ep2_spec) - bwd_in_args = [ - d_expert_outputs, - ctx.casted_sorted_x_lhs_trans, - ctx.casted_wi_rhs_trans, - ctx.gate_proj_out, - ctx.up_proj_out, - ctx.casted_intermediate_lhs_trans, - ctx.casted_wo_rhs_trans, - ctx.local_group_sizes, - ctx.recv_topk_weights, - ] - - def _ffn_bwd_body(*args): - grads = _ffn_bwd_per_shard( - *args, - ctx.quantizer_sets, - activation_type=activation_type, - apply_topk_weights_early=apply_topk_weights_early, - has_bias=has_bias, - use_cudnn_jax_fusion=use_cudnn_jax_fusion, - cudnn_native_weight_layout=static["cudnn_native_weight_layout"], - ) - ( - d_sorted_x_local, - d_recv_w_local, - d_wi_local, - d_wo_local, - d_wi_0_bias_local, - d_wi_1_bias_local, - d_wo_bias_local, - ) = grads - if data_parallelism_axes: - dp_axes = tuple(data_parallelism_axes) - d_wi_local = jax.lax.psum(d_wi_local, axis_name=dp_axes) - d_wo_local = jax.lax.psum(d_wo_local, axis_name=dp_axes) - if has_bias: - d_wi_0_bias_local = jax.lax.psum(d_wi_0_bias_local, axis_name=dp_axes) - d_wi_1_bias_local = jax.lax.psum(d_wi_1_bias_local, axis_name=dp_axes) - d_wo_bias_local = jax.lax.psum(d_wo_bias_local, axis_name=dp_axes) - return ( - d_sorted_x_local, - d_recv_w_local, - d_wi_local, - d_wo_local, - d_wi_0_bias_local, - d_wi_1_bias_local, - d_wo_bias_local, + with global_shard_guard(mesh_resource): + ep_axis, data_parallelism_axes = _moe_mesh_axes(mesh_resource) + del ( + num_groups, + group_topk, + dtype, + recv_capacity_per_rank, + wi_0_checkpoint_name, + wi_1_checkpoint_name, + wo_checkpoint_name, + dispatch_checkpoint_name, + quant_before_fsdp_ag, + ) # captured / unused in bwd + from jax.experimental.shard_map import shard_map + + # total_recv_tokens is a non-differentiable output; its cotangent is unused. + d_output, d_aux_loss, _d_total_recv_tokens = cotangents + + ctx, static = residuals + has_bias = static["has_bias"] + x_shape = static["x_shape"] + recv_pr = static["recv_pr"] + + mesh = _get_mesh() + if mesh is None or mesh.empty: + raise ValueError("moe(...) requires an active jax.sharding.Mesh.") + B, S, _ = x_shape + K = num_experts_per_tok + if not data_parallelism_axes: + batch_pspec_axis: Any = ep_axis + else: + batch_pspec_axis = (*data_parallelism_axes, ep_axis) + ep3_spec = P(batch_pspec_axis, None, None) + ep2_spec = P(batch_pspec_axis, None) + out_partition_spec = (batch_pspec_axis, None, None) + + # ---------------- Combine bwd (global view) ---------------- + d_output = jax.lax.with_sharding_constraint(d_output, NamedSharding(mesh, ep3_spec)) + grad_pre_combine = tex.ep_combine_bwd(ctx.cfg, ctx.handle_mem, d_output, recv_pr) + grad_pre_combine = jax.lax.with_sharding_constraint( + grad_pre_combine, NamedSharding(mesh, ep3_spec) ) - - if has_bias: - bwd_out_specs = ( - ep3_spec, + if apply_topk_weights_early: + # combine_fwd consumed already-weighted expert_outputs; the recv_w + # cotangent flows through the early-weighting step inside the FFN bwd. + d_expert_outputs = grad_pre_combine + d_recv_w_from_combine = jnp.zeros_like(ctx.recv_topk_weights) + else: + w = ctx.recv_topk_weights[..., None].astype(grad_pre_combine.dtype) + d_expert_outputs = grad_pre_combine * w + d_recv_w_from_combine = (grad_pre_combine * ctx.expert_outputs).sum(axis=-1) + d_recv_w_from_combine = d_recv_w_from_combine.astype(ctx.recv_topk_weights.dtype) + + # ---------------- FFN bwd (per-shard via shard_map) ---------------- + kernel_spec = P(ep_axis, None, None) + bias_spec = P(ep_axis, None) + token_buffer_spec = P(batch_pspec_axis) + token_matrix_spec = P(batch_pspec_axis, None) + expert_buffer_spec = P(ep_axis) + residuals_specs = ( + token_buffer_spec, + expert_buffer_spec, + token_matrix_spec, + token_matrix_spec, + token_buffer_spec, + expert_buffer_spec, ep2_spec, - kernel_spec, - kernel_spec, - bias_spec, - bias_spec, - bias_spec, ) - else: - bwd_out_specs = (ep3_spec, ep2_spec, kernel_spec, kernel_spec, None, None, None) + bwd_in_specs = (ep3_spec, *residuals_specs, ep2_spec) + bwd_in_args = [ + d_expert_outputs, + ctx.casted_sorted_x_lhs_trans, + ctx.casted_wi_rhs_trans, + ctx.gate_proj_out, + ctx.up_proj_out, + ctx.casted_intermediate_lhs_trans, + ctx.casted_wo_rhs_trans, + ctx.local_group_sizes, + ctx.recv_topk_weights, + ] + + def _ffn_bwd_body(*args): + grads = _ffn_bwd_per_shard( + *args, + ctx.quantizer_sets, + activation_type=activation_type, + apply_topk_weights_early=apply_topk_weights_early, + has_bias=has_bias, + use_cudnn_jax_fusion=use_cudnn_jax_fusion, + cudnn_native_weight_layout=static["cudnn_native_weight_layout"], + ) + ( + d_sorted_x_local, + d_recv_w_local, + d_wi_local, + d_wo_local, + d_wi_0_bias_local, + d_wi_1_bias_local, + d_wo_bias_local, + ) = grads + if data_parallelism_axes: + dp_axes = tuple(data_parallelism_axes) + d_wi_local = jax.lax.psum(d_wi_local, axis_name=dp_axes) + d_wo_local = jax.lax.psum(d_wo_local, axis_name=dp_axes) + if has_bias: + d_wi_0_bias_local = jax.lax.psum(d_wi_0_bias_local, axis_name=dp_axes) + d_wi_1_bias_local = jax.lax.psum(d_wi_1_bias_local, axis_name=dp_axes) + d_wo_bias_local = jax.lax.psum(d_wo_bias_local, axis_name=dp_axes) + return ( + d_sorted_x_local, + d_recv_w_local, + d_wi_local, + d_wo_local, + d_wi_0_bias_local, + d_wi_1_bias_local, + d_wo_bias_local, + ) - ( - d_sorted_x, - d_recv_w_from_intermediate, - d_wi, - d_wo, - d_wi_0_bias, - d_wi_1_bias, - d_wo_bias, - ) = shard_map( - _ffn_bwd_body, - mesh=mesh, - in_specs=bwd_in_specs, - out_specs=bwd_out_specs, - check_rep=False, - )( - *bwd_in_args - ) + if has_bias: + bwd_out_specs = ( + ep3_spec, + ep2_spec, + kernel_spec, + kernel_spec, + bias_spec, + bias_spec, + bias_spec, + ) + else: + bwd_out_specs = (ep3_spec, ep2_spec, kernel_spec, kernel_spec, None, None, None) - d_recv_w_total = d_recv_w_from_combine + d_recv_w_from_intermediate - - # ---------------- Dispatch bwd (global view) ---------------- - d_sorted_x = jax.lax.with_sharding_constraint(d_sorted_x, NamedSharding(mesh, ep3_spec)) - d_recv_w_total = jax.lax.with_sharding_constraint(d_recv_w_total, NamedSharding(mesh, ep2_spec)) - d_x_from_dispatch, d_topk_w = tex.ep_dispatch_bwd( - ctx.cfg, - ctx.handle_mem, - d_sorted_x, - d_recv_w_total, - num_local_tokens=(B, S), - out_partition_spec=out_partition_spec, - ) + ( + d_sorted_x, + d_recv_w_from_intermediate, + d_wi, + d_wo, + d_wi_0_bias, + d_wi_1_bias, + d_wo_bias, + ) = shard_map( + _ffn_bwd_body, + mesh=mesh, + in_specs=bwd_in_specs, + out_specs=bwd_out_specs, + check_rep=False, + )( + *bwd_in_args + ) - # ---------------- Routing bwd (global view) ---------------- - # The cotangent on routing_weights is a sparse scatter into sparse_probs - # at the selected_experts indices. - selected_experts = jnp.argsort(ctx.routing_map, axis=-1)[..., -K:] - d_topk_w_flat = d_topk_w.reshape(-1, K) - d_sparse_probs = jnp.zeros(ctx.routing_map.shape, dtype=d_topk_w_flat.dtype) - d_sparse_probs = d_sparse_probs.at[ - jnp.arange(ctx.routing_map.shape[0])[:, None], selected_experts - ].set(d_topk_w_flat) - - d_logits_2d = tex.fused_topk_with_score_function_bwd( - ctx.routing_map, - ctx.saved_scores, - d_sparse_probs.astype(ctx.saved_scores.dtype), - topk=K, - use_pre_softmax=use_pre_softmax, - scaling_factor=scaling_factor, - score_function=score_function, - compute_aux_scores=False, - ) + d_recv_w_total = d_recv_w_from_combine + d_recv_w_from_intermediate - # ---------------- Aux loss bwd (global view, replicated) ---------------- - # Reverse the fwd's all-gather/aux pipeline: aux_loss_bwd produces - # d_aux_probs, then topk_bwd(compute_aux_scores=True) produces the - # extra d_logits contribution. The replicated tensor adds into the - # T-sharded routing-side d_logits via JAX's normal broadcast. - if aux_loss_coeff > 0.0: - T_global = ctx.logits_2d.shape[0] - d_aux_loss_scalar = d_aux_loss.reshape(()).astype(jnp.float32) - d_aux_probs = tex.fused_moe_aux_loss_bwd( - ctx.aux_const_buf, - ctx.aux_tokens_per_expert.astype(jnp.int32), - d_aux_loss_scalar, - num_tokens=int(T_global), + # ---------------- Dispatch bwd (global view) ---------------- + d_sorted_x = jax.lax.with_sharding_constraint(d_sorted_x, NamedSharding(mesh, ep3_spec)) + d_recv_w_total = jax.lax.with_sharding_constraint( + d_recv_w_total, NamedSharding(mesh, ep2_spec) + ) + d_x_from_dispatch, d_topk_w = tex.ep_dispatch_bwd( + ctx.cfg, + ctx.handle_mem, + d_sorted_x, + d_recv_w_total, + num_local_tokens=(B, S), + out_partition_spec=out_partition_spec, ) - # routing_map is ignored by the kernel when compute_aux_scores=True, - # so pass a zero placeholder of the right shape/dtype. - zero_routing_map = jnp.zeros(ctx.aux_saved_scores.shape, dtype=ctx.routing_map.dtype) - d_logits_aux = tex.fused_topk_with_score_function_bwd( - zero_routing_map, - ctx.aux_saved_scores, - d_aux_probs.astype(ctx.aux_saved_scores.dtype), + + # ---------------- Routing bwd (global view) ---------------- + # The cotangent on routing_weights is a sparse scatter into sparse_probs + # at the selected_experts indices. + selected_experts = jnp.argsort(ctx.routing_map, axis=-1)[..., -K:] + d_topk_w_flat = d_topk_w.reshape(-1, K) + d_sparse_probs = jnp.zeros(ctx.routing_map.shape, dtype=d_topk_w_flat.dtype) + d_sparse_probs = d_sparse_probs.at[ + jnp.arange(ctx.routing_map.shape[0])[:, None], selected_experts + ].set(d_topk_w_flat) + + d_logits_2d = tex.fused_topk_with_score_function_bwd( + ctx.routing_map, + ctx.saved_scores, + d_sparse_probs.astype(ctx.saved_scores.dtype), topk=K, - use_pre_softmax=False, - scaling_factor=1.0, + use_pre_softmax=use_pre_softmax, + scaling_factor=scaling_factor, score_function=score_function, - compute_aux_scores=True, + compute_aux_scores=False, ) - d_logits_2d = d_logits_2d + d_logits_aux.astype(d_logits_2d.dtype) - - # ---------------- Gate bwd (global view) ---------------- - d_gate_logits = d_logits_2d.reshape(B, S, num_experts) - gate_kernel_cast = ctx.gate_kernel.astype(ctx.x.dtype) - d_x_from_gate = jnp.einsum("bse,he->bsh", d_gate_logits, gate_kernel_cast) - d_gate_kernel = jnp.einsum("bsh,bse->he", ctx.x, d_gate_logits).astype(ctx.gate_kernel.dtype) - d_x = d_x_from_gate + d_x_from_dispatch - - # Pin output grads to the declared logical axes so downstream - # optimizers see consistent shardings. - d_x = with_sharding_constraint_by_logical_axes(d_x, input_axes) - d_gate_kernel = with_sharding_constraint_by_logical_axes(d_gate_kernel, gate_kernel_axes) - d_wi = with_sharding_constraint_by_logical_axes(d_wi, wi_kernel_axes) - d_wo = with_sharding_constraint_by_logical_axes(d_wo, wo_kernel_axes) - if has_bias: - wi_bias_axes = (wi_kernel_axes[0], *wi_kernel_axes[2:]) - wo_bias_axes = (wo_kernel_axes[0], *wo_kernel_axes[2:]) - d_wi_0_bias = with_sharding_constraint_by_logical_axes(d_wi_0_bias, wi_bias_axes) - d_wi_1_bias = with_sharding_constraint_by_logical_axes(d_wi_1_bias, wi_bias_axes) - d_wo_bias = with_sharding_constraint_by_logical_axes(d_wo_bias, wo_bias_axes) - - # expert_bias has no learnable bwd path through fused_topk: the - # primitive's bwd returns None for the bias slot. Match that with a - # zero cotangent of the right shape so custom_vjp's arity check - # passes. - d_expert_bias = jnp.zeros_like(ctx.expert_bias) - return ( - d_x, - d_gate_kernel, - d_wi, - d_wo, - d_wi_0_bias if has_bias else None, - d_wi_1_bias if has_bias else None, - d_wo_bias if has_bias else None, - d_expert_bias, - ctx.quantizer_sets, - ) + # ---------------- Aux loss bwd (global view, replicated) ---------------- + # Reverse the fwd's all-gather/aux pipeline: aux_loss_bwd produces + # d_aux_probs, then topk_bwd(compute_aux_scores=True) produces the + # extra d_logits contribution. The replicated tensor adds into the + # T-sharded routing-side d_logits via JAX's normal broadcast. + if aux_loss_coeff > 0.0: + T_global = ctx.logits_2d.shape[0] + d_aux_loss_scalar = d_aux_loss.reshape(()).astype(jnp.float32) + d_aux_probs = tex.fused_moe_aux_loss_bwd( + ctx.aux_const_buf, + ctx.aux_tokens_per_expert.astype(jnp.int32), + d_aux_loss_scalar, + num_tokens=int(T_global), + ) + # routing_map is ignored by the kernel when compute_aux_scores=True, + # so pass a zero placeholder of the right shape/dtype. + zero_routing_map = jnp.zeros(ctx.aux_saved_scores.shape, dtype=ctx.routing_map.dtype) + d_logits_aux = tex.fused_topk_with_score_function_bwd( + zero_routing_map, + ctx.aux_saved_scores, + d_aux_probs.astype(ctx.aux_saved_scores.dtype), + topk=K, + use_pre_softmax=False, + scaling_factor=1.0, + score_function=score_function, + compute_aux_scores=True, + ) + d_logits_2d = d_logits_2d + d_logits_aux.astype(d_logits_2d.dtype) + + # ---------------- Gate bwd (global view) ---------------- + d_gate_logits = d_logits_2d.reshape(B, S, num_experts) + gate_kernel_cast = ctx.gate_kernel.astype(ctx.x.dtype) + d_x_from_gate = jnp.einsum("bse,he->bsh", d_gate_logits, gate_kernel_cast) + d_gate_kernel = jnp.einsum("bsh,bse->he", ctx.x, d_gate_logits).astype( + ctx.gate_kernel.dtype + ) + d_x = d_x_from_gate + d_x_from_dispatch + + # Pin output grads to the declared logical axes so downstream + # optimizers see consistent shardings. + d_x = with_sharding_constraint_by_logical_axes(d_x, input_axes) + d_gate_kernel = with_sharding_constraint_by_logical_axes(d_gate_kernel, gate_kernel_axes) + d_wi = with_sharding_constraint_by_logical_axes(d_wi, wi_kernel_axes) + d_wo = with_sharding_constraint_by_logical_axes(d_wo, wo_kernel_axes) + if has_bias: + wi_bias_axes = (wi_kernel_axes[0], *wi_kernel_axes[2:]) + wo_bias_axes = (wo_kernel_axes[0], *wo_kernel_axes[2:]) + d_wi_0_bias = with_sharding_constraint_by_logical_axes(d_wi_0_bias, wi_bias_axes) + d_wi_1_bias = with_sharding_constraint_by_logical_axes(d_wi_1_bias, wi_bias_axes) + d_wo_bias = with_sharding_constraint_by_logical_axes(d_wo_bias, wo_bias_axes) + + # expert_bias has no learnable bwd path through fused_topk: the + # primitive's bwd returns None for the bias slot. Match that with a + # zero cotangent of the right shape so custom_vjp's arity check + # passes. + d_expert_bias = jnp.zeros_like(ctx.expert_bias) + + return ( + d_x, + d_gate_kernel, + d_wi, + d_wo, + d_wi_0_bias if has_bias else None, + d_wi_1_bias if has_bias else None, + d_wo_bias if has_bias else None, + d_expert_bias, + ctx.quantizer_sets, + ) # ============================================================================= @@ -1670,7 +1809,7 @@ def _ffn_bwd_body(*args): # ============================================================================= -@partial(jax.custom_vjp, nondiff_argnums=tuple(range(9, 33))) +@partial(jax.custom_vjp, nondiff_argnums=tuple(range(9, 32))) def _moe( x, gate_kernel, @@ -1690,8 +1829,7 @@ def _moe( group_topk, scaling_factor, aux_loss_coeff, - ep_axis, - data_parallelism_axes, + mesh_resource, input_axes, gate_kernel_axes, wi_kernel_axes, @@ -1704,7 +1842,7 @@ def _moe( wi_1_checkpoint_name, wo_checkpoint_name, dispatch_checkpoint_name, - weight_gather, + quant_before_fsdp_ag, ): primal, _ = _moe_fwd_rule( x, @@ -1725,8 +1863,7 @@ def _moe( group_topk, scaling_factor, aux_loss_coeff, - ep_axis, - data_parallelism_axes, + mesh_resource, input_axes, gate_kernel_axes, wi_kernel_axes, @@ -1739,7 +1876,7 @@ def _moe( wi_1_checkpoint_name, wo_checkpoint_name, dispatch_checkpoint_name, - weight_gather, + quant_before_fsdp_ag, ) return primal @@ -1771,8 +1908,10 @@ def moe( noop_quantizer_set, noop_quantizer_set, ), - ep_axis: str, - data_parallelism_axes: Tuple[str, ...] = (), + mesh_resource: Optional[MeshResource] = None, + quant_before_fsdp_ag: bool = False, + ep_axis: Optional[str] = None, + data_parallelism_axes: Optional[Tuple[str, ...]] = None, input_axes: Tuple[Optional[str], ...] = (), gate_kernel_axes: Tuple[Optional[str], ...] = (), wi_kernel_axes: Tuple[Optional[str], ...] = ("exp", "embed", "mlp"), @@ -1783,7 +1922,7 @@ def moe( wi_1_checkpoint_name: Optional[str] = None, wo_checkpoint_name: Optional[str] = None, dispatch_checkpoint_name: Optional[str] = None, - weight_gather: WeightGather = WeightGather.full_precision(), + weight_gather: Optional[WeightGather] = None, ) -> Tuple[jnp.ndarray, Optional[jnp.ndarray], jnp.ndarray]: """Run a full MoE block under a single fused custom_vjp on the TE EP path. @@ -1832,13 +1971,20 @@ def moe( dispatch_checkpoint_name : Optional[str] JAX rematerialization checkpoint name for the EP dispatch outputs. ``None`` leaves these values unnamed; the opaque prepare handle is never named. - weight_gather : WeightGather - Expert-weight gather policy. ``WeightGather.full_precision()`` gathers - weights before quantization (the default). Use - ``WeightGather.quantized()`` to quantize local shards and gather their - MXFP8 data and scales along the active ``MeshResource.fsdp_resource``; - pass ``axis`` to select a physical mesh axis without consulting it. - + mesh_resource : Optional[MeshResource] + Physical parallelism resources. ``None`` uses the active + ``global_shard_guard`` context; an explicit resource takes precedence. + One of these is required. Batch sharding combines ``dp_resource`` then + ``fsdp_resource`` (omitting unset/duplicate axes), with ``ep_resource`` + innermost. EP must be set and distinct from the outer axes. + quant_before_fsdp_ag : bool + Quantize expert-weight shards before gathering their MXFP8 data and + scales on ``mesh_resource.fsdp_resource``. Default ``False`` gathers + full-precision weights first. ``True`` requires an FSDP resource and + MXFP8 kernel quantizers. + ep_axis, data_parallelism_axes, weight_gather : deprecated + Compatibility arguments converted into a MeshResource and boolean, + with a DeprecationWarning. Conflicting old and new arguments raise. Note that the per-expert dispatch-slot alignment is fixed internally at 128 tokens (``_ALIGN_SIZE``); see that constant's docstring for rationale and how to extend if a future recipe needs >128. @@ -1848,136 +1994,137 @@ def moe( The fused path uses 256-token expert alignment. Ineligible calls warn and fall back to TE's regular grouped-GEMM implementation. - Axis-name parameters: - - * ``ep_axis`` and ``data_parallelism_axes`` are *physical mesh - axis names* -- they index ``jax.sharding.Mesh.shape`` directly - (to compute ``num_ep`` / ``dp_size`` and to construct - ``P((dp..., ep), None, None)`` for the physical - ``jax.lax.with_sharding_constraint`` calls that JAX requires - to refer to real mesh axes). - * ``input_axes``, ``gate_kernel_axes``, ``wi_kernel_axes``, - ``wo_kernel_axes`` are *logical axis names* (e.g. - ``"batch"``, ``"embed"``, ``"mlp"``, ``"exp"``) -- they get - resolved via the active Flax logical-axis rules and consumed - by ``with_sharding_constraint_by_logical_axes``. They are - ``Optional[str]`` tuples so a rule of ``None`` means - "replicated on this axis". - - Logical-axis support for ``ep_axis`` / ``data_parallelism_axes`` - is intentionally out of scope: the EP comm-group construction - (``dp_color = rank // ep_size``) and the bootstrap signature - check both require concrete integer sizes, so a logical name - would have to be resolved to a physical one anyway before any - EP primitive is called. If a downstream pipeline needs to plumb - logical names all the way to ``moe()``, do the rule lookup at - the call site. + MeshResource fields name physical mesh axes, not Flax logical axes. + ``input_axes``, ``gate_kernel_axes``, ``wi_kernel_axes`` and + ``wo_kernel_axes`` remain logical-axis tuples resolved through the active + Flax rules. The selected MeshResource is also used by internal TE sharding + in both forward and backward, independently of an enclosing global context. See module docstring for the rest of the parameter semantics and the surrounding design rationale. """ + if ep_axis is not None or data_parallelism_axes is not None or weight_gather is not None: + call_args = locals().copy() + resource, quantize = _resolve_moe_mesh_resource( + mesh_resource, quant_before_fsdp_ag, ep_axis, data_parallelism_axes, weight_gather + ) + for name in ("ep_axis", "data_parallelism_axes", "weight_gather"): + call_args.pop(name) + call_args.update(mesh_resource=resource, quant_before_fsdp_ag=quantize) + return moe(**call_args) + + mesh_resource, quant_before_fsdp_ag = _resolve_moe_mesh_resource( + mesh_resource, quant_before_fsdp_ag + ) + ep_axis, data_parallelism_axes = _moe_mesh_axes(mesh_resource) score_function = _validate_score_function(score_function) - if not isinstance(weight_gather, WeightGather): - raise TypeError("weight_gather must be a WeightGather policy.") - # Enforce ((outer_dp..., ep), None, None) on inbound activations. The - # EP comm groups consecutive global ranks (dp_color = rank // ep_size), - # so ep MUST be innermost in the partition spec. Soft re-pin: free if - # upstream already matches, single reshard otherwise. - mesh = _get_mesh() - if mesh is None or mesh.empty: - raise ValueError("moe(...) requires an active jax.sharding.Mesh.") - expected_leading: Any = (*data_parallelism_axes, ep_axis) if data_parallelism_axes else ep_axis - expected_spec = P(expected_leading, None, None) - actual_spec = getattr(getattr(x, "sharding", None), "spec", None) - if actual_spec is not None and tuple(actual_spec) != tuple(expected_spec): - warnings.warn( - f"moe(...): inbound x sharding {actual_spec} does not match expected " - f"{expected_spec}; inserting a reshard. Apply " - "jax.lax.with_sharding_constraint upstream to avoid this overhead.", - UserWarning, - stacklevel=2, + with global_shard_guard(mesh_resource): + # Enforce ((outer_dp..., ep), None, None) on inbound activations. The + # EP comm groups consecutive global ranks (dp_color = rank // ep_size), + # so ep MUST be innermost in the partition spec. Soft re-pin: free if + # upstream already matches, single reshard otherwise. + mesh = _get_mesh() + if mesh is None or mesh.empty: + raise ValueError("moe(...) requires an active jax.sharding.Mesh.") + for axis in (ep_axis, *data_parallelism_axes): + if axis not in mesh.shape: + raise ValueError(f"TE MoE resource axis {axis!r} is not in the active mesh.") + expected_leading: Any = ( + (*data_parallelism_axes, ep_axis) if data_parallelism_axes else ep_axis ) - x = _with_sharding_constraint_cast_bwd(x, NamedSharding(mesh, expected_spec)) + expected_spec = P(expected_leading, None, None) + actual_spec = getattr(getattr(x, "sharding", None), "spec", None) + if actual_spec is not None and tuple(actual_spec) != tuple(expected_spec): + warnings.warn( + f"moe(...): inbound x sharding {actual_spec} does not match expected " + f"{expected_spec}; inserting a reshard. Apply " + "jax.lax.with_sharding_constraint upstream to avoid this overhead.", + UserWarning, + stacklevel=2, + ) + x = _with_sharding_constraint_cast_bwd(x, NamedSharding(mesh, expected_spec)) - # custom_vjp can't trace through None args; lower expert_bias to an - # empty shape-(0,) tensor that fused_topk_with_score_function treats - # as "no bias". - if expert_bias is None: - expert_bias_arg = jnp.zeros((0,), dtype=jnp.float32) - else: - expert_bias_arg = expert_bias.astype(jnp.float32) + # custom_vjp can't trace through None args; lower expert_bias to an + # empty shape-(0,) tensor that fused_topk_with_score_function treats + # as "no bias". + if expert_bias is None: + expert_bias_arg = jnp.zeros((0,), dtype=jnp.float32) + else: + expert_bias_arg = expert_bias.astype(jnp.float32) + + use_cudnn_jax_fusion = False + cudnn_native_weight_layout = wi.ndim == 3 and wi.shape[-1] == x.shape[-1] + if _use_cudnn_cutedsl_fusion_from_env(): + rejection_reasons = _cudnn_jax_fusion_rejection_reasons( + x, + wi, + wi_0_bias, + wi_1_bias, + quantizer_sets, + num_experts=num_experts, + activation_type=activation_type, + ep_axis=ep_axis, + ) + if rejection_reasons: + if cudnn_native_weight_layout: + raise ValueError( + "cuDNN-native MoE weight layout requires the fused cuDNN grouped-GEMM " + "path, which is unsupported for this moe() call: " + + "; ".join(rejection_reasons) + ) + warnings.warn( + f"{_CUDNN_JAX_ENV}=1 is unsupported for this moe() call; falling back to " + "the regular TE grouped-GEMM path: " + + "; ".join(rejection_reasons), + UserWarning, + stacklevel=2, + ) + else: + use_cudnn_jax_fusion = True + elif cudnn_native_weight_layout: + raise ValueError( + "cuDNN-native MoE weight layout requires the fused cuDNN grouped-GEMM path; " + f"set {_CUDNN_JAX_ENV}=1." + ) - use_cudnn_jax_fusion = False - cudnn_native_weight_layout = wi.ndim == 3 and wi.shape[-1] == x.shape[-1] - if _use_cudnn_cutedsl_fusion_from_env(): - rejection_reasons = _cudnn_jax_fusion_rejection_reasons( + output, aux_loss, total_recv_tokens = _moe( x, + gate_kernel, wi, + wo, wi_0_bias, wi_1_bias, + wo_bias, + expert_bias_arg, quantizer_sets, - num_experts=num_experts, - activation_type=activation_type, - ep_axis=ep_axis, + num_experts, + num_experts_per_tok, + activation_type, + score_function, + use_pre_softmax, + num_groups, + group_topk, + scaling_factor, + float(aux_loss_coeff), + mesh_resource, + input_axes, + gate_kernel_axes, + wi_kernel_axes, + wo_kernel_axes, + dtype, + apply_topk_weights_early, + recv_capacity_per_rank, + use_cudnn_jax_fusion, + wi_0_checkpoint_name, + wi_1_checkpoint_name, + wo_checkpoint_name, + dispatch_checkpoint_name, + quant_before_fsdp_ag, ) - if rejection_reasons: - if cudnn_native_weight_layout: - raise ValueError( - "cuDNN-native MoE weight layout requires the fused cuDNN grouped-GEMM " - "path, which is unsupported for this moe() call: " - + "; ".join(rejection_reasons) - ) - warnings.warn( - f"{_CUDNN_JAX_ENV}=1 is unsupported for this moe() call; falling back to " - "the regular TE grouped-GEMM path: " + "; ".join(rejection_reasons), - UserWarning, - stacklevel=2, - ) - else: - use_cudnn_jax_fusion = True - elif cudnn_native_weight_layout: - raise ValueError( - "cuDNN-native MoE weight layout requires the fused cuDNN grouped-GEMM path; " - f"set {_CUDNN_JAX_ENV}=1." - ) - - output, aux_loss, total_recv_tokens = _moe( - x, - gate_kernel, - wi, - wo, - wi_0_bias, - wi_1_bias, - wo_bias, - expert_bias_arg, - quantizer_sets, - num_experts, - num_experts_per_tok, - activation_type, - score_function, - use_pre_softmax, - num_groups, - group_topk, - scaling_factor, - float(aux_loss_coeff), - ep_axis, - data_parallelism_axes, - input_axes, - gate_kernel_axes, - wi_kernel_axes, - wo_kernel_axes, - dtype, - apply_topk_weights_early, - recv_capacity_per_rank, - use_cudnn_jax_fusion, - wi_0_checkpoint_name, - wi_1_checkpoint_name, - wo_checkpoint_name, - dispatch_checkpoint_name, - weight_gather, - ) - if aux_loss_coeff <= 0.0: - aux_loss = None - assert output.dtype == x.dtype, f"moe() output dtype {output.dtype} != input dtype {x.dtype}" - return output, aux_loss, total_recv_tokens + if aux_loss_coeff <= 0.0: + aux_loss = None + assert ( + output.dtype == x.dtype + ), f"moe() output dtype {output.dtype} != input dtype {x.dtype}" + return output, aux_loss, total_recv_tokens From f63ebbb6aea6bdbf474caecae2a3a3ce9f39ad85 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Mon, 5 Oct 2026 12:34:57 -0700 Subject: [PATCH 22/38] Support expert-axis quantized FSDP gathers in JAX MoE Signed-off-by: Jeremy Berchtold --- docs/api/jax.rst | 8 ++ tests/jax/test_te_ep_moe.py | 22 +++-- transformer_engine/jax/moe.py | 146 ++++++++++++++++++++++------------ 3 files changed, 121 insertions(+), 55 deletions(-) diff --git a/docs/api/jax.rst b/docs/api/jax.rst index 947d22a2ea7..5370ba8a9fe 100644 --- a/docs/api/jax.rst +++ b/docs/api/jax.rst @@ -90,6 +90,14 @@ all-gathering their MXFP8 data and scales on the resource's FSDP axis. It requires an FSDP resource and MXFP8 kernel quantizers. The default gathers full-precision weights before quantization. +With active Flax logical-axis rules, ``wi_kernel_axes`` and ``wo_kernel_axes`` +also select the forward FFN weight input specs when quantized gathering is +enabled. FSDP can shard complete experts together with EP on the first weight +dimension, or split a matrix dimension within each expert. Complete-expert +gathering concatenates the quantized matrices and their existing scale blocks; +matrix-shard gathering reconstructs each expert's scale blocks. These specs +are resolved at trace time and do not require JAX Explicit sharding mode. + The old ``ep_axis``, ``data_parallelism_axes`` and ``weight_gather`` arguments remain accepted with a ``DeprecationWarning``. They are translated into a resource and the boolean before calling the new API. Conflicting old and diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index 8c949b5644f..00ddf7a7d37 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -148,6 +148,7 @@ def _read_mp_options(): LOGICAL_AXIS_RULES = ( ("exp", EP_AXIS), + ("expert_weight_fsdp", (EP_AXIS, FSDP_AXIS)), ("embed", FSDP_AXIS), ("mlp", None), ("batch", (FSDP_AXIS, EP_AXIS)), @@ -586,26 +587,35 @@ def test_weight_gather_policy_axis_resolution(mesh): @pytest.mark.parametrize("native_weight_layout", [False, True]) +@pytest.mark.parametrize("expert_fsdp", [False, True]) def test_quantized_weight_gather_matches_full_precision_gather( - mesh, monkeypatch, native_weight_layout + mesh, monkeypatch, native_weight_layout, expert_fsdp ): """The FP8 weight gather retains forward and backward MoE semantics.""" if native_weight_layout: if not _use_cudnn_cutedsl_fusion_from_env(): pytest.skip("Native weight layout requires cuDNN grouped GEMM fusion") + if native_weight_layout or expert_fsdp: from transformer_engine.jax import cpp_extensions as tex flax_moe_module = importlib.import_module("transformer_engine.jax.flax.moe") original_moe = flax_moe_module.moe - def native_moe(*args, **kwargs): + def layout_moe(*args, **kwargs): args = list(args) - gate, up = jnp.split(args[2], 2, axis=-1) - args[2] = tex.pack_swiglu_pair(gate, up).transpose(0, 2, 1) - kwargs["wi_kernel_axes"] = ("exp", "mlp", "embed") + if native_weight_layout: + gate, up = jnp.split(args[2], 2, axis=-1) + args[2] = tex.pack_swiglu_pair(gate, up).transpose(0, 2, 1) + kwargs["wi_kernel_axes"] = ("exp", "mlp", "embed") + if expert_fsdp: + kwargs["wi_kernel_axes"] = ("expert_weight_fsdp", None, None) + kwargs["wo_kernel_axes"] = ("expert_weight_fsdp", None, None) + spec = P((EP_AXIS, FSDP_AXIS), None, None) + args[2] = jax.lax.with_sharding_constraint(args[2], NamedSharding(mesh, spec)) + args[3] = jax.lax.with_sharding_constraint(args[3], NamedSharding(mesh, spec)) return original_moe(*args, **kwargs) - monkeypatch.setattr(flax_moe_module, "moe", native_moe) + monkeypatch.setattr(flax_moe_module, "moe", layout_moe) x = _make_inputs(jax.random.PRNGKey(41)) baseline = _make_block(quantization_recipe=MXFP8BlockScaling()) quantized_ag = _make_block(quantization_recipe=MXFP8BlockScaling(), quant_before_fsdp_ag=True) diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index cc0470ca85c..41389c611e6 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -38,6 +38,7 @@ from typing import Any, Literal, Optional, Tuple, Union import flax.struct +import flax.linen as nn import jax import jax.numpy as jnp from jax.ad_checkpoint import checkpoint_name @@ -619,7 +620,8 @@ def _gather_quantized_weight(tensor, fsdp_axis: str, fsdp_size: int, sharded_axi Grouped tensor data and scales are flat, with scales independently padded and swizzled for each expert. Reassemble each expert in logical order, - then pad and swizzle its gathered scales for grouped GEMM. + then pad and swizzle its gathered scales for grouped GEMM. Whole-expert + shards can instead concatenate their already-swizzled scale blocks. """ if isinstance(tensor, ScaledTensor2x): return ScaledTensor2x( @@ -632,7 +634,9 @@ def _gather_quantized_weight(tensor, fsdp_axis: str, fsdp_size: int, sharded_axi local_shape = tensor.original_shape num_experts = local_shape[0] # The T layout swaps the two matrix dimensions within each expert. - data_axis = 3 - sharded_axis if tensor.data_layout == "T" else sharded_axis + data_axis = ( + 3 - sharded_axis if tensor.data_layout == "T" and sharded_axis != 0 else sharded_axis + ) global_shape = list(local_shape) global_shape[data_axis] *= fsdp_size global_shape = tuple(global_shape) @@ -651,6 +655,48 @@ def _gather_quantized_weight(tensor, fsdp_axis: str, fsdp_size: int, sharded_axi is_padded=True, flatten_axis=tensor.flatten_axis - 1, ) + if sharded_axis == 0: + # The grouped allocation includes a worst-case padding tail. Gather + # only actual expert scale blocks, then allocate the gathered tail. + local_scale_size = math.prod(local_scale_shape) + scale_inv = jax.lax.all_gather( + tensor.scale_inv[: num_experts * local_scale_size], + fsdp_axis, + axis=0, + tiled=True, + ) + else: + scale_inv = _gather_quantized_matrix_scales( + tensor, fsdp_axis, data_axis, local_scale_shape, global_matrix + ) + expected_scale_size = tensor.scaling_mode.get_grouped_scale_shape( + global_shape, + global_shape[0], + tensor.is_colwise, + is_padded=True, + flatten_axis=tensor.flatten_axis, + )[0] + scale_inv = jnp.pad(scale_inv, (0, expected_scale_size - scale_inv.size)) + return GroupedScaledTensor1x( + data=data, + scale_inv=scale_inv, + amax=tensor.amax, + first_dims=None, + last_dims=None, + scaling_mode=tensor.scaling_mode, + dq_dtype=tensor.dq_dtype, + _dq_func=tensor._dq_func, + is_colwise=tensor.is_colwise, + data_layout=tensor.data_layout, + flatten_axis=tensor.flatten_axis, + original_shape=global_shape, + pre_swizzled=True, + ) + + +def _gather_quantized_matrix_scales(tensor, fsdp_axis, data_axis, local_scale_shape, global_matrix): + """Reassemble scales when FSDP splits each expert's matrix dimension.""" + local_matrix = tensor.original_shape[1:] local_unpadded_shape = tensor.scaling_mode.get_scale_shape( local_matrix, data_layout=tensor.data_layout, @@ -675,7 +721,7 @@ def _gather_quantized_weight(tensor, fsdp_axis: str, fsdp_size: int, sharded_axi scale_axis = data_axis - 1 local_scale_size = math.prod(local_scale_shape) gathered_scales = [] - for expert in range(num_experts): + for expert in range(tensor.original_shape[0]): local_swizzled = jax.lax.dynamic_slice_in_dim( tensor.scale_inv, expert * local_scale_size, local_scale_size ) @@ -693,29 +739,18 @@ def _gather_quantized_weight(tensor, fsdp_axis: str, fsdp_size: int, sharded_axi ), ) gathered_scales.append(swizzled_scale(full_padded, 1, tensor.is_colwise).reshape(-1)) - scale_inv = jnp.concatenate(gathered_scales) - expected_scale_size = tensor.scaling_mode.get_grouped_scale_shape( - global_shape, - num_experts, - tensor.is_colwise, - is_padded=True, - flatten_axis=tensor.flatten_axis, - )[0] - scale_inv = jnp.pad(scale_inv, (0, expected_scale_size - scale_inv.size)) - return GroupedScaledTensor1x( - data=data, - scale_inv=scale_inv, - amax=tensor.amax, - first_dims=None, - last_dims=None, - scaling_mode=tensor.scaling_mode, - dq_dtype=tensor.dq_dtype, - _dq_func=tensor._dq_func, - is_colwise=tensor.is_colwise, - data_layout=tensor.data_layout, - flatten_axis=tensor.flatten_axis, - original_shape=global_shape, - pre_swizzled=True, + return jnp.concatenate(gathered_scales) + + +def _weight_fsdp_axis(spec, fsdp_axis): + """Find the tensor dimension partitioned by the physical FSDP resource.""" + return next( + ( + i + for i, axes in enumerate(spec) + if fsdp_axis in (axes if isinstance(axes, tuple) else (axes,)) + ), + None, ) @@ -741,6 +776,8 @@ def _ffn_fwd_per_shard( quant_before_fsdp_ag: bool, fsdp_axis: Optional[str], fsdp_size: int, + wi_fsdp_axis: Optional[int], + wo_fsdp_axis: Optional[int], ): """Run the grouped FFN on one shard's EP receive buffer.""" hidden = recv_tokens_local.shape[-1] @@ -779,12 +816,12 @@ def _ffn_fwd_per_shard( flatten_axis=-1, ) casted_wi = tex.grouped_quantize(wi_for_gemm, fc1_quantizer_set.kernel, flatten_axis=-1) - if quant_before_fsdp_ag: + if quant_before_fsdp_ag and wi_fsdp_axis is not None: casted_wi = _gather_quantized_weight( casted_wi, fsdp_axis, fsdp_size, - 2 if cudnn_native_weight_layout else 1, + wi_fsdp_axis, ) casted_intermediate = None if use_cudnn_jax_fusion: @@ -911,8 +948,8 @@ def _ffn_fwd_per_shard( flatten_axis=-1, ) casted_wo = tex.grouped_quantize(wo, fc2_quantizer_set.kernel, flatten_axis=-1) - if quant_before_fsdp_ag: - casted_wo = _gather_quantized_weight(casted_wo, fsdp_axis, fsdp_size, 2) + if quant_before_fsdp_ag and wo_fsdp_axis is not None: + casted_wo = _gather_quantized_weight(casted_wo, fsdp_axis, fsdp_size, wo_fsdp_axis) expert_outputs = tex.grouped_gemm( casted_intermediate.get_tensor(usage=TensorUsage.LHS), casted_wo.get_tensor(usage=TensorUsage.RHS), @@ -1174,7 +1211,7 @@ def _moe_fwd_rule( """ with global_shard_guard(mesh_resource): ep_axis, data_parallelism_axes = _moe_mesh_axes(mesh_resource) - del gate_kernel_axes, wi_kernel_axes, wo_kernel_axes # used in bwd only + del gate_kernel_axes # used in bwd only from jax.experimental.shard_map import shard_map x = with_sharding_constraint_by_logical_axes(x, input_axes) @@ -1201,17 +1238,37 @@ def _moe_fwd_rule( B, S, H = x.shape K = num_experts_per_tok cudnn_native_weight_layout = wi.ndim == 3 and wi.shape[-1] == H - wi_hidden_axis = 2 if cudnn_native_weight_layout else 1 + kernel_spec = P(ep_axis, None, None) + wi_input_spec = wo_input_spec = kernel_spec + wi_fsdp_axis = wo_fsdp_axis = None if quant_before_fsdp_ag: + # Logical axes describe parameter storage even under Auto mode, + # where tracer types do not expose the physical input sharding. + if nn.get_logical_axis_rules(): + wi_input_spec = nn.logical_to_mesh_axes(wi_kernel_axes) + wo_input_spec = nn.logical_to_mesh_axes(wo_kernel_axes) + else: + # Preserve the existing no-Flax-rules convention. + wi_input_spec = ( + P(ep_axis, None, mesh_resource.fsdp_resource) + if cudnn_native_weight_layout + else P(ep_axis, mesh_resource.fsdp_resource, None) + ) + wo_input_spec = P(ep_axis, None, mesh_resource.fsdp_resource) + wi_fsdp_axis = _weight_fsdp_axis(wi_input_spec, mesh_resource.fsdp_resource) + wo_fsdp_axis = _weight_fsdp_axis(wo_input_spec, mesh_resource.fsdp_resource) if mesh_resource.fsdp_resource not in data_parallelism_axes: raise ValueError( "Quantized weight all-gather requires its FSDP axis among the outer batch axes." ) if any(quantizer_set.kernel is None for quantizer_set in quantizer_sets): raise ValueError("Quantized weight all-gather requires MXFP8 kernel quantizers.") - if wi.shape[wi_hidden_axis] % ( - mesh.shape[mesh_resource.fsdp_resource] * 32 - ) or wo.shape[2] % (mesh.shape[mesh_resource.fsdp_resource] * 32): + if any( + axis is not None + and axis != 0 + and weight.shape[axis] % (mesh.shape[mesh_resource.fsdp_resource] * 32) + for weight, axis in ((wi, wi_fsdp_axis), (wo, wo_fsdp_axis)) + ): raise ValueError("FSDP weight shards must be divisible by the MXFP8 block size 32.") if B % num_procs != 0: @@ -1369,19 +1426,6 @@ def _moe_fwd_rule( # ---------------- FFN (per-shard via shard_map) ---------------- has_bias = wi_0_bias is not None - kernel_spec = P(ep_axis, None, None) - wi_input_spec = ( - ( - P(ep_axis, None, mesh_resource.fsdp_resource) - if cudnn_native_weight_layout - else P(ep_axis, mesh_resource.fsdp_resource, None) - ) - if quant_before_fsdp_ag - else kernel_spec - ) - wo_input_spec = ( - P(ep_axis, None, mesh_resource.fsdp_resource) if quant_before_fsdp_ag else kernel_spec - ) bias_spec = P(ep_axis, None) ffn_in_specs = (ep3_spec, ep2_spec, ep2_spec, wi_input_spec, wo_input_spec) ffn_in_args = [recv_tokens, recv_topk_weights, token_counts, wi, wo] @@ -1438,6 +1482,8 @@ def _ffn_fwd_body(*args): if mesh_resource.fsdp_resource is not None else 1 ), + wi_fsdp_axis=wi_fsdp_axis, + wo_fsdp_axis=wo_fsdp_axis, ) expert_outputs, ffn_residuals = shard_map( @@ -1981,7 +2027,9 @@ def moe( Quantize expert-weight shards before gathering their MXFP8 data and scales on ``mesh_resource.fsdp_resource``. Default ``False`` gathers full-precision weights first. ``True`` requires an FSDP resource and - MXFP8 kernel quantizers. + MXFP8 kernel quantizers. With active Flax logical-axis rules, the weight + input specs follow ``wi_kernel_axes`` and ``wo_kernel_axes``; FSDP may + shard whole experts or a matrix dimension within each expert. ep_axis, data_parallelism_axes, weight_gather : deprecated Compatibility arguments converted into a MeshResource and boolean, with a DeprecationWarning. Conflicting old and new arguments raise. From ccf8a9d9245e4d4e42aee63235c74776d02df815 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Tue, 6 Oct 2026 11:39:36 -0700 Subject: [PATCH 23/38] Add capability-based fallbacks for cuDNN JAX MoE fusion Signed-off-by: Jeremy Berchtold --- .../jax/test_grouped_gemm_swiglu_fallback.py | 251 ++++++++++++++++++ tests/jax/test_te_ep_moe.py | 62 ++++- .../jax/cpp_extensions/grouped_gemm_swiglu.py | 73 +++-- transformer_engine/jax/moe.py | 105 ++++---- 4 files changed, 426 insertions(+), 65 deletions(-) create mode 100644 tests/jax/test_grouped_gemm_swiglu_fallback.py diff --git a/tests/jax/test_grouped_gemm_swiglu_fallback.py b/tests/jax/test_grouped_gemm_swiglu_fallback.py new file mode 100644 index 00000000000..996079c9f7c --- /dev/null +++ b/tests/jax/test_grouped_gemm_swiglu_fallback.py @@ -0,0 +1,251 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. +"""cuDNN MoE API compatibility and ordered dispatch without kernel compilation.""" + +import importlib +import inspect +import warnings +from types import SimpleNamespace + +import pytest +import transformer_engine_jax + +adapter = importlib.import_module("transformer_engine.jax.cpp_extensions.grouped_gemm_swiglu") +moe = importlib.import_module("transformer_engine.jax.moe") + + +def api(arguments, change=None): + def function(**kwargs): + raise AssertionError("Probes must not execute kernels") + + parameters = [ + inspect.Parameter( + name, + ( + inspect.Parameter.POSITIONAL_ONLY + if change == "positional" and name == "a_tensor" + else inspect.Parameter.POSITIONAL_OR_KEYWORD + ), + ) + for name in arguments + if not (change == "renamed" and name == "discrete_col_sfd") + ] + if change in ("optional", "mandatory"): + parameters.append( + inspect.Parameter( + "new_arg", + inspect.Parameter.KEYWORD_ONLY, + default=None if change == "optional" else inspect.Parameter.empty, + ) + ) + function.__signature__ = inspect.Signature(parameters) + return function + + +@pytest.fixture +def frontend(monkeypatch): + module = SimpleNamespace( + grouped_gemm_glu=api(adapter._FORWARD_ARGS), + grouped_gemm_swiglu=api(adapter._FORWARD_ARGS), + grouped_gemm_dswiglu=api(adapter._BACKWARD_ARGS), + ) + original = importlib.import_module + + def import_module(name, package=None): + if name == "cudnn.jax": + return module + if name == "cutlass.jax": + return SimpleNamespace(is_available=lambda: True) + return original(name, package) + + monkeypatch.setattr(adapter.importlib, "import_module", import_module) + return module + + +@pytest.mark.parametrize("rubin", [False, True]) +@pytest.mark.parametrize("operation", ["forward", "backward"]) +@pytest.mark.parametrize( + "change", + [ + "missing", + "not_callable", + "renamed", + "mandatory", + "optional", + "positional", + "opaque", + ], +) +def test_api_signature(frontend, rubin, operation, change): + name = ( + ("grouped_gemm_glu" if rubin else "grouped_gemm_swiglu") + if operation == "forward" + else "grouped_gemm_dswiglu" + ) + args = adapter._FORWARD_ARGS if operation == "forward" else adapter._BACKWARD_ARGS + if change == "missing": + delattr(frontend, name) + elif change == "not_callable": + setattr(frontend, name, 42) + elif change == "opaque": + setattr(frontend, name, lambda **kwargs: None) + else: + setattr(frontend, name, api(args, change)) + available, reason = adapter.grouped_gemm_swiglu_dependencies_available(rubin) + assert available == (change == "optional") + assert bool(reason) == (change != "optional") + + +@pytest.mark.parametrize("capability", [90, 100, 103, 107, 120]) +@pytest.mark.parametrize( + "missing", + [None, "grouped_gemm_glu", "grouped_gemm_swiglu", "grouped_gemm_dswiglu", "all"], +) +def test_ordered_fallback(frontend, monkeypatch, capability, missing): + if missing == "all": + vars(frontend).clear() + elif missing: + delattr(frontend, missing) + monkeypatch.setattr( + transformer_engine_jax, "get_device_compute_capability", lambda _: capability + ) + expected = False + if capability >= 100 and missing not in ("grouped_gemm_dswiglu", "all"): + if capability == 107 and missing != "grouped_gemm_glu": + expected = "rubin" + elif missing != "grouped_gemm_swiglu": + expected = "blackwell" + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + assert moe._select_cudnn_jax_fusion([]) == expected + if expected == "rubin": + assert not caught + else: + assert len(caught) == 1 + message = str(caught[0].message) + assert "falling back" in message and "1.31.0" in message + assert ("generic Blackwell+ fused" if expected else "unfused TE") in message + + +@pytest.mark.parametrize("name", ["cudnn.jax", "cutlass.jax"]) +def test_missing_dependency(monkeypatch, name): + original = importlib.import_module + + def import_module(module, package=None): + if module == name: + raise ModuleNotFoundError(f"No module named {module}") + return original(module, package) + + monkeypatch.setattr(adapter.importlib, "import_module", import_module) + available, reason = adapter.grouped_gemm_swiglu_dependencies_available() + assert not available and name in reason + + +def test_ineligible_call_and_device_query_failure(frontend, monkeypatch): + monkeypatch.setattr(transformer_engine_jax, "get_device_compute_capability", lambda _: 107) + with pytest.warns(UserWarning, match="bias unsupported"): + assert moe._select_cudnn_jax_fusion(["bias unsupported"]) is False + + def fail(_): + raise RuntimeError("device query failed") + + monkeypatch.setattr(transformer_engine_jax, "get_device_compute_capability", fail) + with pytest.warns(UserWarning, match="device query failed"): + assert moe._select_cudnn_jax_fusion([]) is False + + +def test_installed_frontend_contract(): + for rubin in (False, True): + available, reason = adapter.grouped_gemm_swiglu_dependencies_available(rubin) + assert available or reason + + +@pytest.mark.parametrize("native_layout", [False, True]) +@pytest.mark.parametrize( + "missing,expected", + [(None, "rubin"), ("grouped_gemm_glu", "blackwell"), ("grouped_gemm_dswiglu", False)], +) +def test_public_moe_passes_selected_path(frontend, monkeypatch, native_layout, missing, expected): + import jax + import jax.numpy as jnp + import numpy as np + from jax.sharding import Mesh + from transformer_engine.jax.sharding import MeshResource + + if missing: + delattr(frontend, missing) + monkeypatch.setenv(moe._CUDNN_JAX_ENV, "1") + monkeypatch.setattr(transformer_engine_jax, "get_device_compute_capability", lambda _: 107) + monkeypatch.setattr(moe, "_cudnn_jax_fusion_rejection_reasons", lambda *args, **kwargs: []) + mesh = Mesh(np.asarray(jax.devices()[:1]), ("ep",)) + monkeypatch.setattr(moe, "_get_mesh", lambda: mesh) + monkeypatch.setattr(moe, "_with_sharding_constraint_cast_bwd", lambda x, _: x) + received = [] + signature = inspect.signature(moe._moe) + + def execute(*args): + received.append(signature.bind(*args).arguments["use_cudnn_jax_fusion"]) + return args[0], None, jnp.asarray(0) + + monkeypatch.setattr(moe, "_moe", execute) + x = jnp.ones((1, 1, 128), jnp.bfloat16) + wi = jnp.ones((2, 256, 128) if native_layout else (2, 128, 256), jnp.bfloat16) + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + output, _, _ = moe.moe( + x, + jnp.ones((128, 2)), + wi, + jnp.ones((2, 128, 128)), + num_experts=2, + num_experts_per_tok=1, + mesh_resource=MeshResource(ep_resource="ep"), + ) + assert received == [expected] + assert output is x + assert len(caught) == (0 if expected == "rubin" else 1) + + +@pytest.mark.parametrize("path", ["rubin", "blackwell"]) +def test_forward_uses_selected_kernel(monkeypatch, path): + import jax.numpy as jnp + + class SelectedKernel(Exception): + pass + + def selected(*args, **kwargs): + raise SelectedKernel + + def rejected(*args, **kwargs): + raise AssertionError("Forward ignored the selected fallback") + + monkeypatch.setattr(moe.tex, "grouped_gemm_glu", selected if path == "rubin" else rejected) + monkeypatch.setattr( + moe.tex, "grouped_gemm_swiglu", selected if path == "blackwell" else rejected + ) + + def quantize(data, *args, **kwargs): + return SimpleNamespace( + get_tensor=lambda **kwargs: SimpleNamespace( + data=data, scale_inv=jnp.ones((1,), dtype=jnp.uint8) + ) + ) + + monkeypatch.setattr(moe.tex, "grouped_quantize", quantize) + quantizers = SimpleNamespace(x=SimpleNamespace(q_dtype=jnp.float8_e4m3fn), kernel=None) + kwargs = dict.fromkeys(inspect.signature(moe._ffn_fwd_per_shard).parameters) + kwargs.update( + recv_tokens_local=jnp.ones((1, 256, 128), jnp.bfloat16), + recv_topk_weights_local=jnp.ones((1, 256)), + token_counts_local=jnp.asarray([256]), + wi=jnp.ones((1, 128, 256), jnp.bfloat16), + wo=jnp.ones((1, 128, 128), jnp.bfloat16), + quantizer_sets=(quantizers, quantizers), + num_local_experts=1, + use_cudnn_jax_fusion=path, + cudnn_native_weight_layout=False, + quant_before_fsdp_ag=False, + ) + with pytest.raises(SelectedKernel): + moe._ffn_fwd_per_shard(**kwargs) diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index 00ddf7a7d37..8e1c44e3332 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -867,6 +867,49 @@ def recorded_checkpoint_name(value, name): class TestTeEpMoeCudnnCutedslFusion: """End-to-end MXFP8 coverage for cuDNN's grouped GLU JAX APIs.""" + @pytest.mark.parametrize("native_layout", [False, True]) + def test_missing_apis_unfused_forward_and_backward(self, mesh, monkeypatch, native_layout): + """Unfused fallback retains bootstrap capacity and native-layout gradients.""" + if not _use_cudnn_cutedsl_fusion_from_env(): + pytest.skip("Requires fusion requested at bootstrap") + import cudnn.jax as cudnn_jax + from transformer_engine.jax import cpp_extensions as tex + + monkeypatch.setattr(cudnn_jax, "grouped_gemm_glu", None) + monkeypatch.setattr(cudnn_jax, "grouped_gemm_swiglu", None) + block = _make_block(quantization_recipe=MXFP8BlockScaling()) + x = _make_inputs(jax.random.PRNGKey(51)) + variables, baseline_output, _ = _init_apply(block, mesh, x, jax.random.PRNGKey(52)) + baseline_grads, baseline_dx = _grad_step(block, variables, mesh, x) + + if native_layout: + flax_moe = importlib.import_module("transformer_engine.jax.flax.moe") + original_moe = flax_moe.moe + + def native_moe(*args, **kwargs): + args = list(args) + gate, up = jnp.split(args[2], 2, axis=-1) + args[2] = tex.pack_swiglu_pair(gate, up).transpose(0, 2, 1) + kwargs["wi_kernel_axes"] = ("exp", "mlp", "embed") + return original_moe(*args, **kwargs) + + monkeypatch.setattr(flax_moe, "moe", native_moe) + with _ctx(mesh): + output, _, _ = jax.jit(block.apply)(variables, _shard_inputs(x, mesh)) + output.block_until_ready() + grads, dx = _grad_step(block, variables, mesh, x) + np.testing.assert_array_equal( + _to_global_numpy(output, mesh), _to_global_numpy(baseline_output, mesh) + ) + np.testing.assert_array_equal( + _to_global_numpy(dx, mesh), _to_global_numpy(baseline_dx, mesh) + ) + for name in ("gate_kernel", "wi", "wo"): + np.testing.assert_array_equal( + _to_global_numpy(_unwrap(grads["params"][name]), mesh), + _to_global_numpy(_unwrap(baseline_grads["params"][name]), mesh), + ) + @pytest.mark.parametrize("apply_topk_weights_early", [False, True]) def test_mxfp8_forward_and_backward(self, mesh, apply_topk_weights_early, monkeypatch): if not _use_cudnn_cutedsl_fusion_from_env(): @@ -913,12 +956,25 @@ def checked_glu(*args, **kwargs): # The dedicated SwiGLU path is the reference for this kernel # substitution. Its MXFP8 gradients can differ from pure JAX by # more than the strict unfused test threshold on Rubin. - monkeypatch.setattr(tex, "grouped_gemm_glu", tex.grouped_gemm_swiglu) + import cudnn.jax as cudnn_jax + + # Simulate a frontend without the Rubin API. Selection must fall + # back to generic SwiGLU even though the GPU itself is Rubin. + monkeypatch.setattr(cudnn_jax, "grouped_gemm_glu", None) + generic_calls = [] + original_swiglu = tex.grouped_gemm_swiglu + + def checked_swiglu(*args, **kwargs): + generic_calls.append(True) + return original_swiglu(*args, **kwargs) + + monkeypatch.setattr(tex, "grouped_gemm_swiglu", checked_swiglu) with _ctx(mesh): x_sh = _shard_inputs(x, mesh) baseline_output, _, _ = jax.jit(block.apply)(variables, x_sh) baseline_output.block_until_ready() baseline_grads, baseline_grad_x = _grad_step(block, variables, mesh, x) + assert generic_calls, "Missing Rubin API did not select generic SwiGLU" np.testing.assert_allclose( output_np, _to_global_numpy(baseline_output, mesh).astype(np.float32), @@ -997,7 +1053,9 @@ def test_cudnn_fused_with_checkpoint_names(self, mesh, monkeypatch, use_regular_ moe_module = importlib.import_module("transformer_engine.jax.moe") flax_moe_module = importlib.import_module("transformer_engine.jax.flax.moe") if use_regular_swiglu: - monkeypatch.setattr(moe_module, "_is_rubin_device", lambda: False) + import cudnn.jax as cudnn_jax + + monkeypatch.setattr(cudnn_jax, "grouped_gemm_glu", None) selected_calls = [] fused_op_name = "grouped_gemm_swiglu" if use_regular_swiglu else "grouped_gemm_glu" diff --git a/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py b/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py index 95516e5df3d..671381172db 100644 --- a/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py +++ b/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py @@ -1,10 +1,13 @@ -# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # # See LICENSE for license information. """cuDNN Frontend JAX API adapter for fused grouped GEMM + SwiGLU.""" from __future__ import annotations +import importlib +import inspect + import jax import jax.numpy as jnp @@ -70,21 +73,53 @@ def _compact_sf(scale: jax.Array, shape: tuple[int, ...], name: str) -> jax.Arra return scale.reshape(-1)[:size].reshape(shape) -def grouped_gemm_swiglu_dependencies_available(rubin: bool = False) -> tuple[bool, str]: - """Check the public cuDNN JAX API without compiling a kernel.""" - try: - import cutlass.jax - from cudnn.jax import ( # noqa: F401 - grouped_gemm_dswiglu, - grouped_gemm_swiglu, - ) +# These contracts match cuDNN Frontend 1.31.0. Probe APIs rather than the version: +# newer optional parameters are compatible, but TE must supply every required one. +_FORWARD_ARGS = ( + "a_tensor", + "b_tensor", + "sfa_tensor", + "sfb_tensor", + "padded_offsets", + "alpha_tensor", + "prob_tensor", + "norm_const_tensor", + "c_dtype", + "d_dtype", + "discrete_col_sfd", +) +_BACKWARD_ARGS = tuple(arg for arg in _FORWARD_ARGS if arg != "c_dtype") + ( + "c_tensor", + "beta_tensor", +) - if rubin: - from cudnn.jax import grouped_gemm_glu # noqa: F401 - if not cutlass.jax.is_available(): +def grouped_gemm_swiglu_dependencies_available(rubin: bool = False) -> tuple[bool, str]: + """Check that the selected forward and shared backward accept TE's keywords.""" + try: + cutlass_jax = importlib.import_module("cutlass.jax") + cudnn_jax = importlib.import_module("cudnn.jax") + if not cutlass_jax.is_available(): return False, "CuTeDSL JAX support is unavailable" - except (ImportError, ModuleNotFoundError, RuntimeError, AttributeError) as exc: + forward = "grouped_gemm_glu" if rubin else "grouped_gemm_swiglu" + for name, arguments in ( + (forward, _FORWARD_ARGS), + ("grouped_gemm_dswiglu", _BACKWARD_ARGS), + ): + api = getattr(cudnn_jax, name, None) + if not callable(api): + return False, f"cudnn.jax.{name} is unavailable" + signature = inspect.signature(api) + missing = set(arguments) - signature.parameters.keys() + if missing: + return ( + False, + f"cudnn.jax.{name} is missing parameters: {', '.join(sorted(missing))}", + ) + # Binding detects missing/renamed keywords, positional-only arguments, + # and new mandatory arguments, while allowing new optional arguments. + signature.bind(**dict.fromkeys(arguments)) + except (ImportError, RuntimeError, AttributeError, TypeError, ValueError) as exc: return False, str(exc) return True, "" @@ -111,7 +146,9 @@ def grouped_gemm_swiglu( if b.ndim != 3: raise ValueError(f"Expected physical B[E,N,K], got {b.shape}") - from cudnn.jax import grouped_gemm_swiglu as cudnn_grouped_gemm_swiglu + cudnn_grouped_gemm_swiglu = getattr( + importlib.import_module("cudnn.jax"), "grouped_gemm_swiglu" + ) return _grouped_gemm_forward( cudnn_grouped_gemm_swiglu, @@ -138,7 +175,9 @@ def grouped_gemm_glu( output_dtype, ) -> tuple[jax.Array, jax.Array, jax.Array, jax.Array, jax.Array]: """Run the Rubin cuDNN grouped MXFP8 GEMM + SwiGLU kernel.""" - from cudnn.jax import grouped_gemm_glu as cudnn_grouped_gemm_glu + cudnn_grouped_gemm_glu = getattr( + importlib.import_module("cudnn.jax"), "grouped_gemm_glu" + ) return _grouped_gemm_forward( cudnn_grouped_gemm_glu, @@ -212,7 +251,9 @@ def grouped_gemm_dswiglu( if c.ndim != 2: raise ValueError(f"Expected C[M,2N], got {c.shape}") - from cudnn.jax import grouped_gemm_dswiglu as cudnn_grouped_gemm_dswiglu + cudnn_grouped_gemm_dswiglu = getattr( + importlib.import_module("cudnn.jax"), "grouped_gemm_dswiglu" + ) rows, hidden = a.shape experts, intermediate, b_hidden = b.shape diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index 41389c611e6..9dee05f0e48 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -254,12 +254,46 @@ def _use_cudnn_cutedsl_fusion_from_env() -> bool: return value == "1" -def _is_rubin_device() -> bool: - """SM107 is reported as 107 by TransformerEngine's CUDA utility.""" +def _select_cudnn_jax_fusion(rejection_reasons: list[str]) -> str | bool: + """Select Rubin GLU, generic Blackwell+ SwiGLU, or unfused TE, in order.""" from transformer_engine_jax import get_device_compute_capability - return get_device_compute_capability(0) == 107 - + reasons = list(rejection_reasons) + path = False + try: + capability = get_device_compute_capability(0) + except RuntimeError as exc: + reasons.append(f"could not query GPU compute capability: {exc}") + else: + if not reasons: + for candidate, supported in ( + ("rubin", capability == 107), + ("blackwell", capability >= 100), + ): + if not supported: + requirement = "SM107" if candidate == "rubin" else "SM100+" + reasons.append( + f"{candidate} fused kernel requires {requirement}, got SM{capability}" + ) + continue + available, error = tex.grouped_gemm_swiglu_dependencies_available( + rubin=candidate == "rubin" + ) + if available: + path = candidate + break + reasons.append(f"{candidate} fused API is incompatible: {error}") + if reasons: + destination = "generic Blackwell+ fused kernel" if path else "unfused TE grouped-GEMM path" + warnings.warn( + f"{_CUDNN_JAX_ENV}=1: falling back to the {destination}: " + + "; ".join(reasons) + + ". Install cuDNN Frontend with compatible JAX APIs (TE's signatures match " + "cuDNN Frontend 1.31.0) and CuTeDSL JAX support; use supported GPU hardware.", + UserWarning, + stacklevel=2, + ) + return path def _cudnn_jax_fusion_rejection_reasons( x, @@ -273,17 +307,7 @@ def _cudnn_jax_fusion_rejection_reasons( ep_axis, ) -> list[str]: """Return reasons this call cannot use cuDNN's grouped SwiGLU JAX API.""" - from transformer_engine_jax import get_device_compute_capability - errors = [] - compute_capability = None - try: - compute_capability = get_device_compute_capability(0) - except RuntimeError as exc: - errors.append(f"could not query GPU compute capability: {exc}") - else: - if compute_capability < 100: - errors.append(f"requires an SM100+ GPU, got SM{compute_capability}") if str(activation_type).lower() != "silu": errors.append("requires activation_type='silu'") if wi_0_bias is not None or wi_1_bias is not None: @@ -330,11 +354,6 @@ def _cudnn_jax_fusion_rejection_reasons( if num_local_experts > 1024: errors.append(f"requires at most 1024 local experts, got {num_local_experts}") - dependencies_available, dependency_error = tex.grouped_gemm_swiglu_dependencies_available( - rubin=compute_capability == 107 - ) - if not dependencies_available: - errors.append(f"could not load cuDNN's grouped SwiGLU JAX API: {dependency_error}") return errors @@ -768,7 +787,7 @@ def _ffn_fwd_per_shard( num_local_experts: int, activation_type: str, apply_topk_weights_early: bool, - use_cudnn_jax_fusion: bool, + use_cudnn_jax_fusion: str | bool, wi_0_checkpoint_name: Optional[str], wi_1_checkpoint_name: Optional[str], wo_checkpoint_name: Optional[str], @@ -842,7 +861,7 @@ def _ffn_fwd_per_shard( intermediate_col, intermediate_scale_row, intermediate_scale_col, - ) = (tex.grouped_gemm_glu if _is_rubin_device() else tex.grouped_gemm_swiglu)( + ) = (tex.grouped_gemm_glu if use_cudnn_jax_fusion == "rubin" else tex.grouped_gemm_swiglu)( casted_sorted_x_lhs.data.reshape(sorted_x.shape[0], hidden, 1), ( casted_wi_rhs.data.reshape(num_local_experts, combined, hidden) @@ -994,7 +1013,7 @@ def _ffn_bwd_per_shard( activation_type: str, apply_topk_weights_early: bool, has_bias: bool, - use_cudnn_jax_fusion: bool, + use_cudnn_jax_fusion: str | bool, cudnn_native_weight_layout: bool, ): """Backward mirror of :func:`_ffn_fwd_per_shard`.""" @@ -1123,9 +1142,9 @@ def _ffn_bwd_per_shard( d_sorted_x = tex.grouped_gemm( casted_d_combined.get_tensor(usage=TensorUsage.LHS), casted_wi_rhs_trans, - contracting_dims=((1,), (1 if cudnn_native_weight_layout else 2,)), + contracting_dims=((1,), (1 if use_cudnn_jax_fusion and cudnn_native_weight_layout else 2,)), ) - if cudnn_native_weight_layout: + if use_cudnn_jax_fusion and cudnn_native_weight_layout: # dY^T @ X directly produces [E,2N,K], matching the persistent native # parameter, without a post-GEMM transpose or de-interleave/repack. d_wi_combined = tex.grouped_gemm( @@ -1276,7 +1295,13 @@ def _moe_fwd_rule( # Per-rank send capacity: B/num_procs rows x S tokens per rank. max_tokens_per_rank = (B // num_procs) * S - dispatch_alignment = _CUDNN_JAX_ALIGN_SIZE if use_cudnn_jax_fusion else _ALIGN_SIZE + # Keep capacity and alignment consistent with an EP bootstrap sized + # for requested fusion, even when API/hardware checks choose unfused. + dispatch_alignment = ( + _CUDNN_JAX_ALIGN_SIZE + if use_cudnn_jax_fusion or _use_cudnn_cutedsl_fusion_from_env() + else _ALIGN_SIZE + ) worst_case_recv_pr = get_moe_recv_capacity_per_rank( num_experts=num_experts, num_experts_per_tok=K, @@ -2033,14 +2058,15 @@ def moe( ep_axis, data_parallelism_axes, weight_gather : deprecated Compatibility arguments converted into a MeshResource and boolean, with a DeprecationWarning. Conflicting old and new arguments raise. - Note that the per-expert dispatch-slot alignment is fixed internally - at 128 tokens (``_ALIGN_SIZE``); see that constant's docstring for - rationale and how to extend if a future recipe needs >128. + Per-expert dispatch-slot alignment defaults to 128 tokens (``_ALIGN_SIZE``). + Requesting cuDNN fusion reserves 256 tokens, also when falling back, to + preserve compatibility with EP bootstrap buffer sizing. Set ``NVTE_JAX_TEMP_FLAG_FOR_ABHINAV_CUDNN_GROUPED_GEMM_FUSION=1`` to use cuDNN's - dedicated JAX grouped MXFP8 GEMM + SwiGLU API for eligible SM100 calls. - The fused path uses 256-token expert alignment. Ineligible calls warn - and fall back to TE's regular grouped-GEMM implementation. + JAX grouped MXFP8 APIs: Rubin GLU first, then generic SM100+ SwiGLU. + Ineligible calls warn + and fall back to TE's regular grouped-GEMM implementation. API signatures + and GPU capability determine support; fallbacks emit an actionable warning. MeshResource fields name physical mesh axes, not Flax logical axes. ``input_axes``, ``gate_kernel_axes``, ``wi_kernel_axes`` and @@ -2114,22 +2140,7 @@ def moe( activation_type=activation_type, ep_axis=ep_axis, ) - if rejection_reasons: - if cudnn_native_weight_layout: - raise ValueError( - "cuDNN-native MoE weight layout requires the fused cuDNN grouped-GEMM " - "path, which is unsupported for this moe() call: " - + "; ".join(rejection_reasons) - ) - warnings.warn( - f"{_CUDNN_JAX_ENV}=1 is unsupported for this moe() call; falling back to " - "the regular TE grouped-GEMM path: " - + "; ".join(rejection_reasons), - UserWarning, - stacklevel=2, - ) - else: - use_cudnn_jax_fusion = True + use_cudnn_jax_fusion = _select_cudnn_jax_fusion(rejection_reasons) elif cudnn_native_weight_layout: raise ValueError( "cuDNN-native MoE weight layout requires the fused cuDNN grouped-GEMM path; " From ae3ec09cf1f158768f0174fec112e0d4ac5ba3a7 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Tue, 6 Oct 2026 11:51:00 -0700 Subject: [PATCH 24/38] Replace the temporary MoE fusion flag with an explicit boolean Signed-off-by: Jeremy Berchtold --- tests/jax/conftest.py | 7 +- tests/jax/run_te_ep_moe.sh | 24 +++---- .../jax/test_grouped_gemm_swiglu_fallback.py | 71 +++++++++++++++++-- tests/jax/test_te_ep_moe.py | 69 +++++++++++++++--- transformer_engine/jax/flax/moe.py | 13 ++-- transformer_engine/jax/moe.py | 55 +++++++------- 6 files changed, 176 insertions(+), 63 deletions(-) diff --git a/tests/jax/conftest.py b/tests/jax/conftest.py index d729bfd1c7b..c7cebc7fdb1 100644 --- a/tests/jax/conftest.py +++ b/tests/jax/conftest.py @@ -98,7 +98,12 @@ def pytest_addoption(parser): """ parser.addoption("--num-process", action="store", default=0) parser.addoption("--process-id", action="store", default=0) - + parser.addoption( + "--use-cudnn-fusion", + choices=("0", "1"), + default="1", + help="Pass use_cudnn_fusion to MoE blocks and EP capacity calculation", + ) def pytest_configure(config): config.addinivalue_line( diff --git a/tests/jax/run_te_ep_moe.sh b/tests/jax/run_te_ep_moe.sh index 2445cee359e..c8233605ec9 100755 --- a/tests/jax/run_te_ep_moe.sh +++ b/tests/jax/run_te_ep_moe.sh @@ -67,7 +67,7 @@ trap cleanup EXIT INT TERM run_phase() { local phase_name="$1" - local fusion_env="$2" + local use_cudnn_fusion="$2" shift 2 local -a phase_args=("$@") local phase_log_dir="$LOG_DIR/$phase_name" @@ -78,7 +78,7 @@ run_phase() { echo echo "============================================================" echo "Phase: $phase_name" - echo " NVTE_JAX_TEMP_FLAG_FOR_ABHINAV_CUDNN_GROUPED_GEMM_FUSION=$fusion_env" + echo " use_cudnn_fusion=$use_cudnn_fusion" echo " phase pytest args : ${phase_args[*]:-}" echo " logs : $phase_log_dir" echo "============================================================" @@ -92,15 +92,14 @@ run_phase() { -v -s --num-process="$NUM_GPUS" --process-id="$i" + --use-cudnn-fusion="$use_cudnn_fusion" "${phase_args[@]}" ) if [ "$i" -eq 0 ]; then echo "=== Live output from process 0 ($phase_name) ===" - env NVTE_JAX_TEMP_FLAG_FOR_ABHINAV_CUDNN_GROUPED_GEMM_FUSION="$fusion_env" \ - "${pytest_cmd[@]}" 2>&1 | tee "$log_file" & + "${pytest_cmd[@]}" 2>&1 | tee "$log_file" & else - env NVTE_JAX_TEMP_FLAG_FOR_ABHINAV_CUDNN_GROUPED_GEMM_FUSION="$fusion_env" \ - "${pytest_cmd[@]}" > "$log_file" 2>&1 & + "${pytest_cmd[@]}" > "$log_file" 2>&1 & fi PIDS+=("$!") done @@ -136,15 +135,10 @@ run_phase() { fi } -if [ "${NVTE_JAX_TEMP_FLAG_FOR_ABHINAV_CUDNN_GROUPED_GEMM_FUSION:-0}" = "1" ]; then - # Keep ordinary CUDA C++ and cuDNN JAX coverage in separate Python process - # groups. TE EP/NCCL caches layer alignment process-wide, so 128-token and - # 256-token dispatch-alignment tests cannot safely share one interpreter. - run_phase "ordinary" "0" -k "not TestTeEpMoeCudnnCutedslFusion" "$@" - run_phase "cutedsl" "1" -k "TestTeEpMoeCudnnCutedslFusion" "$@" -else - run_phase "ordinary" "0" "$@" -fi +# Keep ordinary CUDA C++ and cuDNN JAX coverage in separate process groups. +# TE EP/NCCL caches alignment process-wide; these use 128 and 256 respectively. +run_phase "ordinary" "0" -k "not TestTeEpMoeCudnnCutedslFusion" "$@" +run_phase "cutedsl" "1" -k "TestTeEpMoeCudnnCutedslFusion" "$@" echo if [ "$PHASE_FAILED" -eq 0 ]; then diff --git a/tests/jax/test_grouped_gemm_swiglu_fallback.py b/tests/jax/test_grouped_gemm_swiglu_fallback.py index 996079c9f7c..edabdbd84c3 100644 --- a/tests/jax/test_grouped_gemm_swiglu_fallback.py +++ b/tests/jax/test_grouped_gemm_swiglu_fallback.py @@ -161,12 +161,15 @@ def test_installed_frontend_contract(): assert available or reason +@pytest.mark.parametrize("request_fusion", [None, True, False]) @pytest.mark.parametrize("native_layout", [False, True]) @pytest.mark.parametrize( "missing,expected", [(None, "rubin"), ("grouped_gemm_glu", "blackwell"), ("grouped_gemm_dswiglu", False)], ) -def test_public_moe_passes_selected_path(frontend, monkeypatch, native_layout, missing, expected): +def test_public_moe_passes_selected_path( + frontend, monkeypatch, native_layout, missing, expected, request_fusion +): import jax import jax.numpy as jnp import numpy as np @@ -175,7 +178,6 @@ def test_public_moe_passes_selected_path(frontend, monkeypatch, native_layout, m if missing: delattr(frontend, missing) - monkeypatch.setenv(moe._CUDNN_JAX_ENV, "1") monkeypatch.setattr(transformer_engine_jax, "get_device_compute_capability", lambda _: 107) monkeypatch.setattr(moe, "_cudnn_jax_fusion_rejection_reasons", lambda *args, **kwargs: []) mesh = Mesh(np.asarray(jax.devices()[:1]), ("ep",)) @@ -185,12 +187,17 @@ def test_public_moe_passes_selected_path(frontend, monkeypatch, native_layout, m signature = inspect.signature(moe._moe) def execute(*args): - received.append(signature.bind(*args).arguments["use_cudnn_jax_fusion"]) + bound = signature.bind(*args).arguments + assert bound["use_cudnn_fusion"] is (request_fusion is not False) + received.append(bound["use_cudnn_jax_fusion"]) return args[0], None, jnp.asarray(0) monkeypatch.setattr(moe, "_moe", execute) x = jnp.ones((1, 1, 128), jnp.bfloat16) wi = jnp.ones((2, 256, 128) if native_layout else (2, 128, 256), jnp.bfloat16) + kwargs = {} if request_fusion is None else {"use_cudnn_fusion": request_fusion} + if request_fusion is False: + expected = False with warnings.catch_warnings(record=True) as caught: warnings.simplefilter("always") output, _, _ = moe.moe( @@ -201,10 +208,11 @@ def execute(*args): num_experts=2, num_experts_per_tok=1, mesh_resource=MeshResource(ep_resource="ep"), + **kwargs, ) assert received == [expected] assert output is x - assert len(caught) == (0 if expected == "rubin" else 1) + assert len(caught) == (0 if request_fusion is False or expected == "rubin" else 1) @pytest.mark.parametrize("path", ["rubin", "blackwell"]) @@ -249,3 +257,58 @@ def quantize(data, *args, **kwargs): ) with pytest.raises(SelectedKernel): moe._ffn_fwd_per_shard(**kwargs) + + +@pytest.mark.parametrize("request_fusion", [None, True, False]) +def test_flax_forwards_fusion_bool(monkeypatch, request_fusion): + import jax + import jax.numpy as jnp + from transformer_engine.jax.flax import _MoEBlock + from transformer_engine.jax.sharding import MeshResource + + flax_moe = importlib.import_module("transformer_engine.jax.flax.moe") + received = [] + + def execute(inputs, *args, **kwargs): + received.append(kwargs["use_cudnn_fusion"]) + return inputs, None, jnp.asarray(0) + + monkeypatch.setattr(flax_moe, "moe", execute) + kwargs = {} if request_fusion is None else {"use_cudnn_fusion": request_fusion} + block = _MoEBlock( + num_experts=2, intermediate_size=32, mesh_resource=MeshResource(ep_resource="ep"), **kwargs + ) + block.init(jax.random.PRNGKey(0), jnp.ones((1, 1, 32))) + assert received == [request_fusion is not False] + + +@pytest.mark.parametrize("request_fusion", [True, False]) +def test_capacity_follows_explicit_fusion_bool(request_fusion): + kwargs = dict(num_experts=8, num_experts_per_tok=2, max_tokens_per_rank=64, ep_size=2) + alignment = moe._CUDNN_JAX_ALIGN_SIZE if request_fusion else moe._ALIGN_SIZE + expected = moe.get_moe_recv_capacity_per_rank(**kwargs, alignment=alignment) + assert moe.get_moe_recv_capacity_per_rank(**kwargs, use_cudnn_fusion=request_fusion) == expected + if request_fusion: + assert moe.get_moe_recv_capacity_per_rank(**kwargs) == expected + + +@pytest.mark.parametrize("invalid", [0, 1, None, "true"]) +def test_fusion_argument_requires_bool(invalid): + with pytest.raises(TypeError, match="use_cudnn_fusion must be a bool"): + moe.moe( + None, None, None, None, num_experts=2, num_experts_per_tok=1, use_cudnn_fusion=invalid + ) + with pytest.raises(TypeError, match="use_cudnn_fusion must be a bool"): + moe.get_moe_recv_capacity_per_rank( + num_experts=2, + num_experts_per_tok=1, + max_tokens_per_rank=16, + ep_size=1, + use_cudnn_fusion=invalid, + ) + + +def test_vjp_bool_is_static_and_defaults_to_true(): + assert inspect.signature(moe._moe).parameters["use_cudnn_fusion"].default is True + assert inspect.signature(moe.moe).parameters["use_cudnn_fusion"].default is True + assert 32 in moe._moe.nondiff_argnums diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index 8e1c44e3332..65fd68e22f9 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -99,6 +99,16 @@ def _read_mp_options(): return num, pid +def _read_cudnn_fusion_option() -> bool: + for index, argument in enumerate(sys.argv): + if argument.startswith("--use-cudnn-fusion="): + return argument.split("=", 1)[1] == "1" + if argument == "--use-cudnn-fusion" and index + 1 < len(sys.argv): + return sys.argv[index + 1] == "1" + return True + + +_USE_CUDNN_FUSION = _read_cudnn_fusion_option() _MP_NUM_PROCESS, _MP_PROCESS_ID = _read_mp_options() _MP_ACTIVE = _init_distributed(_MP_NUM_PROCESS, _MP_PROCESS_ID) @@ -123,7 +133,6 @@ def _read_mp_options(): from transformer_engine.jax.moe import ( _ALIGN_SIZE, _CUDNN_JAX_ALIGN_SIZE, - _use_cudnn_cutedsl_fusion_from_env, WeightGather, get_moe_recv_capacity_per_rank, moe, @@ -209,7 +218,7 @@ def mesh(): # Worst-case recv capacity per rank # TODO(jberchtold) support configurations other than worst-case by refactoring tests # but if possible avoid bootstrap/teardown for each test - alignment = _CUDNN_JAX_ALIGN_SIZE if _use_cudnn_cutedsl_fusion_from_env() else _ALIGN_SIZE + alignment = _CUDNN_JAX_ALIGN_SIZE if _USE_CUDNN_FUSION else _ALIGN_SIZE recv_capacity_per_rank = get_moe_recv_capacity_per_rank( num_experts=NUM_EXPERTS, num_experts_per_tok=TOPK, @@ -368,6 +377,7 @@ def _make_block( dispatch_checkpoint_name=None, quant_before_fsdp_ag=False, mesh_resource=None, + use_cudnn_fusion=_USE_CUDNN_FUSION, ): kwargs = dict( num_experts=NUM_EXPERTS, @@ -383,6 +393,7 @@ def _make_block( quantization_recipe=quantization_recipe, dispatch_checkpoint_name=dispatch_checkpoint_name, quant_before_fsdp_ag=quant_before_fsdp_ag, + use_cudnn_fusion=use_cudnn_fusion, ) # Custom expert_bias_init lets tests inject a non-zero expert_bias without # poking variables['params'] post-init. @@ -593,7 +604,7 @@ def test_quantized_weight_gather_matches_full_precision_gather( ): """The FP8 weight gather retains forward and backward MoE semantics.""" if native_weight_layout: - if not _use_cudnn_cutedsl_fusion_from_env(): + if not _USE_CUDNN_FUSION: pytest.skip("Native weight layout requires cuDNN grouped GEMM fusion") if native_weight_layout or expert_fsdp: from transformer_engine.jax import cpp_extensions as tex @@ -838,7 +849,7 @@ def loss_fn(params, x): def test_ep_checkpoint_names(mesh, monkeypatch): - if _use_cudnn_cutedsl_fusion_from_env(): + if _USE_CUDNN_FUSION: pytest.skip("BF16 fallback uses a different EP alignment than the cuDNN bootstrap") moe_module = importlib.import_module("transformer_engine.jax.moe") named_values = {} @@ -864,13 +875,55 @@ def recorded_checkpoint_name(value, name): assert np.all(np.isfinite(_to_global_numpy(_unwrap(grads["params"][name]), mesh))) +def test_explicitly_disabled_fusion_with_native_weights(mesh, monkeypatch): + """Disabling fusion bypasses dependency probing and preserves native gradients.""" + if _USE_CUDNN_FUSION: + pytest.skip("Requires the ordinary 128-token EP bootstrap") + from transformer_engine.jax import cpp_extensions as tex + + moe_module = importlib.import_module("transformer_engine.jax.moe") + + def unexpected_selection(*args, **kwargs): + raise AssertionError("Explicitly disabled fusion must not probe cuDNN") + + monkeypatch.setattr(moe_module, "_select_cudnn_jax_fusion", unexpected_selection) + block = _make_block(quantization_recipe=MXFP8BlockScaling(), use_cudnn_fusion=False) + x = _make_inputs(jax.random.PRNGKey(53)) + variables, baseline_output, _ = _init_apply(block, mesh, x, jax.random.PRNGKey(54)) + baseline_grads, baseline_dx = _grad_step(block, variables, mesh, x) + flax_moe = importlib.import_module("transformer_engine.jax.flax.moe") + original_moe = flax_moe.moe + + def native_moe(*args, **kwargs): + args = list(args) + gate, up = jnp.split(args[2], 2, axis=-1) + args[2] = tex.pack_swiglu_pair(gate, up).transpose(0, 2, 1) + kwargs["wi_kernel_axes"] = ("exp", "mlp", "embed") + return original_moe(*args, **kwargs) + + monkeypatch.setattr(flax_moe, "moe", native_moe) + with _ctx(mesh): + output, _, _ = jax.jit(block.apply)(variables, _shard_inputs(x, mesh)) + output.block_until_ready() + grads, dx = _grad_step(block, variables, mesh, x) + np.testing.assert_array_equal( + _to_global_numpy(output, mesh), _to_global_numpy(baseline_output, mesh) + ) + np.testing.assert_array_equal(_to_global_numpy(dx, mesh), _to_global_numpy(baseline_dx, mesh)) + for name in ("gate_kernel", "wi", "wo"): + np.testing.assert_array_equal( + _to_global_numpy(_unwrap(grads["params"][name]), mesh), + _to_global_numpy(_unwrap(baseline_grads["params"][name]), mesh), + ) + + class TestTeEpMoeCudnnCutedslFusion: """End-to-end MXFP8 coverage for cuDNN's grouped GLU JAX APIs.""" @pytest.mark.parametrize("native_layout", [False, True]) def test_missing_apis_unfused_forward_and_backward(self, mesh, monkeypatch, native_layout): """Unfused fallback retains bootstrap capacity and native-layout gradients.""" - if not _use_cudnn_cutedsl_fusion_from_env(): + if not _USE_CUDNN_FUSION: pytest.skip("Requires fusion requested at bootstrap") import cudnn.jax as cudnn_jax from transformer_engine.jax import cpp_extensions as tex @@ -912,9 +965,9 @@ def native_moe(*args, **kwargs): @pytest.mark.parametrize("apply_topk_weights_early", [False, True]) def test_mxfp8_forward_and_backward(self, mesh, apply_topk_weights_early, monkeypatch): - if not _use_cudnn_cutedsl_fusion_from_env(): + if not _USE_CUDNN_FUSION: pytest.skip( - "run separately with NVTE_JAX_TEMP_FLAG_FOR_ABHINAV_CUDNN_GROUPED_GEMM_FUSION=1" + "run separately with --use-cudnn-fusion=1" ) rubin_calls = [] if get_device_compute_capability(0) == 107: @@ -1043,7 +1096,7 @@ def loss_fn(params, inputs): @pytest.mark.parametrize("use_regular_swiglu", [False, True]) def test_cudnn_fused_with_checkpoint_names(self, mesh, monkeypatch, use_regular_swiglu): - if not _use_cudnn_cutedsl_fusion_from_env(): + if not _USE_CUDNN_FUSION: pytest.skip("cuDNN grouped GEMM fusion is disabled") if not use_regular_swiglu and get_device_compute_capability(0) != 107: pytest.skip("Rubin grouped GLU requires SM107") diff --git a/transformer_engine/jax/flax/moe.py b/transformer_engine/jax/flax/moe.py index 429cce08e24..cf9280fa54d 100644 --- a/transformer_engine/jax/flax/moe.py +++ b/transformer_engine/jax/flax/moe.py @@ -103,6 +103,11 @@ class _MoEBlock(TransformerEngineBase): ep_axis, data_parallelism_axes, weight_gather : deprecated Compatibility arguments converted into MeshResource and the boolean with a DeprecationWarning. + use_cudnn_fusion : bool + Defaults to ``True``: try Rubin fused GLU, then generic Blackwell+ fused + SwiGLU, then unfused TE grouped GEMM. Unsupported fused paths warn. + ``False`` selects unfused execution. Use the same value when calculating + EP bootstrap receive capacity with ``get_moe_recv_capacity_per_rank``. apply_topk_weights_early : bool If ``True``, multiply expert outputs by their top-k weights *inside* each shard before ``ep_combine`` (saves one global @@ -114,10 +119,8 @@ class _MoEBlock(TransformerEngineBase): JAX rematerialization checkpoint name for the EP dispatch outputs. ``None`` leaves them unnamed. - The per-expert dispatch-slot alignment is fixed internally at 128 - tokens (see ``moe._ALIGN_SIZE``) -- the value required by NCCL EP - HT and satisfied by every current TE grouped-GEMM recipe -- and is - therefore not exposed as a per-instance knob. + Per-expert dispatch-slot alignment is 256 tokens when fusion is requested, + including fallbacks, and 128 tokens with ``use_cudnn_fusion=False``. dtype : jnp.dtype Compute / parameter dtype. @@ -161,6 +164,7 @@ class _MoEBlock(TransformerEngineBase): data_parallelism_axes: Optional[Tuple[str, ...]] = None # MoE knobs forwarded to ``moe()`` + use_cudnn_fusion: bool = True apply_topk_weights_early: bool = False recv_capacity_per_rank: Optional[int] = None dispatch_checkpoint_name: Optional[str] = None @@ -324,6 +328,7 @@ def make_grouped_quantizer_set(postfix): group_topk=self.group_topk, scaling_factor=self.scaling_factor, aux_loss_coeff=self.aux_loss_coeff, + use_cudnn_fusion=self.use_cudnn_fusion, apply_topk_weights_early=self.apply_topk_weights_early, quantizer_sets=quantizer_sets, recv_capacity_per_rank=self.recv_capacity_per_rank, diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index 9dee05f0e48..e56355967ff 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -31,7 +31,6 @@ """ import math -import os import warnings from dataclasses import dataclass, fields, replace from functools import partial @@ -244,14 +243,6 @@ def _resolve_moe_mesh_resource( # same 128-token tile, so a single constant covers every supported path. _ALIGN_SIZE = 128 _CUDNN_JAX_ALIGN_SIZE = 256 -_CUDNN_JAX_ENV = "NVTE_JAX_TEMP_FLAG_FOR_ABHINAV_CUDNN_GROUPED_GEMM_FUSION" - - -def _use_cudnn_cutedsl_fusion_from_env() -> bool: - value = os.getenv(_CUDNN_JAX_ENV, "0") - if value not in ("0", "1"): - raise ValueError(f"{_CUDNN_JAX_ENV} must be '0' or '1', got {value!r}") - return value == "1" def _select_cudnn_jax_fusion(rejection_reasons: list[str]) -> str | bool: @@ -286,7 +277,7 @@ def _select_cudnn_jax_fusion(rejection_reasons: list[str]) -> str | bool: if reasons: destination = "generic Blackwell+ fused kernel" if path else "unfused TE grouped-GEMM path" warnings.warn( - f"{_CUDNN_JAX_ENV}=1: falling back to the {destination}: " + f"use_cudnn_fusion=True: falling back to the {destination}: " + "; ".join(reasons) + ". Install cuDNN Frontend with compatible JAX APIs (TE's signatures match " "cuDNN Frontend 1.31.0) and CuTeDSL JAX support; use supported GPU hardware.", @@ -365,6 +356,7 @@ def get_moe_recv_capacity_per_rank( ep_size: int, recv_capacity_factor: Optional[float] = None, alignment: Optional[int] = None, + use_cudnn_fusion: bool = True, ) -> int: """Return the aligned receive capacity for one EP rank. @@ -372,12 +364,14 @@ def get_moe_recv_capacity_per_rank( factor >= 1 scales the capacity needed by perfectly balanced routing and is capped at the worst case. The balanced baseline includes the independent per-local-expert alignment required by NCCL EP. When ``alignment`` is not - supplied, it follows the active MoE implementation: 256 for the cuDNN - grouped-SwiGLU fusion and 128 for the regular TE grouped GEMM. This keeps + supplied, ``use_cudnn_fusion=True`` (default) reserves 256-token alignment, + including fallbacks; ``False`` reserves 128 for regular TE grouped GEMM. This keeps eager bootstrap callers in sync with the later compiled ``moe()`` call. """ + if not isinstance(use_cudnn_fusion, bool): + raise TypeError("use_cudnn_fusion must be a bool") if alignment is None: - alignment = _CUDNN_JAX_ALIGN_SIZE if _use_cudnn_cutedsl_fusion_from_env() else _ALIGN_SIZE + alignment = _CUDNN_JAX_ALIGN_SIZE if use_cudnn_fusion else _ALIGN_SIZE if num_experts <= 0 or num_experts_per_tok <= 0 or max_tokens_per_rank <= 0: raise ValueError( "num_experts, num_experts_per_tok, and max_tokens_per_rank must be positive" @@ -1222,6 +1216,7 @@ def _moe_fwd_rule( wo_checkpoint_name, dispatch_checkpoint_name, quant_before_fsdp_ag, + use_cudnn_fusion: bool = True, ): """Forward: gate -> topk -> ep_dispatch -> FFN -> ep_combine. @@ -1297,11 +1292,7 @@ def _moe_fwd_rule( max_tokens_per_rank = (B // num_procs) * S # Keep capacity and alignment consistent with an EP bootstrap sized # for requested fusion, even when API/hardware checks choose unfused. - dispatch_alignment = ( - _CUDNN_JAX_ALIGN_SIZE - if use_cudnn_jax_fusion or _use_cudnn_cutedsl_fusion_from_env() - else _ALIGN_SIZE - ) + dispatch_alignment = _CUDNN_JAX_ALIGN_SIZE if use_cudnn_fusion else _ALIGN_SIZE worst_case_recv_pr = get_moe_recv_capacity_per_rank( num_experts=num_experts, num_experts_per_tok=K, @@ -1615,6 +1606,7 @@ def _moe_bwd_rule( wo_checkpoint_name, dispatch_checkpoint_name, quant_before_fsdp_ag, + use_cudnn_fusion, residuals, cotangents, ): @@ -1631,6 +1623,7 @@ def _moe_bwd_rule( wo_checkpoint_name, dispatch_checkpoint_name, quant_before_fsdp_ag, + use_cudnn_fusion, ) # captured / unused in bwd from jax.experimental.shard_map import shard_map @@ -1880,7 +1873,7 @@ def _ffn_bwd_body(*args): # ============================================================================= -@partial(jax.custom_vjp, nondiff_argnums=tuple(range(9, 32))) +@partial(jax.custom_vjp, nondiff_argnums=tuple(range(9, 33))) def _moe( x, gate_kernel, @@ -1914,6 +1907,7 @@ def _moe( wo_checkpoint_name, dispatch_checkpoint_name, quant_before_fsdp_ag, + use_cudnn_fusion: bool = True, ): primal, _ = _moe_fwd_rule( x, @@ -1948,6 +1942,7 @@ def _moe( wo_checkpoint_name, dispatch_checkpoint_name, quant_before_fsdp_ag, + use_cudnn_fusion, ) return primal @@ -1994,6 +1989,7 @@ def moe( wo_checkpoint_name: Optional[str] = None, dispatch_checkpoint_name: Optional[str] = None, weight_gather: Optional[WeightGather] = None, + use_cudnn_fusion: bool = True, ) -> Tuple[jnp.ndarray, Optional[jnp.ndarray], jnp.ndarray]: """Run a full MoE block under a single fused custom_vjp on the TE EP path. @@ -2062,11 +2058,11 @@ def moe( Requesting cuDNN fusion reserves 256 tokens, also when falling back, to preserve compatibility with EP bootstrap buffer sizing. - Set ``NVTE_JAX_TEMP_FLAG_FOR_ABHINAV_CUDNN_GROUPED_GEMM_FUSION=1`` to use cuDNN's - JAX grouped MXFP8 APIs: Rubin GLU first, then generic SM100+ SwiGLU. - Ineligible calls warn - and fall back to TE's regular grouped-GEMM implementation. API signatures - and GPU capability determine support; fallbacks emit an actionable warning. + use_cudnn_fusion : bool + Defaults to ``True``: try cuDNN's JAX grouped MXFP8 APIs, Rubin GLU first, + then generic SM100+ SwiGLU. ``False`` uses unfused TE grouped GEMM. + Ineligible calls warn and fall back to TE's regular grouped-GEMM + implementation. API signatures and GPU capability determine support. MeshResource fields name physical mesh axes, not Flax logical axes. ``input_axes``, ``gate_kernel_axes``, ``wi_kernel_axes`` and @@ -2077,6 +2073,8 @@ def moe( See module docstring for the rest of the parameter semantics and the surrounding design rationale. """ + if not isinstance(use_cudnn_fusion, bool): + raise TypeError("use_cudnn_fusion must be a bool") if ep_axis is not None or data_parallelism_axes is not None or weight_gather is not None: call_args = locals().copy() resource, quantize = _resolve_moe_mesh_resource( @@ -2128,8 +2126,7 @@ def moe( expert_bias_arg = expert_bias.astype(jnp.float32) use_cudnn_jax_fusion = False - cudnn_native_weight_layout = wi.ndim == 3 and wi.shape[-1] == x.shape[-1] - if _use_cudnn_cutedsl_fusion_from_env(): + if use_cudnn_fusion: rejection_reasons = _cudnn_jax_fusion_rejection_reasons( x, wi, @@ -2141,11 +2138,6 @@ def moe( ep_axis=ep_axis, ) use_cudnn_jax_fusion = _select_cudnn_jax_fusion(rejection_reasons) - elif cudnn_native_weight_layout: - raise ValueError( - "cuDNN-native MoE weight layout requires the fused cuDNN grouped-GEMM path; " - f"set {_CUDNN_JAX_ENV}=1." - ) output, aux_loss, total_recv_tokens = _moe( x, @@ -2180,6 +2172,7 @@ def moe( wo_checkpoint_name, dispatch_checkpoint_name, quant_before_fsdp_ag, + use_cudnn_fusion, ) if aux_loss_coeff <= 0.0: aux_loss = None From fdc1e3c3157fde7f202778fd2bae240c358b544c Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Tue, 6 Oct 2026 12:25:29 -0700 Subject: [PATCH 25/38] Disambiguate JAX MoE weight layouts and fix FSDP fallback axes Signed-off-by: Jeremy Berchtold --- .../jax/test_grouped_gemm_swiglu_fallback.py | 48 ++++- tests/jax/test_moe_fsdp_layout.py | 174 ++++++++++++++++++ transformer_engine/jax/flax/moe.py | 1 + transformer_engine/jax/moe.py | 70 ++++++- 4 files changed, 279 insertions(+), 14 deletions(-) create mode 100644 tests/jax/test_moe_fsdp_layout.py diff --git a/tests/jax/test_grouped_gemm_swiglu_fallback.py b/tests/jax/test_grouped_gemm_swiglu_fallback.py index edabdbd84c3..eb3dce90148 100644 --- a/tests/jax/test_grouped_gemm_swiglu_fallback.py +++ b/tests/jax/test_grouped_gemm_swiglu_fallback.py @@ -163,12 +163,17 @@ def test_installed_frontend_contract(): @pytest.mark.parametrize("request_fusion", [None, True, False]) @pytest.mark.parametrize("native_layout", [False, True]) +@pytest.mark.parametrize("square", [False, True]) @pytest.mark.parametrize( "missing,expected", - [(None, "rubin"), ("grouped_gemm_glu", "blackwell"), ("grouped_gemm_dswiglu", False)], + [ + (None, "rubin"), + ("grouped_gemm_glu", "blackwell"), + ("grouped_gemm_dswiglu", False), + ], ) def test_public_moe_passes_selected_path( - frontend, monkeypatch, native_layout, missing, expected, request_fusion + frontend, monkeypatch, native_layout, square, missing, expected, request_fusion ): import jax import jax.numpy as jnp @@ -189,6 +194,7 @@ def test_public_moe_passes_selected_path( def execute(*args): bound = signature.bind(*args).arguments assert bound["use_cudnn_fusion"] is (request_fusion is not False) + assert bound["cudnn_native_weight_layout"] is native_layout received.append(bound["use_cudnn_jax_fusion"]) return args[0], None, jnp.asarray(0) @@ -196,6 +202,9 @@ def execute(*args): x = jnp.ones((1, 1, 128), jnp.bfloat16) wi = jnp.ones((2, 256, 128) if native_layout else (2, 128, 256), jnp.bfloat16) kwargs = {} if request_fusion is None else {"use_cudnn_fusion": request_fusion} + if square: + wi = jnp.ones((2, 128, 128), jnp.bfloat16) + kwargs["cudnn_native_weight_layout"] = native_layout if request_fusion is False: expected = False with warnings.catch_warnings(record=True) as caught: @@ -216,13 +225,16 @@ def execute(*args): @pytest.mark.parametrize("path", ["rubin", "blackwell"]) -def test_forward_uses_selected_kernel(monkeypatch, path): +@pytest.mark.parametrize("native_layout", [False, True]) +@pytest.mark.parametrize("gather_gated_dimension", [False, True]) +def test_forward_uses_selected_kernel(monkeypatch, path, native_layout, gather_gated_dimension): import jax.numpy as jnp class SelectedKernel(Exception): pass def selected(*args, **kwargs): + assert args[1].shape == (1, 256, 128) raise SelectedKernel def rejected(*args, **kwargs): @@ -241,19 +253,30 @@ def quantize(data, *args, **kwargs): ) monkeypatch.setattr(moe.tex, "grouped_quantize", quantize) + + def gather(tensor, fsdp_axis, fsdp_size, sharded_axis): + data = tensor.get_tensor().data + assert sharded_axis == (1 if native_layout else 2) + return quantize(jnp.concatenate([data, data], axis=sharded_axis)) + + monkeypatch.setattr(moe, "_gather_quantized_weight", gather) quantizers = SimpleNamespace(x=SimpleNamespace(q_dtype=jnp.float8_e4m3fn), kernel=None) kwargs = dict.fromkeys(inspect.signature(moe._ffn_fwd_per_shard).parameters) + combined = 128 if gather_gated_dimension else 256 kwargs.update( recv_tokens_local=jnp.ones((1, 256, 128), jnp.bfloat16), recv_topk_weights_local=jnp.ones((1, 256)), token_counts_local=jnp.asarray([256]), - wi=jnp.ones((1, 128, 256), jnp.bfloat16), + wi=jnp.ones((1, combined, 128) if native_layout else (1, 128, combined), jnp.bfloat16), wo=jnp.ones((1, 128, 128), jnp.bfloat16), quantizer_sets=(quantizers, quantizers), num_local_experts=1, use_cudnn_jax_fusion=path, - cudnn_native_weight_layout=False, - quant_before_fsdp_ag=False, + cudnn_native_weight_layout=native_layout, + quant_before_fsdp_ag=gather_gated_dimension, + wi_fsdp_axis=(1 if native_layout else 2), + fsdp_axis="fsdp", + fsdp_size=2, ) with pytest.raises(SelectedKernel): moe._ffn_fwd_per_shard(**kwargs) @@ -276,7 +299,10 @@ def execute(inputs, *args, **kwargs): monkeypatch.setattr(flax_moe, "moe", execute) kwargs = {} if request_fusion is None else {"use_cudnn_fusion": request_fusion} block = _MoEBlock( - num_experts=2, intermediate_size=32, mesh_resource=MeshResource(ep_resource="ep"), **kwargs + num_experts=2, + intermediate_size=32, + mesh_resource=MeshResource(ep_resource="ep"), + **kwargs, ) block.init(jax.random.PRNGKey(0), jnp.ones((1, 1, 32))) assert received == [request_fusion is not False] @@ -296,7 +322,13 @@ def test_capacity_follows_explicit_fusion_bool(request_fusion): def test_fusion_argument_requires_bool(invalid): with pytest.raises(TypeError, match="use_cudnn_fusion must be a bool"): moe.moe( - None, None, None, None, num_experts=2, num_experts_per_tok=1, use_cudnn_fusion=invalid + None, + None, + None, + None, + num_experts=2, + num_experts_per_tok=1, + use_cudnn_fusion=invalid, ) with pytest.raises(TypeError, match="use_cudnn_fusion must be a bool"): moe.get_moe_recv_capacity_per_rank( diff --git a/tests/jax/test_moe_fsdp_layout.py b/tests/jax/test_moe_fsdp_layout.py new file mode 100644 index 00000000000..5ee64c52b57 --- /dev/null +++ b/tests/jax/test_moe_fsdp_layout.py @@ -0,0 +1,174 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. +"""Isolate MoE weight layout errors using real packing and FSDP collectives. + +The fallback tests replace FP8 quantization and grouped GEMM with lossless +wrappers and JAX matmul to isolate layout/axis handling from kernel accuracy. +Known failures are strict xfails; remove the marks when fixing the bugs. +""" + +import importlib +from types import SimpleNamespace + +import jax +import jax.numpy as jnp +import numpy as np +import pytest +from jax.experimental.shard_map import shard_map +from jax.sharding import Mesh, PartitionSpec as P + +moe = importlib.import_module("transformer_engine.jax.moe") + + +def fsdp_mesh(): + if len(jax.devices()) < 2: + pytest.skip("Requires two JAX devices for a real FSDP all-gather") + return Mesh(np.asarray(jax.devices()[:2]), ("fsdp",)) + + +@pytest.mark.xfail( + strict=True, + reason="Packing contiguous gated-dimension shards changes global gate/up pairing", + raises=AssertionError, +) +def test_standard_gated_shards_match_unsharded_swiglu(): + mesh = fsdp_mesh() + # Each rank receives 64 contiguous columns: rank 0 gets only gates (1), + # rank 1 only ups (3). Each local half still meets pack's 32-column rule. + gate = jnp.ones((1, 32, 64), jnp.float32) + up = 3 * gate + wi = jnp.concatenate((gate, up), axis=-1) + x = jnp.ones((1, 32), jnp.float32) / 32 + wo = jnp.ones((64, 1), jnp.float32) / 64 + + def output(packed): + global_gate, global_up = moe.tex.unpack_swiglu_pair(packed) + return (jax.nn.silu(x @ global_gate[0]) * (x @ global_up[0])) @ wo + + expected = output(moe.tex.pack_swiglu_pair(gate, up)) + + def local_forward(local_wi): + local_gate, local_up = jnp.split(local_wi, 2, axis=-1) + packed = moe.tex.pack_swiglu_pair(local_gate, local_up) + gathered = jax.lax.all_gather(packed, "fsdp", axis=2, tiled=True) + return output(gathered) + + actual = shard_map( + local_forward, + mesh=mesh, + in_specs=P(None, None, "fsdp"), + out_specs=P(), + check_rep=False, + )(wi) + # expected=3*silu(1)=2.193176; actual=(silu(1)+3*silu(3))/2=4.652113. + np.testing.assert_allclose(actual, expected, rtol=1e-6) + + +@pytest.mark.parametrize("native_layout", [False, True]) +def test_unfused_k_sharded_weight_matches_unsharded(monkeypatch, native_layout): + mesh = fsdp_mesh() + + class LosslessTensor: + def __init__(self, data): + self.data = data + + def get_tensor(self, **kwargs): + return self + + def checkpoint(self, quantizer): + return self + + def quantize(data, *args, **kwargs): + return LosslessTensor(data) + + def gather(tensor, fsdp_axis, fsdp_size, sharded_axis): + assert fsdp_size == 2 + return LosslessTensor( + jax.lax.all_gather(tensor.data, fsdp_axis, axis=sharded_axis, tiled=True) + ) + + def gemm(lhs, rhs, *, contracting_dims, bias): + assert contracting_dims == ((1,), (1,)) + # One expert is sufficient to reproduce the matrix-axis error. + return lhs.data @ rhs.data[0] + + monkeypatch.setattr(moe.tex, "grouped_quantize", quantize) + monkeypatch.setattr(moe, "_gather_quantized_weight", gather) + monkeypatch.setattr(moe.tex, "grouped_gemm", gemm) + quantizers = SimpleNamespace(x=None, kernel=None) + gate = jnp.ones((1, 64, 64), jnp.float32) / 64 + up = 2 * gate + standard_wi = jnp.concatenate((gate, up), axis=-1) + wi = moe.tex.pack_swiglu_pair(gate, up).transpose(0, 2, 1) if native_layout else standard_wi + x = jnp.ones((1, 1, 64), jnp.float32) + wo = jnp.ones((1, 64, 64), jnp.float32) / 64 + + def forward(weight, quant_before_fsdp_ag): + result, _ = moe._ffn_fwd_per_shard( + x, + jnp.ones((1, 1)), + jnp.asarray([1]), + weight, + wo, + None, + None, + None, + (quantizers, quantizers), + num_local_experts=1, + activation_type="silu", + apply_topk_weights_early=False, + use_cudnn_jax_fusion=False, + wi_0_checkpoint_name=None, + wi_1_checkpoint_name=None, + wo_checkpoint_name=None, + cudnn_native_weight_layout=native_layout, + quant_before_fsdp_ag=quant_before_fsdp_ag, + fsdp_axis="fsdp", + fsdp_size=2, + wi_fsdp_axis=(2 if native_layout else 1), + wo_fsdp_axis=None, + ) + return result + + expected = forward(wi, False) + independent_reference = (jax.nn.silu(x @ gate[0]) * (x @ up[0])) @ wo[0] + np.testing.assert_allclose(expected, independent_reference, rtol=1e-6) + actual = shard_map( + lambda local_wi: forward(local_wi, True), + mesh=mesh, + in_specs=P(None, None, "fsdp") if native_layout else P(None, "fsdp", None), + out_specs=P(), + check_rep=False, + )(wi) + np.testing.assert_allclose(actual, expected, rtol=1e-6) + + +@pytest.mark.parametrize( + "shape,flag,expected", + [ + ((2, 128, 256), None, False), + ((2, 256, 128), None, True), + ((2, 128, 128), False, False), + ((2, 128, 128), True, True), + ((2, 128, 128), None, ValueError), + ((2, 128, 256), True, ValueError), + ((2, 256, 128), False, ValueError), + ((2, 128, 256), "false", TypeError), + ], +) +def test_layout_resolution(shape, flag, expected): + wi = jnp.zeros(shape) + if isinstance(expected, type): + with pytest.raises(expected): + moe._resolve_cudnn_native_weight_layout(wi, 128, flag) + else: + assert moe._resolve_cudnn_native_weight_layout(wi, 128, flag) is expected + + +def test_layout_argument_is_optional_and_static(): + import inspect + + assert inspect.signature(moe.moe).parameters["cudnn_native_weight_layout"].default is None + assert inspect.signature(moe._moe).parameters["cudnn_native_weight_layout"].default is None + assert 33 in moe._moe.nondiff_argnums diff --git a/transformer_engine/jax/flax/moe.py b/transformer_engine/jax/flax/moe.py index cf9280fa54d..d65116c903f 100644 --- a/transformer_engine/jax/flax/moe.py +++ b/transformer_engine/jax/flax/moe.py @@ -329,6 +329,7 @@ def make_grouped_quantizer_set(postfix): scaling_factor=self.scaling_factor, aux_loss_coeff=self.aux_loss_coeff, use_cudnn_fusion=self.use_cudnn_fusion, + cudnn_native_weight_layout=False, apply_topk_weights_early=self.apply_topk_weights_early, quantizer_sets=quantizer_sets, recv_capacity_per_rank=self.recv_capacity_per_rank, diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index e56355967ff..8d9afff9c9d 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -286,6 +286,7 @@ def _select_cudnn_jax_fusion(rejection_reasons: list[str]) -> str | bool: ) return path + def _cudnn_jax_fusion_rejection_reasons( x, wi, @@ -755,6 +756,25 @@ def _gather_quantized_matrix_scales(tensor, fsdp_axis, data_axis, local_scale_sh return jnp.concatenate(gathered_scales) +def _resolve_cudnn_native_weight_layout(wi, hidden, layout=None): + """Resolve FC1 parameter storage using its global shape or an explicit flag.""" + if layout is not None and not isinstance(layout, bool): + raise TypeError("cudnn_native_weight_layout must be a bool or None") + standard = wi.ndim == 3 and wi.shape[1] == hidden and wi.shape[2] % 2 == 0 + native = wi.ndim == 3 and wi.shape[2] == hidden and wi.shape[1] % 64 == 0 + if layout is None: + if standard and native: + raise ValueError( + "Ambiguous square FC1 weight layout; pass cudnn_native_weight_layout=" + "False for standard [E,K,2N] or True for native [E,2N,K] storage." + ) + if standard or native: + return native + elif native if layout else standard: + return layout + raise ValueError(f"Invalid FC1 weight shape {wi.shape} for K={hidden} and layout={layout}") + + def _weight_fsdp_axis(spec, fsdp_axis): """Find the tensor dimension partitioned by the physical FSDP resource.""" return next( @@ -815,6 +835,8 @@ def _ffn_fwd_per_shard( wi_interleaved = wi.transpose(0, 2, 1) wi_gate, wi_up = tex.unpack_swiglu_pair(wi_interleaved) wi_for_gemm = jnp.concatenate((wi_gate, wi_up), axis=-1) + if wi_fsdp_axis in (1, 2): + wi_fsdp_axis = 3 - wi_fsdp_axis else: wi_for_gemm = wi wi_combined_bias = ( @@ -842,7 +864,8 @@ def _ffn_fwd_per_shard( casted_wi_rhs = casted_wi.get_tensor( usage=TensorUsage.LHS if cudnn_native_weight_layout else TensorUsage.RHS ) - combined = wi_for_gemm.shape[-2] if cudnn_native_weight_layout else wi_for_gemm.shape[-1] + # Quantized gathering may have expanded the gated dimension since packing. + combined = casted_wi_rhs.data.size // (num_local_experts * hidden) padded_offsets = jnp.cumsum(group_sizes, dtype=jnp.int32) prob = ( recv_w_flat[:, None, None] @@ -1136,7 +1159,10 @@ def _ffn_bwd_per_shard( d_sorted_x = tex.grouped_gemm( casted_d_combined.get_tensor(usage=TensorUsage.LHS), casted_wi_rhs_trans, - contracting_dims=((1,), (1 if use_cudnn_jax_fusion and cudnn_native_weight_layout else 2,)), + contracting_dims=( + (1,), + (1 if use_cudnn_jax_fusion and cudnn_native_weight_layout else 2,), + ), ) if use_cudnn_jax_fusion and cudnn_native_weight_layout: # dY^T @ X directly produces [E,2N,K], matching the persistent native @@ -1217,6 +1243,7 @@ def _moe_fwd_rule( dispatch_checkpoint_name, quant_before_fsdp_ag, use_cudnn_fusion: bool = True, + cudnn_native_weight_layout: Optional[bool] = None, ): """Forward: gate -> topk -> ep_dispatch -> FFN -> ep_combine. @@ -1251,7 +1278,9 @@ def _moe_fwd_rule( ) B, S, H = x.shape K = num_experts_per_tok - cudnn_native_weight_layout = wi.ndim == 3 and wi.shape[-1] == H + cudnn_native_weight_layout = _resolve_cudnn_native_weight_layout( + wi, H, cudnn_native_weight_layout + ) kernel_spec = P(ep_axis, None, None) wi_input_spec = wo_input_spec = kernel_spec wi_fsdp_axis = wo_fsdp_axis = None @@ -1607,6 +1636,7 @@ def _moe_bwd_rule( dispatch_checkpoint_name, quant_before_fsdp_ag, use_cudnn_fusion, + cudnn_native_weight_layout, residuals, cotangents, ): @@ -1624,6 +1654,7 @@ def _moe_bwd_rule( dispatch_checkpoint_name, quant_before_fsdp_ag, use_cudnn_fusion, + cudnn_native_weight_layout, ) # captured / unused in bwd from jax.experimental.shard_map import shard_map @@ -1741,7 +1772,15 @@ def _ffn_bwd_body(*args): bias_spec, ) else: - bwd_out_specs = (ep3_spec, ep2_spec, kernel_spec, kernel_spec, None, None, None) + bwd_out_specs = ( + ep3_spec, + ep2_spec, + kernel_spec, + kernel_spec, + None, + None, + None, + ) ( d_sorted_x, @@ -1873,7 +1912,7 @@ def _ffn_bwd_body(*args): # ============================================================================= -@partial(jax.custom_vjp, nondiff_argnums=tuple(range(9, 33))) +@partial(jax.custom_vjp, nondiff_argnums=tuple(range(9, 34))) def _moe( x, gate_kernel, @@ -1908,6 +1947,7 @@ def _moe( dispatch_checkpoint_name, quant_before_fsdp_ag, use_cudnn_fusion: bool = True, + cudnn_native_weight_layout: Optional[bool] = None, ): primal, _ = _moe_fwd_rule( x, @@ -1943,6 +1983,7 @@ def _moe( dispatch_checkpoint_name, quant_before_fsdp_ag, use_cudnn_fusion, + cudnn_native_weight_layout, ) return primal @@ -1990,6 +2031,7 @@ def moe( dispatch_checkpoint_name: Optional[str] = None, weight_gather: Optional[WeightGather] = None, use_cudnn_fusion: bool = True, + cudnn_native_weight_layout: Optional[bool] = None, ) -> Tuple[jnp.ndarray, Optional[jnp.ndarray], jnp.ndarray]: """Run a full MoE block under a single fused custom_vjp on the TE EP path. @@ -2064,6 +2106,13 @@ def moe( Ineligible calls warn and fall back to TE's regular grouped-GEMM implementation. API signatures and GPU capability determine support. + cudnn_native_weight_layout : Optional[bool] + ``False`` selects standard contiguous gate/up [E,K,2N] storage; + ``True`` selects native [E,2N,K] storage with alternating 32-column + gate/up blocks. Defaults to ``None``, inferring storage from the global + shape. Ambiguous square shapes require an explicit flag. This option + describes parameter storage independently of the selected execution path. + MeshResource fields name physical mesh axes, not Flax logical axes. ``input_axes``, ``gate_kernel_axes``, ``wi_kernel_axes`` and ``wo_kernel_axes`` remain logical-axis tuples resolved through the active @@ -2078,7 +2127,11 @@ def moe( if ep_axis is not None or data_parallelism_axes is not None or weight_gather is not None: call_args = locals().copy() resource, quantize = _resolve_moe_mesh_resource( - mesh_resource, quant_before_fsdp_ag, ep_axis, data_parallelism_axes, weight_gather + mesh_resource, + quant_before_fsdp_ag, + ep_axis, + data_parallelism_axes, + weight_gather, ) for name in ("ep_axis", "data_parallelism_axes", "weight_gather"): call_args.pop(name) @@ -2117,6 +2170,10 @@ def moe( ) x = _with_sharding_constraint_cast_bwd(x, NamedSharding(mesh, expected_spec)) + cudnn_native_weight_layout = _resolve_cudnn_native_weight_layout( + wi, x.shape[-1], cudnn_native_weight_layout + ) + # custom_vjp can't trace through None args; lower expert_bias to an # empty shape-(0,) tensor that fused_topk_with_score_function treats # as "no bias". @@ -2173,6 +2230,7 @@ def moe( dispatch_checkpoint_name, quant_before_fsdp_ag, use_cudnn_fusion, + cudnn_native_weight_layout, ) if aux_loss_coeff <= 0.0: aux_loss = None From 600bedabf49fdafa2c9a86b44d7ce6ebc34a0da0 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Tue, 6 Oct 2026 12:41:44 -0700 Subject: [PATCH 26/38] Preserve gate/up pairing in quantized JAX MoE FSDP gathers Standard FC1 weights sharded across the gated dimension must not be split into local gate/up halves before gathering. Native fallback weights must likewise retain their interleaved column order until gathering completes. Quantize each shard, gather its FP8 data and scales, then permute both into the global execution layout without a BF16 all-gather or requantization. Extract FC1 preparation, grouped tensor reconstruction, and shared scale reading/padding/swizzling into helpers to simplify the FFN forward body. Replace the expected-failure isolation with a passing MXFP8 regression. Check exact FP8 data and active scale bytes across both layouts, execution paths, and expert/K/gated sharding. Include reference-encoded 32- and 96-column shards, and extend the four-GPU MoE test to gated sharding with forward and backward parity in fused and unfused modes. Validation: 166 focused tests passed; six four-GPU parity cases passed in each execution mode. Common/JAX pylint and changed-file pre-commit passed. The repository license check reports existing failures in nccl_ep/ and tests/jax/repro_cutedsl_moe_shardy.py; changed-file headers passed. Signed-off-by: Jeremy Berchtold --- .../jax/test_grouped_gemm_swiglu_fallback.py | 7 + tests/jax/test_moe_fsdp_layout.py | 226 ++++++++++++++++-- tests/jax/test_te_ep_moe.py | 31 ++- transformer_engine/jax/moe.py | 220 +++++++++++------ 4 files changed, 378 insertions(+), 106 deletions(-) diff --git a/tests/jax/test_grouped_gemm_swiglu_fallback.py b/tests/jax/test_grouped_gemm_swiglu_fallback.py index eb3dce90148..1d784358ac4 100644 --- a/tests/jax/test_grouped_gemm_swiglu_fallback.py +++ b/tests/jax/test_grouped_gemm_swiglu_fallback.py @@ -260,6 +260,13 @@ def gather(tensor, fsdp_axis, fsdp_size, sharded_axis): return quantize(jnp.concatenate([data, data], axis=sharded_axis)) monkeypatch.setattr(moe, "_gather_quantized_weight", gather) + + def reorder(tensor, *, interleave): + assert interleave and gather_gated_dimension and not native_layout + gate, up = jnp.split(tensor.get_tensor().data, 2, axis=-1) + return quantize(moe.tex.pack_swiglu_pair(gate, up)) + + monkeypatch.setattr(moe, "_reorder_quantized_swiglu_weight", reorder) quantizers = SimpleNamespace(x=SimpleNamespace(q_dtype=jnp.float8_e4m3fn), kernel=None) kwargs = dict.fromkeys(inspect.signature(moe._ffn_fwd_per_shard).parameters) combined = 128 if gather_gated_dimension else 256 diff --git a/tests/jax/test_moe_fsdp_layout.py b/tests/jax/test_moe_fsdp_layout.py index 5ee64c52b57..b961fe32ac9 100644 --- a/tests/jax/test_moe_fsdp_layout.py +++ b/tests/jax/test_moe_fsdp_layout.py @@ -5,7 +5,7 @@ The fallback tests replace FP8 quantization and grouped GEMM with lossless wrappers and JAX matmul to isolate layout/axis handling from kernel accuracy. -Known failures are strict xfails; remove the marks when fixing the bugs. +Weight preparation tests use real MXFP8 quantization, scales, and all-gathers. """ import importlib @@ -17,6 +17,8 @@ import pytest from jax.experimental.shard_map import shard_map from jax.sharding import Mesh, PartitionSpec as P +from transformer_engine.jax.quantize import QuantizerFactory, QuantizeLayout, ScalingMode +from transformer_engine.jax.quantize.dequantizer import Dequantizer moe = importlib.import_module("transformer_engine.jax.moe") @@ -27,44 +29,232 @@ def fsdp_mesh(): return Mesh(np.asarray(jax.devices()[:2]), ("fsdp",)) -@pytest.mark.xfail( - strict=True, - reason="Packing contiguous gated-dimension shards changes global gate/up pairing", - raises=AssertionError, -) def test_standard_gated_shards_match_unsharded_swiglu(): mesh = fsdp_mesh() - # Each rank receives 64 contiguous columns: rank 0 gets only gates (1), + # Each rank receives 128 contiguous columns: rank 0 gets only gates (1), # rank 1 only ups (3). Each local half still meets pack's 32-column rule. - gate = jnp.ones((1, 32, 64), jnp.float32) + gate = jnp.ones((1, 128, 128), jnp.bfloat16) up = 3 * gate wi = jnp.concatenate((gate, up), axis=-1) - x = jnp.ones((1, 32), jnp.float32) / 32 - wo = jnp.ones((64, 1), jnp.float32) / 64 + x = jnp.ones((1, 128), jnp.float32) / 128 + wo = jnp.ones((128, 1), jnp.float32) / 128 + quantizer = weight_quantizer(1) - def output(packed): + def output(tensor): + packed = jnp.concatenate(tensor.rowwise_tensor.dequantize()) global_gate, global_up = moe.tex.unpack_swiglu_pair(packed) return (jax.nn.silu(x @ global_gate[0]) * (x @ global_up[0])) @ wo - expected = output(moe.tex.pack_swiglu_pair(gate, up)) + expected = output(prepare_weight(wi, quantizer, fused=True, native=False)) def local_forward(local_wi): - local_gate, local_up = jnp.split(local_wi, 2, axis=-1) - packed = moe.tex.pack_swiglu_pair(local_gate, local_up) - gathered = jax.lax.all_gather(packed, "fsdp", axis=2, tiled=True) - return output(gathered) + return prepare_weight(local_wi, quantizer, fused=True, native=False, axis=2) - actual = shard_map( + gathered = shard_map( local_forward, mesh=mesh, in_specs=P(None, None, "fsdp"), out_specs=P(), check_rep=False, )(wi) - # expected=3*silu(1)=2.193176; actual=(silu(1)+3*silu(3))/2=4.652113. + actual = output(gathered) + # Both paths now yield 3*silu(1)=2.193176, rather than FSDP yielding 4.652113. np.testing.assert_allclose(actual, expected, rtol=1e-6) +def weight_quantizer(experts): + return QuantizerFactory.create( + scaling_mode=ScalingMode.MXFP8_1D_SCALING, + q_dtype=jnp.float8_e4m3fn, + q_layout=QuantizeLayout.ROWWISE_COLWISE, + n_groups=experts, + ) + + +def prepare_weight(wi, quantizer, *, fused, native, axis=None): + return moe._prepare_fc1_weight( + wi, + quantizer, + fused=fused, + native_layout=native, + quant_before_fsdp_ag=axis is not None, + fsdp_axis="fsdp", + fsdp_size=2, + sharded_axis=axis, + ) + + +def plain_weight_scales(tensor): + """Ignore unused allocation padding while checking every active scale byte.""" + kwargs = dict( + data_layout=tensor.data_layout, + is_colwise=tensor.is_colwise, + flatten_axis=tensor.flatten_axis - 1, + ) + padded = tensor.scaling_mode.get_scale_shape( + tensor.original_shape[1:], is_padded=True, **kwargs + ) + unpadded = tensor.scaling_mode.get_scale_shape( + tensor.original_shape[1:], is_padded=False, **kwargs + ) + scale_size = int(np.prod(padded)) + return jnp.stack( + [ + moe._unswizzle_mxfp8_grouped_scale( + tensor.scale_inv[e * scale_size : (e + 1) * scale_size], padded, tensor.is_colwise + )[: unpadded[0], : unpadded[1]] + for e in range(tensor.original_shape[0]) + ] + ).view(jnp.uint8) + + +def reference_quantize_weight(data, *args, **kwargs): + """Pure-JAX MXFP8 encoding for positive powers of two, including small shards. + + This avoids the legacy grouped quantizer's host-copy path for dimensions + below 128. The production gather, FP8 permutation, and swizzle still run. + """ + copies = [] + mode = ScalingMode.MXFP8_1D_SCALING + experts, rows, columns = data.shape + for colwise in (False, True): + blocks = ( + data.reshape(experts, rows // 32, 32, columns) + if colwise + else data.reshape(experts, rows, columns // 32, 32) + ) + maxima = jnp.max(blocks.astype(jnp.float32), axis=2 if colwise else 3) + exponent = jnp.log2(maxima).astype(jnp.int32) - 8 + inverse_scale = jnp.exp2(exponent.astype(jnp.float32)) + broadcast_scale = jnp.repeat(inverse_scale, 32, axis=1 if colwise else 2) + quantized = (data / broadcast_scale).astype(jnp.float8_e4m3fn) + scales = (exponent + 127).astype(jnp.uint8) + padded = mode.get_scale_shape((rows, columns), is_colwise=colwise, is_padded=True) + swizzled = [ + moe.swizzled_scale( + jnp.pad(scale, ((0, padded[0] - scale.shape[0]), (0, padded[1] - scale.shape[1]))), + 1, + colwise, + ).reshape(-1) + for scale in scales + ] + scale_inv = jnp.concatenate(swizzled) + scale_size = mode.get_grouped_scale_shape(data.shape, experts, colwise, flatten_axis=2)[0] + scale_inv = jnp.pad(scale_inv, (0, scale_size - scale_inv.size)) + copies.append( + moe.GroupedScaledTensor1x( + data=quantized.reshape(-1), + scale_inv=scale_inv, + amax=jnp.empty((0,), jnp.float32), + first_dims=None, + last_dims=None, + scaling_mode=mode, + dq_dtype=data.dtype, + _dq_func=Dequantizer.grouped_dequantize, + is_colwise=colwise, + data_layout="N", + flatten_axis=2, + original_shape=data.shape, + pre_swizzled=True, + ) + ) + return moe.ScaledTensor2x(*copies) + + +@pytest.mark.parametrize("native", [False, True]) +@pytest.mark.parametrize("width", [64, 192]) +def test_small_gated_shards_preserve_quantized_pairs(monkeypatch, native, width): + mesh = fsdp_mesh() + source = jnp.exp2( + (jnp.arange(64) // 32)[None, :, None] + + (jnp.arange(width) // 32)[None, None, :] + + jnp.arange(2)[:, None, None] + ).astype(jnp.bfloat16) + gate, up = jnp.split(source, 2, axis=-1) + wi = moe.tex.pack_swiglu_pair(gate, up).transpose(0, 2, 1) if native else source + monkeypatch.setattr(moe.tex, "grouped_quantize", reference_quantize_weight) + expected = prepare_weight(wi, None, fused=not native, native=native) + actual = shard_map( + lambda shard: prepare_weight( + shard, None, fused=not native, native=native, axis=1 if native else 2 + ), + mesh=mesh, + in_specs=P(None, "fsdp", None) if native else P(None, None, "fsdp"), + out_specs=P(), + check_rep=False, + )(wi) + for actual_copy, expected_copy in zip( + (actual.rowwise_tensor, actual.colwise_tensor), + (expected.rowwise_tensor, expected.colwise_tensor), + ): + np.testing.assert_array_equal( + actual_copy.data.view(jnp.uint8), expected_copy.data.view(jnp.uint8) + ) + np.testing.assert_array_equal( + plain_weight_scales(actual_copy), plain_weight_scales(expected_copy) + ) + decoded = jnp.concatenate(actual_copy.dequantize()) + execution_weight = moe.tex.pack_swiglu_pair(gate, up) if not native else source + np.testing.assert_array_equal(decoded, execution_weight) + + +@pytest.mark.parametrize("native", [False, True]) +@pytest.mark.parametrize("fused", [False, True]) +@pytest.mark.parametrize("shard_dimension", ["expert", "k", "gated"]) +@pytest.mark.parametrize("width", [256, 768]) +def test_quantized_fc1_gather_matches_unsharded(monkeypatch, native, fused, shard_dimension, width): + mesh = fsdp_mesh() + # Distinct column-block magnitudes, row magnitudes, and expert magnitudes + # expose scale permutations that equal-valued/constant-scale inputs hide. + source = jax.random.normal(jax.random.PRNGKey(9), (2, 256, width)) + column_magnitude = jnp.exp2(jnp.arange(width) // 32 % 7 - 3) + row_magnitude = jnp.exp2(jnp.arange(256) // 32 - 3) + source = ( + source + * column_magnitude + * row_magnitude[None, :, None] + * jnp.asarray([1.0, 4.0])[:, None, None] + ).astype(jnp.bfloat16) + gate, up = jnp.split(source, 2, axis=-1) + wi = moe.tex.pack_swiglu_pair(gate, up).transpose(0, 2, 1) if native else source + axis = {"expert": 0, "k": 2 if native else 1, "gated": 1 if native else 2}[shard_dimension] + quantizer = weight_quantizer(2) + expected = prepare_weight(wi, quantizer, fused=fused, native=native) + spec = [None, None, None] + spec[axis] = "fsdp" + quantized_shapes = [] + original_quantize = moe.tex.grouped_quantize + + def record_quantize(data, *args, **kwargs): + quantized_shapes.append(data.shape) + return original_quantize(data, *args, **kwargs) + + monkeypatch.setattr(moe.tex, "grouped_quantize", record_quantize) + actual = shard_map( + lambda shard: prepare_weight(shard, quantizer, fused=fused, native=native, axis=axis), + mesh=mesh, + in_specs=P(*spec), + out_specs=P(), + check_rep=False, + )(wi) + # Only a local shard is quantized, once; no global BF16 requantization. + assert quantized_shapes + assert all(int(np.prod(shape)) == wi.size // 2 for shape in quantized_shapes) + for actual_copy, expected_copy in zip( + (actual.rowwise_tensor, actual.colwise_tensor), + (expected.rowwise_tensor, expected.colwise_tensor), + ): + np.testing.assert_array_equal( + actual_copy.data.view(jnp.uint8), expected_copy.data.view(jnp.uint8) + ) + np.testing.assert_array_equal( + plain_weight_scales(actual_copy), plain_weight_scales(expected_copy) + ) + np.testing.assert_array_equal( + jnp.concatenate(actual_copy.dequantize()), jnp.concatenate(expected_copy.dequantize()) + ) + + @pytest.mark.parametrize("native_layout", [False, True]) def test_unfused_k_sharded_weight_matches_unsharded(monkeypatch, native_layout): mesh = fsdp_mesh() diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index 65fd68e22f9..3eb6b4a170f 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -158,6 +158,7 @@ def _read_cudnn_fusion_option() -> bool: LOGICAL_AXIS_RULES = ( ("exp", EP_AXIS), ("expert_weight_fsdp", (EP_AXIS, FSDP_AXIS)), + ("gated_weight_fsdp", FSDP_AXIS), ("embed", FSDP_AXIS), ("mlp", None), ("batch", (FSDP_AXIS, EP_AXIS)), @@ -598,15 +599,12 @@ def test_weight_gather_policy_axis_resolution(mesh): @pytest.mark.parametrize("native_weight_layout", [False, True]) -@pytest.mark.parametrize("expert_fsdp", [False, True]) +@pytest.mark.parametrize("fsdp_dimension", ["k", "gated", "expert"]) def test_quantized_weight_gather_matches_full_precision_gather( - mesh, monkeypatch, native_weight_layout, expert_fsdp + mesh, monkeypatch, native_weight_layout, fsdp_dimension ): """The FP8 weight gather retains forward and backward MoE semantics.""" - if native_weight_layout: - if not _USE_CUDNN_FUSION: - pytest.skip("Native weight layout requires cuDNN grouped GEMM fusion") - if native_weight_layout or expert_fsdp: + if native_weight_layout or fsdp_dimension != "k": from transformer_engine.jax import cpp_extensions as tex flax_moe_module = importlib.import_module("transformer_engine.jax.flax.moe") @@ -618,7 +616,20 @@ def layout_moe(*args, **kwargs): gate, up = jnp.split(args[2], 2, axis=-1) args[2] = tex.pack_swiglu_pair(gate, up).transpose(0, 2, 1) kwargs["wi_kernel_axes"] = ("exp", "mlp", "embed") - if expert_fsdp: + kwargs["cudnn_native_weight_layout"] = True + if fsdp_dimension == "gated": + kwargs["wi_kernel_axes"] = ( + ("exp", "gated_weight_fsdp", None) + if native_weight_layout + else ("exp", None, "gated_weight_fsdp") + ) + spec = ( + P(EP_AXIS, FSDP_AXIS, None) + if native_weight_layout + else P(EP_AXIS, None, FSDP_AXIS) + ) + args[2] = jax.lax.with_sharding_constraint(args[2], NamedSharding(mesh, spec)) + if fsdp_dimension == "expert": kwargs["wi_kernel_axes"] = ("expert_weight_fsdp", None, None) kwargs["wo_kernel_axes"] = ("expert_weight_fsdp", None, None) spec = P((EP_AXIS, FSDP_AXIS), None, None) @@ -899,6 +910,7 @@ def native_moe(*args, **kwargs): gate, up = jnp.split(args[2], 2, axis=-1) args[2] = tex.pack_swiglu_pair(gate, up).transpose(0, 2, 1) kwargs["wi_kernel_axes"] = ("exp", "mlp", "embed") + kwargs["cudnn_native_weight_layout"] = True return original_moe(*args, **kwargs) monkeypatch.setattr(flax_moe, "moe", native_moe) @@ -944,6 +956,7 @@ def native_moe(*args, **kwargs): gate, up = jnp.split(args[2], 2, axis=-1) args[2] = tex.pack_swiglu_pair(gate, up).transpose(0, 2, 1) kwargs["wi_kernel_axes"] = ("exp", "mlp", "embed") + kwargs["cudnn_native_weight_layout"] = True return original_moe(*args, **kwargs) monkeypatch.setattr(flax_moe, "moe", native_moe) @@ -966,9 +979,7 @@ def native_moe(*args, **kwargs): @pytest.mark.parametrize("apply_topk_weights_early", [False, True]) def test_mxfp8_forward_and_backward(self, mesh, apply_topk_weights_early, monkeypatch): if not _USE_CUDNN_FUSION: - pytest.skip( - "run separately with --use-cudnn-fusion=1" - ) + pytest.skip("run separately with --use-cudnn-fusion=1") rubin_calls = [] if get_device_compute_capability(0) == 107: from transformer_engine.jax import cpp_extensions as tex diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index 8d9afff9c9d..d51b2deed91 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -629,6 +629,25 @@ def _validate_moe_quantizer_sets( ) +def _with_grouped_weight_buffers(tensor, data, scale_inv, original_shape): + """Rebuild a dense grouped weight while preserving its quantization metadata.""" + return GroupedScaledTensor1x( + data=data, + scale_inv=scale_inv, + amax=tensor.amax, + first_dims=None, + last_dims=None, + scaling_mode=tensor.scaling_mode, + dq_dtype=tensor.dq_dtype, + _dq_func=tensor._dq_func, + is_colwise=tensor.is_colwise, + data_layout=tensor.data_layout, + flatten_axis=tensor.flatten_axis, + original_shape=original_shape, + pre_swizzled=True, + ) + + def _gather_quantized_weight(tensor, fsdp_axis: str, fsdp_size: int, sharded_axis: int): """Gather an MXFP8 grouped weight without gathering its BF16 source. @@ -691,68 +710,51 @@ def _gather_quantized_weight(tensor, fsdp_axis: str, fsdp_size: int, sharded_axi flatten_axis=tensor.flatten_axis, )[0] scale_inv = jnp.pad(scale_inv, (0, expected_scale_size - scale_inv.size)) - return GroupedScaledTensor1x( - data=data, - scale_inv=scale_inv, - amax=tensor.amax, - first_dims=None, - last_dims=None, - scaling_mode=tensor.scaling_mode, - dq_dtype=tensor.dq_dtype, - _dq_func=tensor._dq_func, - is_colwise=tensor.is_colwise, - data_layout=tensor.data_layout, - flatten_axis=tensor.flatten_axis, - original_shape=global_shape, - pre_swizzled=True, + return _with_grouped_weight_buffers(tensor, data, scale_inv, global_shape) + + +def _weight_scale_shapes(tensor, matrix_shape): + """Return padded and active per-expert MXFP8 scale shapes.""" + kwargs = { + "data_layout": tensor.data_layout, + "is_colwise": tensor.is_colwise, + "flatten_axis": tensor.flatten_axis - 1, + } + return ( + tensor.scaling_mode.get_scale_shape(matrix_shape, is_padded=True, **kwargs), + tensor.scaling_mode.get_scale_shape(matrix_shape, is_padded=False, **kwargs), ) +def _read_weight_scale(tensor, expert, padded_shape, plain_shape): + """Read one expert's active scales in matrix order.""" + scale_size = math.prod(padded_shape) + swizzled = jax.lax.dynamic_slice_in_dim(tensor.scale_inv, expert * scale_size, scale_size) + plain = _unswizzle_mxfp8_grouped_scale(swizzled, padded_shape, tensor.is_colwise) + return plain[: plain_shape[0], : plain_shape[1]] + + +def _swizzle_weight_scale(plain, padded_shape, is_colwise): + """Pad and swizzle one expert's scale matrix for grouped GEMM.""" + padded = jnp.pad( + plain, ((0, padded_shape[0] - plain.shape[0]), (0, padded_shape[1] - plain.shape[1])) + ) + return swizzled_scale(padded, 1, is_colwise).reshape(-1) + + def _gather_quantized_matrix_scales(tensor, fsdp_axis, data_axis, local_scale_shape, global_matrix): """Reassemble scales when FSDP splits each expert's matrix dimension.""" - local_matrix = tensor.original_shape[1:] - local_unpadded_shape = tensor.scaling_mode.get_scale_shape( - local_matrix, - data_layout=tensor.data_layout, - is_colwise=tensor.is_colwise, - is_padded=False, - flatten_axis=tensor.flatten_axis - 1, - ) - global_unpadded_shape = tensor.scaling_mode.get_scale_shape( - global_matrix, - data_layout=tensor.data_layout, - is_colwise=tensor.is_colwise, - is_padded=False, - flatten_axis=tensor.flatten_axis - 1, - ) - global_padded_shape = tensor.scaling_mode.get_scale_shape( - global_matrix, - data_layout=tensor.data_layout, - is_colwise=tensor.is_colwise, - is_padded=True, - flatten_axis=tensor.flatten_axis - 1, - ) + _, local_unpadded_shape = _weight_scale_shapes(tensor, tensor.original_shape[1:]) + global_padded_shape, global_unpadded_shape = _weight_scale_shapes(tensor, global_matrix) scale_axis = data_axis - 1 - local_scale_size = math.prod(local_scale_shape) gathered_scales = [] for expert in range(tensor.original_shape[0]): - local_swizzled = jax.lax.dynamic_slice_in_dim( - tensor.scale_inv, expert * local_scale_size, local_scale_size - ) - local_plain = _unswizzle_mxfp8_grouped_scale( - local_swizzled, local_scale_shape, tensor.is_colwise - ) - local_plain = local_plain[: local_unpadded_shape[0], : local_unpadded_shape[1]] + local_plain = _read_weight_scale(tensor, expert, local_scale_shape, local_unpadded_shape) full_plain = jax.lax.all_gather(local_plain, fsdp_axis, axis=scale_axis, tiled=True) assert full_plain.shape == global_unpadded_shape - full_padded = jnp.pad( - full_plain, - ( - (0, global_padded_shape[0] - global_unpadded_shape[0]), - (0, global_padded_shape[1] - global_unpadded_shape[1]), - ), + gathered_scales.append( + _swizzle_weight_scale(full_plain, global_padded_shape, tensor.is_colwise) ) - gathered_scales.append(swizzled_scale(full_padded, 1, tensor.is_colwise).reshape(-1)) return jnp.concatenate(gathered_scales) @@ -787,6 +789,81 @@ def _weight_fsdp_axis(spec, fsdp_axis): ) +def _swiglu_column_indices(width, interleave): + """Map contiguous gate/up halves to 32-column pairs, or undo that mapping.""" + if width % 64: + raise ValueError("The global gated dimension must be divisible by 64") + blocks = jnp.arange(width // 32) + blocks = ( + blocks.reshape(2, -1).T.reshape(-1) if interleave else blocks.reshape(-1, 2).T.reshape(-1) + ) + columns = (blocks[:, None] * 32 + jnp.arange(32)[None, :]).reshape(-1) + return blocks, columns + + +def _reorder_quantized_swiglu_weight(tensor, *, interleave): + """Permute gathered MXFP8 columns and their scales without requantization. + + Input logical storage is [E,K,2N]. Both rowwise and colwise copies may use + physical N or T layout. Gate/up boundaries and permutations are 32-aligned, + so each quantization block keeps its original data and inverse scale. + """ + if isinstance(tensor, ScaledTensor2x): + return ScaledTensor2x( + _reorder_quantized_swiglu_weight(tensor.rowwise_tensor, interleave=interleave), + _reorder_quantized_swiglu_weight(tensor.colwise_tensor, interleave=interleave), + ) + if not isinstance(tensor, GroupedScaledTensor1x): + raise TypeError("SwiGLU weight reordering requires grouped MXFP8 tensors") + if tensor.scaling_mode != ScalingMode.MXFP8_1D_SCALING or not tensor.pre_swizzled: + raise ValueError("SwiGLU weight reordering requires pre-swizzled MXFP8 scales") + + shape = tensor.original_shape + data_axis = 1 if tensor.data_layout == "T" else 2 + block_indices, column_indices = _swiglu_column_indices(shape[data_axis], interleave) + data = jnp.take(tensor.data.reshape(shape), column_indices, axis=data_axis).reshape(-1) + padded_shape, plain_shape = _weight_scale_shapes(tensor, shape[1:]) + scale_axis = data_axis - 1 + # Rowwise scales cover 32 columns; colwise scales have one entry per column. + scale_indices = ( + block_indices if plain_shape[scale_axis] == shape[data_axis] // 32 else column_indices + ) + reordered_scales = [] + for expert in range(shape[0]): + plain = _read_weight_scale(tensor, expert, padded_shape, plain_shape) + reordered = jnp.take(plain, scale_indices, axis=scale_axis) + reordered_scales.append(_swizzle_weight_scale(reordered, padded_shape, tensor.is_colwise)) + scale_inv = jnp.concatenate(reordered_scales) + scale_inv = jnp.pad(scale_inv, (0, tensor.scale_inv.size - scale_inv.size)) + return _with_grouped_weight_buffers(tensor, data, scale_inv, shape) + + +def _prepare_fc1_weight( + wi, quantizer, *, fused, native_layout, quant_before_fsdp_ag, fsdp_axis, fsdp_size, sharded_axis +): + """Convert FC1 storage and gather quantized shards in the correct logical order.""" + gated_axis = 1 if native_layout else 2 + convert_layout = bool(fused) != native_layout + convert_after_gather = convert_layout and quant_before_fsdp_ag and sharded_axis == gated_axis + if native_layout and not fused: + wi = wi.transpose(0, 2, 1) + if sharded_axis in (1, 2): + sharded_axis = 3 - sharded_axis + if not convert_after_gather: + gate, up = tex.unpack_swiglu_pair(wi) + wi = jnp.concatenate((gate, up), axis=-1) + elif fused and not native_layout and not convert_after_gather: + gate, up = jnp.split(wi, 2, axis=-1) + wi = tex.pack_swiglu_pair(gate, up) + + casted = tex.grouped_quantize(wi, quantizer, flatten_axis=-1) + if quant_before_fsdp_ag and sharded_axis is not None: + casted = _gather_quantized_weight(casted, fsdp_axis, fsdp_size, sharded_axis) + if convert_after_gather: + casted = _reorder_quantized_swiglu_weight(casted, interleave=bool(fused)) + return casted + + def _ffn_fwd_per_shard( recv_tokens_local: jnp.ndarray, recv_topk_weights_local: jnp.ndarray, @@ -821,24 +898,6 @@ def _ffn_fwd_per_shard( wi = wi.astype(sorted_x.dtype) wo = wo.astype(sorted_x.dtype) - # The cuDNN-native parameter is persistent [E,2N,K] storage with alternating - # 32-column gate/up blocks. Standard TE storage is [E,K,2N] with contiguous - # gate/up halves. Keep conversion only as a compatibility fallback. - if use_cudnn_jax_fusion: - if cudnn_native_weight_layout: - wi_for_gemm = wi - else: - wi_gate, wi_up = jnp.split(wi, 2, axis=-1) - wi_for_gemm = tex.pack_swiglu_pair(wi_gate, wi_up) - else: - if cudnn_native_weight_layout: - wi_interleaved = wi.transpose(0, 2, 1) - wi_gate, wi_up = tex.unpack_swiglu_pair(wi_interleaved) - wi_for_gemm = jnp.concatenate((wi_gate, wi_up), axis=-1) - if wi_fsdp_axis in (1, 2): - wi_fsdp_axis = 3 - wi_fsdp_axis - else: - wi_for_gemm = wi wi_combined_bias = ( jnp.concatenate([wi_0_bias, wi_1_bias], axis=-1) if wi_0_bias is not None else None ) @@ -850,14 +909,16 @@ def _ffn_fwd_per_shard( group_sizes, flatten_axis=-1, ) - casted_wi = tex.grouped_quantize(wi_for_gemm, fc1_quantizer_set.kernel, flatten_axis=-1) - if quant_before_fsdp_ag and wi_fsdp_axis is not None: - casted_wi = _gather_quantized_weight( - casted_wi, - fsdp_axis, - fsdp_size, - wi_fsdp_axis, - ) + casted_wi = _prepare_fc1_weight( + wi, + fc1_quantizer_set.kernel, + fused=use_cudnn_jax_fusion, + native_layout=cudnn_native_weight_layout, + quant_before_fsdp_ag=quant_before_fsdp_ag, + fsdp_axis=fsdp_axis, + fsdp_size=fsdp_size, + sharded_axis=wi_fsdp_axis, + ) casted_intermediate = None if use_cudnn_jax_fusion: casted_sorted_x_lhs = casted_sorted_x.get_tensor(usage=TensorUsage.LHS) @@ -2092,7 +2153,10 @@ def moe( full-precision weights first. ``True`` requires an FSDP resource and MXFP8 kernel quantizers. With active Flax logical-axis rules, the weight input specs follow ``wi_kernel_axes`` and ``wo_kernel_axes``; FSDP may - shard whole experts or a matrix dimension within each expert. + shard whole experts or a matrix dimension within each expert. When + FSDP shards the gated dimension, layout conversion permutes gathered + FP8 data and inverse scales, preserving global gate/up pairing without + gathering or requantizing full-precision weights. ep_axis, data_parallelism_axes, weight_gather : deprecated Compatibility arguments converted into a MeshResource and boolean, with a DeprecationWarning. Conflicting old and new arguments raise. From a76e2f9a813cb6d934cb9bb508ce72f226d36a21 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Tue, 6 Oct 2026 13:59:25 -0700 Subject: [PATCH 27/38] Replace JAX MoE WeightGather with a boolean flag Signed-off-by: Jeremy Berchtold --- docs/api/jax.rst | 8 +-- tests/jax/test_moe_resource_api.py | 42 +++++++++--- tests/jax/test_te_ep_moe.py | 19 ------ transformer_engine/jax/flax/moe.py | 8 +-- transformer_engine/jax/moe.py | 100 ++++------------------------- 5 files changed, 51 insertions(+), 126 deletions(-) diff --git a/docs/api/jax.rst b/docs/api/jax.rst index 5370ba8a9fe..a296727125b 100644 --- a/docs/api/jax.rst +++ b/docs/api/jax.rst @@ -98,8 +98,6 @@ gathering concatenates the quantized matrices and their existing scale blocks; matrix-shard gathering reconstructs each expert's scale blocks. These specs are resolved at trace time and do not require JAX Explicit sharding mode. -The old ``ep_axis``, ``data_parallelism_axes`` and ``weight_gather`` arguments -remain accepted with a ``DeprecationWarning``. They are translated into a -resource and the boolean before calling the new API. Conflicting old and -new arguments raise ``ValueError``. ``WeightGather`` remains available only -for this compatibility path; new callers should use the boolean. +The old ``ep_axis`` and ``data_parallelism_axes`` arguments remain accepted +with a ``DeprecationWarning``. They are translated into a resource before +calling the new API. Conflicting old and new arguments raise ``ValueError``. diff --git a/tests/jax/test_moe_resource_api.py b/tests/jax/test_moe_resource_api.py index 4265319d532..65a138054e4 100644 --- a/tests/jax/test_moe_resource_api.py +++ b/tests/jax/test_moe_resource_api.py @@ -15,7 +15,7 @@ from jax.sharding import Mesh from transformer_engine.jax.flax import _MoEBlock -from transformer_engine.jax.moe import WeightGather, _moe_mesh_axes, _resolve_moe_mesh_resource +from transformer_engine.jax.moe import _moe_mesh_axes, _resolve_moe_mesh_resource from transformer_engine.jax.sharding import MeshResource, global_mesh_resource, global_shard_guard @@ -87,11 +87,13 @@ def test_bool_and_resource_types(): def test_legacy_preserves_arbitrary_outer_axis_order(): - with pytest.warns(DeprecationWarning): + with global_shard_guard(MeshResource(fsdp_resource="fsdp", ep_resource="ep")), pytest.warns( + DeprecationWarning + ): resource, quantize = _resolve_moe_mesh_resource( ep_axis="ep", data_parallelism_axes=("outer", "fsdp", "replica"), - weight_gather=WeightGather.quantized(axis="fsdp"), + quant_before_fsdp_ag=True, ) assert _moe_mesh_axes(resource) == ("ep", ("outer", "fsdp", "replica")) assert resource.fsdp_resource == "fsdp" @@ -110,7 +112,6 @@ def test_legacy_defaults_to_no_outer_axes(): [ ({"ep_axis": "other"}, "ep_axis conflicts"), ({"data_parallelism_axes": ("other",)}, "data_parallelism_axes conflicts"), - ({"weight_gather": WeightGather.quantized(axis="other")}, "weight_gather axis conflicts"), ], ) def test_conflicting_old_and_new_args(kwargs, error): @@ -122,7 +123,8 @@ def test_conflicting_old_and_new_args(kwargs, error): @pytest.mark.parametrize("legacy", [False, True]) -def test_public_api_delegates_with_selected_resource(monkeypatch, legacy): +@pytest.mark.parametrize("quantize", [False, True]) +def test_public_api_delegates_with_selected_resource(monkeypatch, legacy, quantize): module = importlib.import_module("transformer_engine.jax.moe") original_moe = module.moe signature = inspect.signature(module._moe) @@ -146,12 +148,12 @@ def delegated_moe(*args, **kwargs): kwargs.update( ep_axis="ep", data_parallelism_axes=("dp", "fsdp"), - weight_gather=WeightGather.quantized(axis="fsdp"), + quant_before_fsdp_ag=quantize, ) else: kwargs.update( mesh_resource=MeshResource(dp_resource="dp", fsdp_resource="fsdp", ep_resource="ep"), - quant_before_fsdp_ag=True, + quant_before_fsdp_ag=quantize, ) with jax.set_mesh(mesh), warnings.catch_warnings(record=True) as recorded: warnings.simplefilter("always", DeprecationWarning) @@ -164,7 +166,31 @@ def delegated_moe(*args, **kwargs): ) assert any("deprecated for TE MoE" in str(w.message) for w in recorded) == legacy assert _moe_mesh_axes(captured["mesh_resource"]) == ("ep", ("dp", "fsdp")) - assert captured["quant_before_fsdp_ag"] is True + assert captured["quant_before_fsdp_ag"] is quantize assert "ep_axis" not in captured with pytest.raises(AssertionError, match="Global mesh resource is not set"): global_mesh_resource() + + +@pytest.mark.parametrize("quantize", [False, True]) +def test_flax_block_forwards_boolean(monkeypatch, quantize): + module = importlib.import_module("transformer_engine.jax.flax.moe") + resource = MeshResource(fsdp_resource="fsdp", ep_resource="ep") + captured = {} + + def fake_moe(inputs, *args, **kwargs): + captured.update(kwargs) + return inputs, None, jnp.zeros((1,), jnp.int32) + + monkeypatch.setattr(module, "moe", fake_moe) + mesh = Mesh(np.asarray(jax.devices()[:1]).reshape(1, 1), ("fsdp", "ep")) + block = _MoEBlock( + num_experts=2, + intermediate_size=4, + mesh_resource=resource, + quant_before_fsdp_ag=quantize, + ) + with jax.set_mesh(mesh): + block.init(jax.random.PRNGKey(0), jnp.ones((1, 1, 4))) + assert captured["quant_before_fsdp_ag"] is quantize + assert _moe_mesh_axes(captured["mesh_resource"]) == ("ep", ("fsdp",)) diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index 3eb6b4a170f..2ba7f78b382 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -133,7 +133,6 @@ def _read_cudnn_fusion_option() -> bool: from transformer_engine.jax.moe import ( _ALIGN_SIZE, _CUDNN_JAX_ALIGN_SIZE, - WeightGather, get_moe_recv_capacity_per_rank, moe, record_ep_bootstrap_signature_for_moe, @@ -586,18 +585,6 @@ def _quantization_recipe(quantization): _QUANTIZATION_CASES.append(pytest.param("mxfp8", id="mxfp8")) -def test_weight_gather_policy_axis_resolution(mesh): - """The policy resolves FSDP context only when no axis was supplied.""" - with pytest.raises(ValueError, match="active MeshResource"): - WeightGather.quantized() - with _ctx(mesh): - assert WeightGather.quantized().axis == FSDP_AXIS - with global_shard_guard(MeshResource()): - assert WeightGather.quantized(axis="custom").axis == "custom" - with pytest.raises(ValueError, match="fsdp_resource"): - WeightGather.quantized() - - @pytest.mark.parametrize("native_weight_layout", [False, True]) @pytest.mark.parametrize("fsdp_dimension", ["k", "gated", "expert"]) def test_quantized_weight_gather_matches_full_precision_gather( @@ -680,12 +667,6 @@ def test_mesh_resource_api_forward_and_backward(mesh, quant_before_fsdp_ag, api) else: candidate = baseline.clone( data_parallelism_axes=(FSDP_AXIS,), - quant_before_fsdp_ag=False, - weight_gather=( - WeightGather.quantized(axis=FSDP_AXIS) - if quant_before_fsdp_ag - else WeightGather.full_precision() - ), ) x = _make_inputs(jax.random.PRNGKey(51)) variables, baseline_out, _ = _init_apply(baseline, mesh, x, jax.random.PRNGKey(52)) diff --git a/transformer_engine/jax/flax/moe.py b/transformer_engine/jax/flax/moe.py index d65116c903f..5528185bbd6 100644 --- a/transformer_engine/jax/flax/moe.py +++ b/transformer_engine/jax/flax/moe.py @@ -34,7 +34,7 @@ from flax import linen as nn from transformer_engine.common.recipe import Recipe -from ..moe import WeightGather, _moe_mesh_axes, _resolve_moe_mesh_resource, moe +from ..moe import _moe_mesh_axes, _resolve_moe_mesh_resource, moe from ..quantize import QuantizerSet from ..router import ScoreFunction from ..sharding import MeshResource, _get_mesh, global_shard_guard @@ -100,8 +100,8 @@ class _MoEBlock(TransformerEngineBase): quant_before_fsdp_ag : bool Quantize MXFP8 weight shards before their FSDP all-gather. Defaults to ``False``; ``True`` requires ``mesh_resource.fsdp_resource``. - ep_axis, data_parallelism_axes, weight_gather : deprecated - Compatibility arguments converted into MeshResource and the boolean + ep_axis, data_parallelism_axes : deprecated + Compatibility axis arguments converted into MeshResource with a DeprecationWarning. use_cudnn_fusion : bool Defaults to ``True``: try Rubin fused GLU, then generic Blackwell+ fused @@ -168,7 +168,6 @@ class _MoEBlock(TransformerEngineBase): apply_topk_weights_early: bool = False recv_capacity_per_rank: Optional[int] = None dispatch_checkpoint_name: Optional[str] = None - weight_gather: Optional[WeightGather] = None # Dtypes / init / misc dtype: DType = jnp.float32 @@ -214,7 +213,6 @@ def __call__(self, inputs: Array) -> Tuple[Array, Optional[Array], Array]: self.quant_before_fsdp_ag, self.ep_axis, self.data_parallelism_axes, - self.weight_gather, ) _, data_parallelism_axes = _moe_mesh_axes(mesh_resource) assert ( diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index d51b2deed91..92161581199 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -34,7 +34,7 @@ import warnings from dataclasses import dataclass, fields, replace from functools import partial -from typing import Any, Literal, Optional, Tuple, Union +from typing import Any, Optional, Tuple, Union import flax.struct import flax.linen as nn @@ -61,56 +61,7 @@ from .router import ScoreFunction, _validate_score_function from .sharding import MeshResource, _get_mesh, global_mesh_resource, global_shard_guard -__all__ = ["WeightGather", "get_moe_recv_capacity_per_rank", "moe"] - - -@dataclass(frozen=True) -class WeightGather: - """Deprecated compatibility policy; use ``quant_before_fsdp_ag`` instead. - - How MoE expert weights are gathered across a sharding mesh axis. - - The quantization recipe determines the wire format for ``quantized``; - this policy only chooses whether quantization precedes the gather. - """ - - mode: Literal["full_precision", "quantized"] = "full_precision" - axis: Optional[str] = None - - def __post_init__(self): - if self.mode not in ("full_precision", "quantized"): - raise ValueError(f"Unsupported weight gather mode: {self.mode!r}") - if self.mode == "quantized" and not self.axis: - raise ValueError("Quantized weight gather requires a mesh axis.") - if self.mode == "full_precision" and self.axis is not None: - raise ValueError("A weight gather axis is only used in quantized mode.") - - @classmethod - def full_precision(cls) -> "WeightGather": - """Gather weights before quantization (the default behavior).""" - return cls() - - @classmethod - def quantized(cls, *, axis: Optional[str] = None) -> "WeightGather": - """Quantize local shards before gathering data and scales. - - With no explicit axis, use the active ``MeshResource.fsdp_resource``. - Passing an axis bypasses the global resource lookup entirely. - """ - if axis is None: - try: - axis = global_mesh_resource().fsdp_resource - except AssertionError as exc: - raise ValueError( - "WeightGather.quantized() requires an active MeshResource " - "with fsdp_resource, or an explicit axis." - ) from exc - if not axis: - raise ValueError( - "WeightGather.quantized() requires MeshResource.fsdp_resource " - "or an explicit axis." - ) - return cls(mode="quantized", axis=axis) +__all__ = ["get_moe_recv_capacity_per_rank", "moe"] @dataclass @@ -144,9 +95,8 @@ def _resolve_moe_mesh_resource( quant_before_fsdp_ag=False, ep_axis=None, data_parallelism_axes=None, - weight_gather=None, ): - """Resolve the canonical API and adapt deprecated axes/gather arguments.""" + """Resolve the canonical API and adapt deprecated axis arguments.""" if not isinstance(quant_before_fsdp_ag, bool): raise TypeError("quant_before_fsdp_ag must be a bool.") if mesh_resource is not None and not isinstance(mesh_resource, MeshResource): @@ -158,20 +108,14 @@ def _resolve_moe_mesh_resource( except AssertionError: mesh_resource = None - legacy = ep_axis is not None or data_parallelism_axes is not None or weight_gather is not None + legacy = ep_axis is not None or data_parallelism_axes is not None if legacy: warnings.warn( - "ep_axis, data_parallelism_axes, and weight_gather are deprecated for TE MoE; " + "ep_axis and data_parallelism_axes are deprecated for TE MoE; " "pass mesh_resource=MeshResource(...) and quant_before_fsdp_ag instead.", DeprecationWarning, stacklevel=3, ) - if weight_gather is not None and not isinstance(weight_gather, WeightGather): - raise TypeError("weight_gather must be a WeightGather policy.") - if weight_gather is not None: - if quant_before_fsdp_ag and weight_gather.mode != "quantized": - raise ValueError("quant_before_fsdp_ag conflicts with weight_gather.") - quant_before_fsdp_ag = weight_gather.mode == "quantized" if explicit_resource: if ep_axis is not None and ep_axis != mesh_resource.ep_resource: raise ValueError("ep_axis conflicts with mesh_resource.ep_resource.") @@ -180,28 +124,10 @@ def _resolve_moe_mesh_resource( and tuple(data_parallelism_axes) != _moe_mesh_axes(mesh_resource)[1] ): raise ValueError("data_parallelism_axes conflicts with mesh_resource.") - if ( - weight_gather is not None - and weight_gather.axis is not None - and weight_gather.axis != mesh_resource.fsdp_resource - ): - raise ValueError("weight_gather axis conflicts with mesh_resource.fsdp_resource.") - elif ep_axis is not None or data_parallelism_axes is not None: + else: # The old functional API defaulted to no outer axes even in a global context. axes = tuple(data_parallelism_axes or ()) - fsdp_axis = ( - weight_gather.axis - if weight_gather is not None and weight_gather.axis is not None - else getattr(mesh_resource, "fsdp_resource", None) - ) - if ( - weight_gather is not None - and weight_gather.axis is not None - and weight_gather.axis not in axes - ): - raise ValueError( - "Quantized weight all-gather requires its FSDP axis among the outer batch axes." - ) + fsdp_axis = getattr(mesh_resource, "fsdp_resource", None) if fsdp_axis not in axes: fsdp_axis = axes[-1] if axes else None dp_axes = tuple(axis for axis in axes if axis != fsdp_axis) @@ -220,8 +146,6 @@ def _resolve_moe_mesh_resource( **resources, _legacy_data_parallelism_axes=axes, ) - elif weight_gather.axis is not None and mesh_resource is not None: - mesh_resource = replace(mesh_resource, fsdp_resource=weight_gather.axis) if mesh_resource is None: raise ValueError( @@ -2090,7 +2014,6 @@ def moe( wi_1_checkpoint_name: Optional[str] = None, wo_checkpoint_name: Optional[str] = None, dispatch_checkpoint_name: Optional[str] = None, - weight_gather: Optional[WeightGather] = None, use_cudnn_fusion: bool = True, cudnn_native_weight_layout: Optional[bool] = None, ) -> Tuple[jnp.ndarray, Optional[jnp.ndarray], jnp.ndarray]: @@ -2157,8 +2080,8 @@ def moe( FSDP shards the gated dimension, layout conversion permutes gathered FP8 data and inverse scales, preserving global gate/up pairing without gathering or requantizing full-precision weights. - ep_axis, data_parallelism_axes, weight_gather : deprecated - Compatibility arguments converted into a MeshResource and boolean, + ep_axis, data_parallelism_axes : deprecated + Compatibility axis arguments converted into a MeshResource, with a DeprecationWarning. Conflicting old and new arguments raise. Per-expert dispatch-slot alignment defaults to 128 tokens (``_ALIGN_SIZE``). Requesting cuDNN fusion reserves 256 tokens, also when falling back, to @@ -2188,16 +2111,15 @@ def moe( """ if not isinstance(use_cudnn_fusion, bool): raise TypeError("use_cudnn_fusion must be a bool") - if ep_axis is not None or data_parallelism_axes is not None or weight_gather is not None: + if ep_axis is not None or data_parallelism_axes is not None: call_args = locals().copy() resource, quantize = _resolve_moe_mesh_resource( mesh_resource, quant_before_fsdp_ag, ep_axis, data_parallelism_axes, - weight_gather, ) - for name in ("ep_axis", "data_parallelism_axes", "weight_gather"): + for name in ("ep_axis", "data_parallelism_axes"): call_args.pop(name) call_args.update(mesh_resource=resource, quant_before_fsdp_ag=quantize) return moe(**call_args) From c7c2651c2d01e17383fc29828b158980de6e7e48 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Tue, 6 Oct 2026 14:01:11 -0700 Subject: [PATCH 28/38] Remove documentation for unused JAX MoE fusion environment variable Signed-off-by: Jeremy Berchtold --- docs/envvars.rst | 11 ----------- 1 file changed, 11 deletions(-) diff --git a/docs/envvars.rst b/docs/envvars.rst index 056938301f2..46b70bbe46a 100644 --- a/docs/envvars.rst +++ b/docs/envvars.rst @@ -556,17 +556,6 @@ JAX-Specific Variables :Default: None :Description: Test level for JAX unit tests (``"L0"``, ``"L1"``, ``"L2"``). Used internally by the test suite. -.. envvar:: NVTE_JAX_MOE_USE_CUDNN_CUTEDSL_FUSION - - :Type: ``int`` (0 or 1) - :Default: ``0`` - :Description: **(JAX only)** Enable the experimental cuDNN frontend JAX fusion - for MXFP8 MoE FC1 grouped GEMM, SwiGLU, and grouped quantization. Forward - uses cuDNN's dedicated ``cudnn.jax.call`` API; backward uses TE's regular - MXFP8 grouped-GEMM path. Explicit opt-in requires an eligible SM100 SwiGLU - MXFP8 MoE call and cuDNN frontend with the grouped SwiGLU JAX entry point; - unsupported calls warn with the full validation reason list and fall back. - JAX Triton Extensions ^^^^^^^^^^^^^^^^^^^^^ From c31ddb0bebabe24de20e8a349ac5d1d145ee223c Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Tue, 6 Oct 2026 14:06:05 -0700 Subject: [PATCH 29/38] Remove obsolete CuTeDSL MoE Shardy reproducer Signed-off-by: Jeremy Berchtold --- tests/jax/repro_cutedsl_moe_shardy.py | 124 -------------------------- 1 file changed, 124 deletions(-) delete mode 100644 tests/jax/repro_cutedsl_moe_shardy.py diff --git a/tests/jax/repro_cutedsl_moe_shardy.py b/tests/jax/repro_cutedsl_moe_shardy.py deleted file mode 100644 index 0821116adc4..00000000000 --- a/tests/jax/repro_cutedsl_moe_shardy.py +++ /dev/null @@ -1,124 +0,0 @@ -# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# -# See LICENSE for license information. -"""Shape-only reproducer for nvbug 6432162 and CuTeDSL MoE outputs. - -This intentionally does not call the kernel. It isolates Shardy propagation -over the exact five output shapes used by the production EP2/FSDP2 slice. -""" - -import math - -import jax -import jax.numpy as jnp -import numpy as np -from jax.experimental.custom_partitioning import ( - CompoundFactor, - SdyShardingRule, - custom_partitioning, -) -from jax.sharding import Mesh, NamedSharding, PartitionSpec as P - - -EXPERTS = 16 -M = 266240 -HIDDEN = 1792 -INTERMEDIATE = 2048 -PAIR = 2 - - -def _ceil_div(x, y): - return (x + y - 1) // y - - -ROW_SCALE_SIZE = 32 * 4 * _ceil_div(M, 128) * 4 * _ceil_div(_ceil_div(INTERMEDIATE, 32), 4) -COL_SCALE_SIZE = 32 * 4 * _ceil_div(INTERMEDIATE, 128) * 4 * _ceil_div(_ceil_div(M, 32), 4) - - -def _shape_only_impl(a, b, sfa, sfb, padded_offsets, prob): - del a, b, sfa, sfb, padded_offsets, prob - return ( - jnp.zeros((M, PAIR * INTERMEDIATE, 1), dtype=jnp.bfloat16), - jnp.zeros((M, INTERMEDIATE, 1), dtype=jnp.float8_e4m3fn), - jnp.zeros((M, INTERMEDIATE, 1), dtype=jnp.float8_e4m3fn), - jnp.zeros((ROW_SCALE_SIZE,), dtype=jnp.float8_e8m0fnu), - jnp.zeros((COL_SCALE_SIZE,), dtype=jnp.float8_e8m0fnu), - ) - - -_shape_only = custom_partitioning(_shape_only_impl) - - -def _partition(mesh, arg_infos, result_infos): - arg_shardings = tuple(info.sharding for info in arg_infos) - result_shardings = tuple(info.sharding for info in result_infos) - return mesh, _shape_only_impl, result_shardings, arg_shardings - - -def _shardy_rule(mesh, value_types, result_types): - del mesh, value_types, result_types - combined = CompoundFactor("intermediate", "swiglu_pair") - return SdyShardingRule( - ( - ("tokens", "hidden", "a_l"), - ("experts", combined, "hidden"), - ("sfa_flat",), - ("sfb_flat",), - ("experts",), - ("tokens", "prob_n", "prob_l"), - ), - ( - ("tokens", combined, "combined_l"), - ("tokens", "intermediate", "row_l"), - ("tokens", "intermediate", "col_l"), - ("row_scale_flat",), - ("col_scale_flat",), - ), - swiglu_pair=PAIR, - ) - - -_shape_only.def_partition(partition=_partition, sharding_rule=_shardy_rule) - - -def main() -> None: - jax.config.update("jax_use_shardy_partitioner", True) - devices = np.asarray(jax.devices()) - if devices.size < 2: - raise RuntimeError("This Shardy reproducer requires at least two JAX devices") - ep = 2 - dp = devices.size // ep - devices = devices[: dp * ep].reshape(dp, ep) - mesh = Mesh(devices, ("data", "expert")) - token_sharding = NamedSharding(mesh, P(("data", "expert"), None, None)) - expert_sharding = NamedSharding(mesh, P("expert", None, None)) - replicated = NamedSharding(mesh, P()) - - args = ( - jax.ShapeDtypeStruct((M, HIDDEN, 1), jnp.float8_e4m3fn, sharding=token_sharding), - jax.ShapeDtypeStruct( - (EXPERTS, PAIR * INTERMEDIATE, HIDDEN), - jnp.float8_e4m3fn, - sharding=expert_sharding, - ), - jax.ShapeDtypeStruct((1,), jnp.float8_e8m0fnu, sharding=replicated), - jax.ShapeDtypeStruct((1,), jnp.float8_e8m0fnu, sharding=replicated), - jax.ShapeDtypeStruct( - (EXPERTS,), - jnp.int32, - sharding=NamedSharding(mesh, P("expert")), - ), - jax.ShapeDtypeStruct((M, 1, 1), jnp.float32, sharding=token_sharding), - ) - with jax.set_mesh(mesh): - lowered = jax.jit(_shape_only).lower(*args) - print( - "Shardy lowering passed for CuTeDSL MoE outputs: " - f"M={M}, I={INTERMEDIATE}, row_scale={ROW_SCALE_SIZE}, " - f"col_scale={COL_SCALE_SIZE}, devices={math.prod(mesh.devices.shape)}" - ) - print(lowered.compiler_ir(dialect="stablehlo")) - - -if __name__ == "__main__": - main() From fd250a2236f21c418cb3d94e0a3db332294471d5 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Tue, 6 Oct 2026 14:08:08 -0700 Subject: [PATCH 30/38] Remove unused CuTeDSL JAX requirements file Signed-off-by: Jeremy Berchtold --- tests/jax/requirements_cutedsl.txt | 6 ------ 1 file changed, 6 deletions(-) delete mode 100644 tests/jax/requirements_cutedsl.txt diff --git a/tests/jax/requirements_cutedsl.txt b/tests/jax/requirements_cutedsl.txt deleted file mode 100644 index 3f32f0a51c8..00000000000 --- a/tests/jax/requirements_cutedsl.txt +++ /dev/null @@ -1,6 +0,0 @@ -# Dependencies for the cuDNN JAX grouped-SwiGLU MoE path. -jax[cuda13]>=0.9.1 -nvidia-cutlass-dsl[cu13]>=4.6.2 - -# Install the sibling cuDNN-FE checkout separately: -# pip install ../cudnn-frontend From 1949887f5eea5ac49288e68ad01044578ca78f08 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Tue, 6 Oct 2026 14:12:51 -0700 Subject: [PATCH 31/38] Consolidate SwiGLU fallback tests into custom call compute suite Signed-off-by: Jeremy Berchtold --- tests/jax/test_custom_call_compute.py | 349 +++++++++++++++++ .../jax/test_grouped_gemm_swiglu_fallback.py | 353 ------------------ 2 files changed, 349 insertions(+), 353 deletions(-) delete mode 100644 tests/jax/test_grouped_gemm_swiglu_fallback.py diff --git a/tests/jax/test_custom_call_compute.py b/tests/jax/test_custom_call_compute.py index 9c3b489c1fc..2c050f521d7 100644 --- a/tests/jax/test_custom_call_compute.py +++ b/tests/jax/test_custom_call_compute.py @@ -2,7 +2,13 @@ # # See LICENSE for license information. +import importlib +import inspect +import warnings +from types import SimpleNamespace + import jax +import numpy as np import jax.numpy as jnp import pytest from jax import jit, value_and_grad @@ -49,6 +55,12 @@ from transformer_engine.jax.dense import dense, grouped_dense from transformer_engine.jax.layernorm_dense import layernorm_dense from transformer_engine.jax.cpp_extensions.topk import topk +import transformer_engine_jax + +_swiglu_adapter = importlib.import_module( + "transformer_engine.jax.cpp_extensions.grouped_gemm_swiglu" +) +_moe_module = importlib.import_module("transformer_engine.jax.moe") GEMM_CASES = [ (256, 256, 512), @@ -2128,3 +2140,340 @@ def test_topk_2d(self, dtype, problem_size): prim_gathered = jnp.take_along_axis(x, prim_idx, axis=1) ref_gathered = jnp.take_along_axis(x, ref_idx, axis=1) assert_allclose(prim_gathered, ref_gathered, dtype=dtype) + + +class TestGroupedGemmSwigluFallback: + """MoE API compatibility and dispatch with mocked kernels on one JAX device.""" + + @staticmethod + def _api(arguments, change=None): + def function(**kwargs): + raise AssertionError("Probes must not execute kernels") + + parameters = [ + inspect.Parameter( + name, + ( + inspect.Parameter.POSITIONAL_ONLY + if change == "positional" and name == "a_tensor" + else inspect.Parameter.POSITIONAL_OR_KEYWORD + ), + ) + for name in arguments + if not (change == "renamed" and name == "discrete_col_sfd") + ] + if change in ("optional", "mandatory"): + parameters.append( + inspect.Parameter( + "new_arg", + inspect.Parameter.KEYWORD_ONLY, + default=None if change == "optional" else inspect.Parameter.empty, + ) + ) + function.__signature__ = inspect.Signature(parameters) + return function + + @pytest.fixture + def frontend(self, monkeypatch): + module = SimpleNamespace( + grouped_gemm_glu=self._api(_swiglu_adapter._FORWARD_ARGS), + grouped_gemm_swiglu=self._api(_swiglu_adapter._FORWARD_ARGS), + grouped_gemm_dswiglu=self._api(_swiglu_adapter._BACKWARD_ARGS), + ) + original = importlib.import_module + + def import_module(name, package=None): + if name == "cudnn.jax": + return module + if name == "cutlass.jax": + return SimpleNamespace(is_available=lambda: True) + return original(name, package) + + monkeypatch.setattr(_swiglu_adapter.importlib, "import_module", import_module) + return module + + @pytest.mark.parametrize("rubin", [False, True]) + @pytest.mark.parametrize("operation", ["forward", "backward"]) + @pytest.mark.parametrize( + "change", + [ + "missing", + "not_callable", + "renamed", + "mandatory", + "optional", + "positional", + "opaque", + ], + ) + def test_api_signature(self, frontend, rubin, operation, change): + name = ( + ("grouped_gemm_glu" if rubin else "grouped_gemm_swiglu") + if operation == "forward" + else "grouped_gemm_dswiglu" + ) + args = ( + _swiglu_adapter._FORWARD_ARGS + if operation == "forward" + else _swiglu_adapter._BACKWARD_ARGS + ) + if change == "missing": + delattr(frontend, name) + elif change == "not_callable": + setattr(frontend, name, 42) + elif change == "opaque": + setattr(frontend, name, lambda **kwargs: None) + else: + setattr(frontend, name, self._api(args, change)) + available, reason = _swiglu_adapter.grouped_gemm_swiglu_dependencies_available(rubin) + assert available == (change == "optional") + assert bool(reason) == (change != "optional") + + @pytest.mark.parametrize("capability", [90, 100, 103, 107, 120]) + @pytest.mark.parametrize( + "missing", + [None, "grouped_gemm_glu", "grouped_gemm_swiglu", "grouped_gemm_dswiglu", "all"], + ) + def test_ordered_fallback(self, frontend, monkeypatch, capability, missing): + if missing == "all": + vars(frontend).clear() + elif missing: + delattr(frontend, missing) + monkeypatch.setattr( + transformer_engine_jax, "get_device_compute_capability", lambda _: capability + ) + expected = False + if capability >= 100 and missing not in ("grouped_gemm_dswiglu", "all"): + if capability == 107 and missing != "grouped_gemm_glu": + expected = "rubin" + elif missing != "grouped_gemm_swiglu": + expected = "blackwell" + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + assert _moe_module._select_cudnn_jax_fusion([]) == expected + if expected == "rubin": + assert not caught + else: + assert len(caught) == 1 + message = str(caught[0].message) + assert "falling back" in message and "1.31.0" in message + assert ("generic Blackwell+ fused" if expected else "unfused TE") in message + + @pytest.mark.parametrize("name", ["cudnn.jax", "cutlass.jax"]) + def test_missing_dependency(self, monkeypatch, name): + original = importlib.import_module + + def import_module(module, package=None): + if module == name: + raise ModuleNotFoundError(f"No module named {module}") + return original(module, package) + + monkeypatch.setattr(_swiglu_adapter.importlib, "import_module", import_module) + available, reason = _swiglu_adapter.grouped_gemm_swiglu_dependencies_available() + assert not available and name in reason + + def test_ineligible_call_and_device_query_failure(self, frontend, monkeypatch): + monkeypatch.setattr(transformer_engine_jax, "get_device_compute_capability", lambda _: 107) + with pytest.warns(UserWarning, match="bias unsupported"): + assert _moe_module._select_cudnn_jax_fusion(["bias unsupported"]) is False + + def fail(_): + raise RuntimeError("device query failed") + + monkeypatch.setattr(transformer_engine_jax, "get_device_compute_capability", fail) + with pytest.warns(UserWarning, match="device query failed"): + assert _moe_module._select_cudnn_jax_fusion([]) is False + + def test_installed_frontend_contract(self): + for rubin in (False, True): + available, reason = _swiglu_adapter.grouped_gemm_swiglu_dependencies_available(rubin) + assert available or reason + + @pytest.mark.parametrize("request_fusion", [None, True, False]) + @pytest.mark.parametrize("native_layout", [False, True]) + @pytest.mark.parametrize("square", [False, True]) + @pytest.mark.parametrize( + "missing,expected", + [ + (None, "rubin"), + ("grouped_gemm_glu", "blackwell"), + ("grouped_gemm_dswiglu", False), + ], + ) + def test_public_moe_passes_selected_path( + self, frontend, monkeypatch, native_layout, square, missing, expected, request_fusion + ): + from jax.sharding import Mesh + from transformer_engine.jax.sharding import MeshResource + + if missing: + delattr(frontend, missing) + monkeypatch.setattr(transformer_engine_jax, "get_device_compute_capability", lambda _: 107) + monkeypatch.setattr( + _moe_module, "_cudnn_jax_fusion_rejection_reasons", lambda *args, **kwargs: [] + ) + mesh = Mesh(np.asarray(jax.devices()[:1]), ("ep",)) + monkeypatch.setattr(_moe_module, "_get_mesh", lambda: mesh) + monkeypatch.setattr(_moe_module, "_with_sharding_constraint_cast_bwd", lambda x, _: x) + received = [] + signature = inspect.signature(_moe_module._moe) + + def execute(*args): + bound = signature.bind(*args).arguments + assert bound["use_cudnn_fusion"] is (request_fusion is not False) + assert bound["cudnn_native_weight_layout"] is native_layout + received.append(bound["use_cudnn_jax_fusion"]) + return args[0], None, jnp.asarray(0) + + monkeypatch.setattr(_moe_module, "_moe", execute) + x = jnp.ones((1, 1, 128), jnp.bfloat16) + wi = jnp.ones((2, 256, 128) if native_layout else (2, 128, 256), jnp.bfloat16) + kwargs = {} if request_fusion is None else {"use_cudnn_fusion": request_fusion} + if square: + wi = jnp.ones((2, 128, 128), jnp.bfloat16) + kwargs["cudnn_native_weight_layout"] = native_layout + if request_fusion is False: + expected = False + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + output, _, _ = _moe_module.moe( + x, + jnp.ones((128, 2)), + wi, + jnp.ones((2, 128, 128)), + num_experts=2, + num_experts_per_tok=1, + mesh_resource=MeshResource(ep_resource="ep"), + **kwargs, + ) + assert received == [expected] + assert output is x + assert len(caught) == (0 if request_fusion is False or expected == "rubin" else 1) + + @pytest.mark.parametrize("path", ["rubin", "blackwell"]) + @pytest.mark.parametrize("native_layout", [False, True]) + @pytest.mark.parametrize("gather_gated_dimension", [False, True]) + def test_forward_uses_selected_kernel( + self, monkeypatch, path, native_layout, gather_gated_dimension + ): + + class SelectedKernel(Exception): + pass + + def selected(*args, **kwargs): + assert args[1].shape == (1, 256, 128) + raise SelectedKernel + + def rejected(*args, **kwargs): + raise AssertionError("Forward ignored the selected fallback") + + monkeypatch.setattr( + _moe_module.tex, "grouped_gemm_glu", selected if path == "rubin" else rejected + ) + monkeypatch.setattr( + _moe_module.tex, "grouped_gemm_swiglu", selected if path == "blackwell" else rejected + ) + + def quantize(data, *args, **kwargs): + return SimpleNamespace( + get_tensor=lambda **kwargs: SimpleNamespace( + data=data, scale_inv=jnp.ones((1,), dtype=jnp.uint8) + ) + ) + + monkeypatch.setattr(_moe_module.tex, "grouped_quantize", quantize) + + def gather(tensor, fsdp_axis, fsdp_size, sharded_axis): + data = tensor.get_tensor().data + assert sharded_axis == (1 if native_layout else 2) + return quantize(jnp.concatenate([data, data], axis=sharded_axis)) + + monkeypatch.setattr(_moe_module, "_gather_quantized_weight", gather) + + def reorder(tensor, *, interleave): + assert interleave and gather_gated_dimension and not native_layout + gate, up = jnp.split(tensor.get_tensor().data, 2, axis=-1) + return quantize(_moe_module.tex.pack_swiglu_pair(gate, up)) + + monkeypatch.setattr(_moe_module, "_reorder_quantized_swiglu_weight", reorder) + quantizers = SimpleNamespace(x=SimpleNamespace(q_dtype=jnp.float8_e4m3fn), kernel=None) + kwargs = dict.fromkeys(inspect.signature(_moe_module._ffn_fwd_per_shard).parameters) + combined = 128 if gather_gated_dimension else 256 + kwargs.update( + recv_tokens_local=jnp.ones((1, 256, 128), jnp.bfloat16), + recv_topk_weights_local=jnp.ones((1, 256)), + token_counts_local=jnp.asarray([256]), + wi=jnp.ones((1, combined, 128) if native_layout else (1, 128, combined), jnp.bfloat16), + wo=jnp.ones((1, 128, 128), jnp.bfloat16), + quantizer_sets=(quantizers, quantizers), + num_local_experts=1, + use_cudnn_jax_fusion=path, + cudnn_native_weight_layout=native_layout, + quant_before_fsdp_ag=gather_gated_dimension, + wi_fsdp_axis=(1 if native_layout else 2), + fsdp_axis="fsdp", + fsdp_size=2, + ) + with pytest.raises(SelectedKernel): + _moe_module._ffn_fwd_per_shard(**kwargs) + + @pytest.mark.parametrize("request_fusion", [None, True, False]) + def test_flax_forwards_fusion_bool(self, monkeypatch, request_fusion): + from transformer_engine.jax.flax import _MoEBlock + from transformer_engine.jax.sharding import MeshResource + + flax_moe = importlib.import_module("transformer_engine.jax.flax.moe") + received = [] + + def execute(inputs, *args, **kwargs): + received.append(kwargs["use_cudnn_fusion"]) + return inputs, None, jnp.asarray(0) + + monkeypatch.setattr(flax_moe, "moe", execute) + kwargs = {} if request_fusion is None else {"use_cudnn_fusion": request_fusion} + block = _MoEBlock( + num_experts=2, + intermediate_size=32, + mesh_resource=MeshResource(ep_resource="ep"), + **kwargs, + ) + block.init(jax.random.PRNGKey(0), jnp.ones((1, 1, 32))) + assert received == [request_fusion is not False] + + @pytest.mark.parametrize("request_fusion", [True, False]) + def test_capacity_follows_explicit_fusion_bool(self, request_fusion): + kwargs = dict(num_experts=8, num_experts_per_tok=2, max_tokens_per_rank=64, ep_size=2) + alignment = _moe_module._CUDNN_JAX_ALIGN_SIZE if request_fusion else _moe_module._ALIGN_SIZE + expected = _moe_module.get_moe_recv_capacity_per_rank(**kwargs, alignment=alignment) + assert ( + _moe_module.get_moe_recv_capacity_per_rank(**kwargs, use_cudnn_fusion=request_fusion) + == expected + ) + if request_fusion: + assert _moe_module.get_moe_recv_capacity_per_rank(**kwargs) == expected + + @pytest.mark.parametrize("invalid", [0, 1, None, "true"]) + def test_fusion_argument_requires_bool(self, invalid): + with pytest.raises(TypeError, match="use_cudnn_fusion must be a bool"): + _moe_module.moe( + None, + None, + None, + None, + num_experts=2, + num_experts_per_tok=1, + use_cudnn_fusion=invalid, + ) + with pytest.raises(TypeError, match="use_cudnn_fusion must be a bool"): + _moe_module.get_moe_recv_capacity_per_rank( + num_experts=2, + num_experts_per_tok=1, + max_tokens_per_rank=16, + ep_size=1, + use_cudnn_fusion=invalid, + ) + + def test_vjp_bool_is_static_and_defaults_to_true(self): + assert inspect.signature(_moe_module._moe).parameters["use_cudnn_fusion"].default is True + assert inspect.signature(_moe_module.moe).parameters["use_cudnn_fusion"].default is True + assert 32 in _moe_module._moe.nondiff_argnums diff --git a/tests/jax/test_grouped_gemm_swiglu_fallback.py b/tests/jax/test_grouped_gemm_swiglu_fallback.py deleted file mode 100644 index 1d784358ac4..00000000000 --- a/tests/jax/test_grouped_gemm_swiglu_fallback.py +++ /dev/null @@ -1,353 +0,0 @@ -# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# -# See LICENSE for license information. -"""cuDNN MoE API compatibility and ordered dispatch without kernel compilation.""" - -import importlib -import inspect -import warnings -from types import SimpleNamespace - -import pytest -import transformer_engine_jax - -adapter = importlib.import_module("transformer_engine.jax.cpp_extensions.grouped_gemm_swiglu") -moe = importlib.import_module("transformer_engine.jax.moe") - - -def api(arguments, change=None): - def function(**kwargs): - raise AssertionError("Probes must not execute kernels") - - parameters = [ - inspect.Parameter( - name, - ( - inspect.Parameter.POSITIONAL_ONLY - if change == "positional" and name == "a_tensor" - else inspect.Parameter.POSITIONAL_OR_KEYWORD - ), - ) - for name in arguments - if not (change == "renamed" and name == "discrete_col_sfd") - ] - if change in ("optional", "mandatory"): - parameters.append( - inspect.Parameter( - "new_arg", - inspect.Parameter.KEYWORD_ONLY, - default=None if change == "optional" else inspect.Parameter.empty, - ) - ) - function.__signature__ = inspect.Signature(parameters) - return function - - -@pytest.fixture -def frontend(monkeypatch): - module = SimpleNamespace( - grouped_gemm_glu=api(adapter._FORWARD_ARGS), - grouped_gemm_swiglu=api(adapter._FORWARD_ARGS), - grouped_gemm_dswiglu=api(adapter._BACKWARD_ARGS), - ) - original = importlib.import_module - - def import_module(name, package=None): - if name == "cudnn.jax": - return module - if name == "cutlass.jax": - return SimpleNamespace(is_available=lambda: True) - return original(name, package) - - monkeypatch.setattr(adapter.importlib, "import_module", import_module) - return module - - -@pytest.mark.parametrize("rubin", [False, True]) -@pytest.mark.parametrize("operation", ["forward", "backward"]) -@pytest.mark.parametrize( - "change", - [ - "missing", - "not_callable", - "renamed", - "mandatory", - "optional", - "positional", - "opaque", - ], -) -def test_api_signature(frontend, rubin, operation, change): - name = ( - ("grouped_gemm_glu" if rubin else "grouped_gemm_swiglu") - if operation == "forward" - else "grouped_gemm_dswiglu" - ) - args = adapter._FORWARD_ARGS if operation == "forward" else adapter._BACKWARD_ARGS - if change == "missing": - delattr(frontend, name) - elif change == "not_callable": - setattr(frontend, name, 42) - elif change == "opaque": - setattr(frontend, name, lambda **kwargs: None) - else: - setattr(frontend, name, api(args, change)) - available, reason = adapter.grouped_gemm_swiglu_dependencies_available(rubin) - assert available == (change == "optional") - assert bool(reason) == (change != "optional") - - -@pytest.mark.parametrize("capability", [90, 100, 103, 107, 120]) -@pytest.mark.parametrize( - "missing", - [None, "grouped_gemm_glu", "grouped_gemm_swiglu", "grouped_gemm_dswiglu", "all"], -) -def test_ordered_fallback(frontend, monkeypatch, capability, missing): - if missing == "all": - vars(frontend).clear() - elif missing: - delattr(frontend, missing) - monkeypatch.setattr( - transformer_engine_jax, "get_device_compute_capability", lambda _: capability - ) - expected = False - if capability >= 100 and missing not in ("grouped_gemm_dswiglu", "all"): - if capability == 107 and missing != "grouped_gemm_glu": - expected = "rubin" - elif missing != "grouped_gemm_swiglu": - expected = "blackwell" - with warnings.catch_warnings(record=True) as caught: - warnings.simplefilter("always") - assert moe._select_cudnn_jax_fusion([]) == expected - if expected == "rubin": - assert not caught - else: - assert len(caught) == 1 - message = str(caught[0].message) - assert "falling back" in message and "1.31.0" in message - assert ("generic Blackwell+ fused" if expected else "unfused TE") in message - - -@pytest.mark.parametrize("name", ["cudnn.jax", "cutlass.jax"]) -def test_missing_dependency(monkeypatch, name): - original = importlib.import_module - - def import_module(module, package=None): - if module == name: - raise ModuleNotFoundError(f"No module named {module}") - return original(module, package) - - monkeypatch.setattr(adapter.importlib, "import_module", import_module) - available, reason = adapter.grouped_gemm_swiglu_dependencies_available() - assert not available and name in reason - - -def test_ineligible_call_and_device_query_failure(frontend, monkeypatch): - monkeypatch.setattr(transformer_engine_jax, "get_device_compute_capability", lambda _: 107) - with pytest.warns(UserWarning, match="bias unsupported"): - assert moe._select_cudnn_jax_fusion(["bias unsupported"]) is False - - def fail(_): - raise RuntimeError("device query failed") - - monkeypatch.setattr(transformer_engine_jax, "get_device_compute_capability", fail) - with pytest.warns(UserWarning, match="device query failed"): - assert moe._select_cudnn_jax_fusion([]) is False - - -def test_installed_frontend_contract(): - for rubin in (False, True): - available, reason = adapter.grouped_gemm_swiglu_dependencies_available(rubin) - assert available or reason - - -@pytest.mark.parametrize("request_fusion", [None, True, False]) -@pytest.mark.parametrize("native_layout", [False, True]) -@pytest.mark.parametrize("square", [False, True]) -@pytest.mark.parametrize( - "missing,expected", - [ - (None, "rubin"), - ("grouped_gemm_glu", "blackwell"), - ("grouped_gemm_dswiglu", False), - ], -) -def test_public_moe_passes_selected_path( - frontend, monkeypatch, native_layout, square, missing, expected, request_fusion -): - import jax - import jax.numpy as jnp - import numpy as np - from jax.sharding import Mesh - from transformer_engine.jax.sharding import MeshResource - - if missing: - delattr(frontend, missing) - monkeypatch.setattr(transformer_engine_jax, "get_device_compute_capability", lambda _: 107) - monkeypatch.setattr(moe, "_cudnn_jax_fusion_rejection_reasons", lambda *args, **kwargs: []) - mesh = Mesh(np.asarray(jax.devices()[:1]), ("ep",)) - monkeypatch.setattr(moe, "_get_mesh", lambda: mesh) - monkeypatch.setattr(moe, "_with_sharding_constraint_cast_bwd", lambda x, _: x) - received = [] - signature = inspect.signature(moe._moe) - - def execute(*args): - bound = signature.bind(*args).arguments - assert bound["use_cudnn_fusion"] is (request_fusion is not False) - assert bound["cudnn_native_weight_layout"] is native_layout - received.append(bound["use_cudnn_jax_fusion"]) - return args[0], None, jnp.asarray(0) - - monkeypatch.setattr(moe, "_moe", execute) - x = jnp.ones((1, 1, 128), jnp.bfloat16) - wi = jnp.ones((2, 256, 128) if native_layout else (2, 128, 256), jnp.bfloat16) - kwargs = {} if request_fusion is None else {"use_cudnn_fusion": request_fusion} - if square: - wi = jnp.ones((2, 128, 128), jnp.bfloat16) - kwargs["cudnn_native_weight_layout"] = native_layout - if request_fusion is False: - expected = False - with warnings.catch_warnings(record=True) as caught: - warnings.simplefilter("always") - output, _, _ = moe.moe( - x, - jnp.ones((128, 2)), - wi, - jnp.ones((2, 128, 128)), - num_experts=2, - num_experts_per_tok=1, - mesh_resource=MeshResource(ep_resource="ep"), - **kwargs, - ) - assert received == [expected] - assert output is x - assert len(caught) == (0 if request_fusion is False or expected == "rubin" else 1) - - -@pytest.mark.parametrize("path", ["rubin", "blackwell"]) -@pytest.mark.parametrize("native_layout", [False, True]) -@pytest.mark.parametrize("gather_gated_dimension", [False, True]) -def test_forward_uses_selected_kernel(monkeypatch, path, native_layout, gather_gated_dimension): - import jax.numpy as jnp - - class SelectedKernel(Exception): - pass - - def selected(*args, **kwargs): - assert args[1].shape == (1, 256, 128) - raise SelectedKernel - - def rejected(*args, **kwargs): - raise AssertionError("Forward ignored the selected fallback") - - monkeypatch.setattr(moe.tex, "grouped_gemm_glu", selected if path == "rubin" else rejected) - monkeypatch.setattr( - moe.tex, "grouped_gemm_swiglu", selected if path == "blackwell" else rejected - ) - - def quantize(data, *args, **kwargs): - return SimpleNamespace( - get_tensor=lambda **kwargs: SimpleNamespace( - data=data, scale_inv=jnp.ones((1,), dtype=jnp.uint8) - ) - ) - - monkeypatch.setattr(moe.tex, "grouped_quantize", quantize) - - def gather(tensor, fsdp_axis, fsdp_size, sharded_axis): - data = tensor.get_tensor().data - assert sharded_axis == (1 if native_layout else 2) - return quantize(jnp.concatenate([data, data], axis=sharded_axis)) - - monkeypatch.setattr(moe, "_gather_quantized_weight", gather) - - def reorder(tensor, *, interleave): - assert interleave and gather_gated_dimension and not native_layout - gate, up = jnp.split(tensor.get_tensor().data, 2, axis=-1) - return quantize(moe.tex.pack_swiglu_pair(gate, up)) - - monkeypatch.setattr(moe, "_reorder_quantized_swiglu_weight", reorder) - quantizers = SimpleNamespace(x=SimpleNamespace(q_dtype=jnp.float8_e4m3fn), kernel=None) - kwargs = dict.fromkeys(inspect.signature(moe._ffn_fwd_per_shard).parameters) - combined = 128 if gather_gated_dimension else 256 - kwargs.update( - recv_tokens_local=jnp.ones((1, 256, 128), jnp.bfloat16), - recv_topk_weights_local=jnp.ones((1, 256)), - token_counts_local=jnp.asarray([256]), - wi=jnp.ones((1, combined, 128) if native_layout else (1, 128, combined), jnp.bfloat16), - wo=jnp.ones((1, 128, 128), jnp.bfloat16), - quantizer_sets=(quantizers, quantizers), - num_local_experts=1, - use_cudnn_jax_fusion=path, - cudnn_native_weight_layout=native_layout, - quant_before_fsdp_ag=gather_gated_dimension, - wi_fsdp_axis=(1 if native_layout else 2), - fsdp_axis="fsdp", - fsdp_size=2, - ) - with pytest.raises(SelectedKernel): - moe._ffn_fwd_per_shard(**kwargs) - - -@pytest.mark.parametrize("request_fusion", [None, True, False]) -def test_flax_forwards_fusion_bool(monkeypatch, request_fusion): - import jax - import jax.numpy as jnp - from transformer_engine.jax.flax import _MoEBlock - from transformer_engine.jax.sharding import MeshResource - - flax_moe = importlib.import_module("transformer_engine.jax.flax.moe") - received = [] - - def execute(inputs, *args, **kwargs): - received.append(kwargs["use_cudnn_fusion"]) - return inputs, None, jnp.asarray(0) - - monkeypatch.setattr(flax_moe, "moe", execute) - kwargs = {} if request_fusion is None else {"use_cudnn_fusion": request_fusion} - block = _MoEBlock( - num_experts=2, - intermediate_size=32, - mesh_resource=MeshResource(ep_resource="ep"), - **kwargs, - ) - block.init(jax.random.PRNGKey(0), jnp.ones((1, 1, 32))) - assert received == [request_fusion is not False] - - -@pytest.mark.parametrize("request_fusion", [True, False]) -def test_capacity_follows_explicit_fusion_bool(request_fusion): - kwargs = dict(num_experts=8, num_experts_per_tok=2, max_tokens_per_rank=64, ep_size=2) - alignment = moe._CUDNN_JAX_ALIGN_SIZE if request_fusion else moe._ALIGN_SIZE - expected = moe.get_moe_recv_capacity_per_rank(**kwargs, alignment=alignment) - assert moe.get_moe_recv_capacity_per_rank(**kwargs, use_cudnn_fusion=request_fusion) == expected - if request_fusion: - assert moe.get_moe_recv_capacity_per_rank(**kwargs) == expected - - -@pytest.mark.parametrize("invalid", [0, 1, None, "true"]) -def test_fusion_argument_requires_bool(invalid): - with pytest.raises(TypeError, match="use_cudnn_fusion must be a bool"): - moe.moe( - None, - None, - None, - None, - num_experts=2, - num_experts_per_tok=1, - use_cudnn_fusion=invalid, - ) - with pytest.raises(TypeError, match="use_cudnn_fusion must be a bool"): - moe.get_moe_recv_capacity_per_rank( - num_experts=2, - num_experts_per_tok=1, - max_tokens_per_rank=16, - ep_size=1, - use_cudnn_fusion=invalid, - ) - - -def test_vjp_bool_is_static_and_defaults_to_true(): - assert inspect.signature(moe._moe).parameters["use_cudnn_fusion"].default is True - assert inspect.signature(moe.moe).parameters["use_cudnn_fusion"].default is True - assert 32 in moe._moe.nondiff_argnums From d7392113bb126a70620c860e32f33547af40e07f Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Tue, 6 Oct 2026 14:20:35 -0700 Subject: [PATCH 32/38] Move MoE FSDP layout tests into dedicated subdirectory Signed-off-by: Jeremy Berchtold --- tests/jax/{ => moe}/test_moe_fsdp_layout.py | 0 1 file changed, 0 insertions(+), 0 deletions(-) rename tests/jax/{ => moe}/test_moe_fsdp_layout.py (100%) diff --git a/tests/jax/test_moe_fsdp_layout.py b/tests/jax/moe/test_moe_fsdp_layout.py similarity index 100% rename from tests/jax/test_moe_fsdp_layout.py rename to tests/jax/moe/test_moe_fsdp_layout.py From f91fc49d2efbfcdcc06633f907760ad8c046e09d Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Tue, 6 Oct 2026 14:21:23 -0700 Subject: [PATCH 33/38] Move MoE resource API tests into dedicated subdirectory Signed-off-by: Jeremy Berchtold --- tests/jax/{ => moe}/test_moe_resource_api.py | 0 1 file changed, 0 insertions(+), 0 deletions(-) rename tests/jax/{ => moe}/test_moe_resource_api.py (100%) diff --git a/tests/jax/test_moe_resource_api.py b/tests/jax/moe/test_moe_resource_api.py similarity index 100% rename from tests/jax/test_moe_resource_api.py rename to tests/jax/moe/test_moe_resource_api.py From 9d8f999ec0dddcaa3dc835ac0ebcea098f9f4658 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Fri, 9 Oct 2026 07:34:36 -0700 Subject: [PATCH 34/38] Adapt JAX MoE fusion to the shared cuDNN GLU API Signed-off-by: Jeremy Berchtold --- docs/api/jax.rst | 7 ++++ tests/jax/test_custom_call_compute.py | 27 +++++++------- tests/jax/test_te_ep_moe.py | 36 ++++++++++++------- .../jax/cpp_extensions/grouped_gemm_swiglu.py | 35 +++++++++--------- transformer_engine/jax/moe.py | 3 +- 5 files changed, 64 insertions(+), 44 deletions(-) diff --git a/docs/api/jax.rst b/docs/api/jax.rst index a296727125b..a3756e81e33 100644 --- a/docs/api/jax.rst +++ b/docs/api/jax.rst @@ -101,3 +101,10 @@ are resolved at trace time and do not require JAX Explicit sharding mode. The old ``ep_axis`` and ``data_parallelism_axes`` arguments remain accepted with a ``DeprecationWarning``. They are translated into a resource before calling the new API. Conflicting old and new arguments raise ``ValueError``. + +The cuDNN fused MoE path uses ``cudnn.grouped_gemm_glu_wrapper_sm100`` with JAX +dispatch, ``act_func="swiglu"``, ``sf_vec_size=32``, and +``generate_c=True`` to retain the pre-activation residual for its custom VJP. +It requires the shared GLU JAX API and ``cudnn.jax.grouped_gemm_dswiglu`` +for backward. The shared forward API selects the Rubin kernel on SM107 +and the Blackwell kernel on other supported SM100+ devices. diff --git a/tests/jax/test_custom_call_compute.py b/tests/jax/test_custom_call_compute.py index 2c050f521d7..487676f15db 100644 --- a/tests/jax/test_custom_call_compute.py +++ b/tests/jax/test_custom_call_compute.py @@ -2176,14 +2176,14 @@ def function(**kwargs): @pytest.fixture def frontend(self, monkeypatch): module = SimpleNamespace( - grouped_gemm_glu=self._api(_swiglu_adapter._FORWARD_ARGS), - grouped_gemm_swiglu=self._api(_swiglu_adapter._FORWARD_ARGS), + __name__="cudnn", + grouped_gemm_glu_wrapper_sm100=self._api(_swiglu_adapter._FORWARD_ARGS), grouped_gemm_dswiglu=self._api(_swiglu_adapter._BACKWARD_ARGS), ) original = importlib.import_module def import_module(name, package=None): - if name == "cudnn.jax": + if name in ("cudnn", "cudnn.jax"): return module if name == "cutlass.jax": return SimpleNamespace(is_available=lambda: True) @@ -2208,9 +2208,7 @@ def import_module(name, package=None): ) def test_api_signature(self, frontend, rubin, operation, change): name = ( - ("grouped_gemm_glu" if rubin else "grouped_gemm_swiglu") - if operation == "forward" - else "grouped_gemm_dswiglu" + "grouped_gemm_glu_wrapper_sm100" if operation == "forward" else "grouped_gemm_dswiglu" ) args = ( _swiglu_adapter._FORWARD_ARGS @@ -2232,21 +2230,22 @@ def test_api_signature(self, frontend, rubin, operation, change): @pytest.mark.parametrize("capability", [90, 100, 103, 107, 120]) @pytest.mark.parametrize( "missing", - [None, "grouped_gemm_glu", "grouped_gemm_swiglu", "grouped_gemm_dswiglu", "all"], + [None, "grouped_gemm_glu_wrapper_sm100", "grouped_gemm_dswiglu", "all"], ) def test_ordered_fallback(self, frontend, monkeypatch, capability, missing): if missing == "all": - vars(frontend).clear() + frontend.grouped_gemm_glu_wrapper_sm100 = None + frontend.grouped_gemm_dswiglu = None elif missing: delattr(frontend, missing) monkeypatch.setattr( transformer_engine_jax, "get_device_compute_capability", lambda _: capability ) expected = False - if capability >= 100 and missing not in ("grouped_gemm_dswiglu", "all"): - if capability == 107 and missing != "grouped_gemm_glu": + if capability >= 100 and missing is None: + if capability == 107: expected = "rubin" - elif missing != "grouped_gemm_swiglu": + else: expected = "blackwell" with warnings.catch_warnings(record=True) as caught: warnings.simplefilter("always") @@ -2256,10 +2255,10 @@ def test_ordered_fallback(self, frontend, monkeypatch, capability, missing): else: assert len(caught) == 1 message = str(caught[0].message) - assert "falling back" in message and "1.31.0" in message + assert "falling back" in message and "shared grouped_gemm_glu API" in message assert ("generic Blackwell+ fused" if expected else "unfused TE") in message - @pytest.mark.parametrize("name", ["cudnn.jax", "cutlass.jax"]) + @pytest.mark.parametrize("name", ["cudnn", "cudnn.jax", "cutlass.jax"]) def test_missing_dependency(self, monkeypatch, name): original = importlib.import_module @@ -2296,7 +2295,7 @@ def test_installed_frontend_contract(self): "missing,expected", [ (None, "rubin"), - ("grouped_gemm_glu", "blackwell"), + ("grouped_gemm_glu_wrapper_sm100", False), ("grouped_gemm_dswiglu", False), ], ) diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index 2ba7f78b382..9174e113887 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -918,11 +918,10 @@ def test_missing_apis_unfused_forward_and_backward(self, mesh, monkeypatch, nati """Unfused fallback retains bootstrap capacity and native-layout gradients.""" if not _USE_CUDNN_FUSION: pytest.skip("Requires fusion requested at bootstrap") - import cudnn.jax as cudnn_jax + import cudnn from transformer_engine.jax import cpp_extensions as tex - monkeypatch.setattr(cudnn_jax, "grouped_gemm_glu", None) - monkeypatch.setattr(cudnn_jax, "grouped_gemm_swiglu", None) + monkeypatch.setattr(cudnn, "grouped_gemm_glu_wrapper_sm100", None) block = _make_block(quantization_recipe=MXFP8BlockScaling()) x = _make_inputs(jax.random.PRNGKey(51)) variables, baseline_output, _ = _init_apply(block, mesh, x, jax.random.PRNGKey(52)) @@ -998,14 +997,20 @@ def checked_glu(*args, **kwargs): assert np.any(grad_x_np != 0) if get_device_compute_capability(0) == 107: - # The dedicated SwiGLU path is the reference for this kernel + # The generic TE adapter is the reference for this adapter # substitution. Its MXFP8 gradients can differ from pure JAX by # more than the strict unfused test threshold on Rubin. - import cudnn.jax as cudnn_jax - - # Simulate a frontend without the Rubin API. Selection must fall - # back to generic SwiGLU even though the GPU itself is Rubin. - monkeypatch.setattr(cudnn_jax, "grouped_gemm_glu", None) + # Exercise the generic TE adapter using the shared cuDNN GLU API. + dependencies_available = tex.grouped_gemm_swiglu_dependencies_available + monkeypatch.setattr( + tex, + "grouped_gemm_swiglu_dependencies_available", + lambda rubin=False: ( + (False, "Rubin path disabled for test") + if rubin + else dependencies_available(rubin=False) + ), + ) generic_calls = [] original_swiglu = tex.grouped_gemm_swiglu @@ -1098,9 +1103,16 @@ def test_cudnn_fused_with_checkpoint_names(self, mesh, monkeypatch, use_regular_ moe_module = importlib.import_module("transformer_engine.jax.moe") flax_moe_module = importlib.import_module("transformer_engine.jax.flax.moe") if use_regular_swiglu: - import cudnn.jax as cudnn_jax - - monkeypatch.setattr(cudnn_jax, "grouped_gemm_glu", None) + dependencies_available = tex.grouped_gemm_swiglu_dependencies_available + monkeypatch.setattr( + tex, + "grouped_gemm_swiglu_dependencies_available", + lambda rubin=False: ( + (False, "Rubin path disabled for test") + if rubin + else dependencies_available(rubin=False) + ), + ) selected_calls = [] fused_op_name = "grouped_gemm_swiglu" if use_regular_swiglu else "grouped_gemm_glu" diff --git a/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py b/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py index 671381172db..a65e05370a4 100644 --- a/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py +++ b/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py @@ -26,9 +26,7 @@ def pack_swiglu_pair(gate: jax.Array, up: jax.Array) -> jax.Array: if gate.shape != up.shape: raise ValueError(f"gate shape {gate.shape} must match up shape {up.shape}") if gate.shape[-1] % 32: - raise ValueError( - f"SwiGLU intermediate dimension {gate.shape[-1]} must be divisible by 32" - ) + raise ValueError(f"SwiGLU intermediate dimension {gate.shape[-1]} must be divisible by 32") blocks = gate.shape[-1] // 32 return jnp.stack( ( @@ -73,7 +71,8 @@ def _compact_sf(scale: jax.Array, shape: tuple[int, ...], name: str) -> jax.Arra return scale.reshape(-1)[:size].reshape(shape) -# These contracts match cuDNN Frontend 1.31.0. Probe APIs rather than the version: +# These contracts use cuDNN Frontend's shared GLU API with JAX dispatch. +# Probe APIs rather than the version: # newer optional parameters are compatible, but TE must supply every required one. _FORWARD_ARGS = ( "a_tensor", @@ -92,6 +91,7 @@ def _compact_sf(scale: jax.Array, shape: tuple[int, ...], name: str) -> jax.Arra "c_tensor", "beta_tensor", ) +_FORWARD_ARGS += ("sf_vec_size", "act_func", "generate_c") def grouped_gemm_swiglu_dependencies_available(rubin: bool = False) -> tuple[bool, str]: @@ -101,20 +101,20 @@ def grouped_gemm_swiglu_dependencies_available(rubin: bool = False) -> tuple[boo cudnn_jax = importlib.import_module("cudnn.jax") if not cutlass_jax.is_available(): return False, "CuTeDSL JAX support is unavailable" - forward = "grouped_gemm_glu" if rubin else "grouped_gemm_swiglu" - for name, arguments in ( - (forward, _FORWARD_ARGS), - ("grouped_gemm_dswiglu", _BACKWARD_ARGS), + cudnn = importlib.import_module("cudnn") + for module, name, arguments in ( + (cudnn, "grouped_gemm_glu_wrapper_sm100", _FORWARD_ARGS), + (cudnn_jax, "grouped_gemm_dswiglu", _BACKWARD_ARGS), ): - api = getattr(cudnn_jax, name, None) + api = getattr(module, name, None) if not callable(api): - return False, f"cudnn.jax.{name} is unavailable" + return False, f"{module.__name__}.{name} is unavailable" signature = inspect.signature(api) missing = set(arguments) - signature.parameters.keys() if missing: return ( False, - f"cudnn.jax.{name} is missing parameters: {', '.join(sorted(missing))}", + f"{module.__name__}.{name} is missing parameters: {', '.join(sorted(missing))}", ) # Binding detects missing/renamed keywords, positional-only arguments, # and new mandatory arguments, while allowing new optional arguments. @@ -135,7 +135,7 @@ def grouped_gemm_swiglu( compute_dtype, output_dtype, ) -> tuple[jax.Array, jax.Array, jax.Array, jax.Array, jax.Array]: - """Run cuDNN's dedicated JAX grouped MXFP8 GEMM + SwiGLU API. + """Run cuDNN's shared GLU API in JAX MXFP8 SwiGLU mode. TE owns conservative flat scale allocations for ragged grouped operations. cuDNN's JAX API accepts the compact physical atom layout, so this adapter @@ -147,7 +147,7 @@ def grouped_gemm_swiglu( raise ValueError(f"Expected physical B[E,N,K], got {b.shape}") cudnn_grouped_gemm_swiglu = getattr( - importlib.import_module("cudnn.jax"), "grouped_gemm_swiglu" + importlib.import_module("cudnn"), "grouped_gemm_glu_wrapper_sm100" ) return _grouped_gemm_forward( @@ -176,7 +176,7 @@ def grouped_gemm_glu( ) -> tuple[jax.Array, jax.Array, jax.Array, jax.Array, jax.Array]: """Run the Rubin cuDNN grouped MXFP8 GEMM + SwiGLU kernel.""" cudnn_grouped_gemm_glu = getattr( - importlib.import_module("cudnn.jax"), "grouped_gemm_glu" + importlib.import_module("cudnn"), "grouped_gemm_glu_wrapper_sm100" ) return _grouped_gemm_forward( @@ -222,6 +222,9 @@ def _grouped_gemm_forward( c_dtype=jnp.dtype(compute_dtype), d_dtype=jnp.dtype(output_dtype), discrete_col_sfd=True, + sf_vec_size=32, + act_func="swiglu", + generate_c=True, ) return ( result["c_tensor"].reshape(rows, combined), @@ -270,9 +273,7 @@ def grouped_gemm_dswiglu( b_tensor=b, c_tensor=c, sfa_tensor=_compact_sf(sfa, _sf_atom_shape(1, rows, hidden), "sfa"), - sfb_tensor=_compact_sf( - sfb, _sf_atom_shape(experts, intermediate, hidden), "sfb" - ), + sfb_tensor=_compact_sf(sfb, _sf_atom_shape(experts, intermediate, hidden), "sfb"), padded_offsets=padded_offsets.astype(jnp.int32), alpha_tensor=alpha, beta_tensor=beta, diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index 92161581199..1b2d4d162b3 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -204,7 +204,8 @@ def _select_cudnn_jax_fusion(rejection_reasons: list[str]) -> str | bool: f"use_cudnn_fusion=True: falling back to the {destination}: " + "; ".join(reasons) + ". Install cuDNN Frontend with compatible JAX APIs (TE's signatures match " - "cuDNN Frontend 1.31.0) and CuTeDSL JAX support; use supported GPU hardware.", + "the shared grouped_gemm_glu API with JAX dispatch) and CuTeDSL JAX support; " + "use supported GPU hardware.", UserWarning, stacklevel=2, ) From 60c71d8a94fca76fd8b2fcc3e0f952df54774549 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Fri, 9 Oct 2026 08:14:49 -0700 Subject: [PATCH 35/38] Use shared cuDNN GLU and dGLU APIs for JAX MoE Signed-off-by: Jeremy Berchtold --- docs/api/jax.rst | 4 +- tests/jax/test_custom_call_compute.py | 77 ++++++------- tests/jax/test_te_ep_moe.py | 102 +++--------------- .../jax/cpp_extensions/__init__.py | 2 +- ...ped_gemm_swiglu.py => grouped_gemm_glu.py} | 69 +++--------- transformer_engine/jax/moe.py | 63 +++++------ 6 files changed, 91 insertions(+), 226 deletions(-) rename transformer_engine/jax/cpp_extensions/{grouped_gemm_swiglu.py => grouped_gemm_glu.py} (81%) diff --git a/docs/api/jax.rst b/docs/api/jax.rst index a3756e81e33..73b25b06087 100644 --- a/docs/api/jax.rst +++ b/docs/api/jax.rst @@ -105,6 +105,6 @@ calling the new API. Conflicting old and new arguments raise ``ValueError``. The cuDNN fused MoE path uses ``cudnn.grouped_gemm_glu_wrapper_sm100`` with JAX dispatch, ``act_func="swiglu"``, ``sf_vec_size=32``, and ``generate_c=True`` to retain the pre-activation residual for its custom VJP. -It requires the shared GLU JAX API and ``cudnn.jax.grouped_gemm_dswiglu`` -for backward. The shared forward API selects the Rubin kernel on SM107 +It requires the shared GLU JAX API and ``cudnn.grouped_gemm_dglu_wrapper_sm100`` +for backward. Both shared APIs select their Rubin kernels on SM107 and the Blackwell kernel on other supported SM100+ devices. diff --git a/tests/jax/test_custom_call_compute.py b/tests/jax/test_custom_call_compute.py index 487676f15db..0a46740e1db 100644 --- a/tests/jax/test_custom_call_compute.py +++ b/tests/jax/test_custom_call_compute.py @@ -57,9 +57,7 @@ from transformer_engine.jax.cpp_extensions.topk import topk import transformer_engine_jax -_swiglu_adapter = importlib.import_module( - "transformer_engine.jax.cpp_extensions.grouped_gemm_swiglu" -) +_glu_adapter = importlib.import_module("transformer_engine.jax.cpp_extensions.grouped_gemm_glu") _moe_module = importlib.import_module("transformer_engine.jax.moe") GEMM_CASES = [ @@ -2142,7 +2140,7 @@ def test_topk_2d(self, dtype, problem_size): assert_allclose(prim_gathered, ref_gathered, dtype=dtype) -class TestGroupedGemmSwigluFallback: +class TestGroupedGemmGluFallback: """MoE API compatibility and dispatch with mocked kernels on one JAX device.""" @staticmethod @@ -2177,8 +2175,8 @@ def function(**kwargs): def frontend(self, monkeypatch): module = SimpleNamespace( __name__="cudnn", - grouped_gemm_glu_wrapper_sm100=self._api(_swiglu_adapter._FORWARD_ARGS), - grouped_gemm_dswiglu=self._api(_swiglu_adapter._BACKWARD_ARGS), + grouped_gemm_glu_wrapper_sm100=self._api(_glu_adapter._FORWARD_ARGS), + grouped_gemm_dglu_wrapper_sm100=self._api(_glu_adapter._BACKWARD_ARGS), ) original = importlib.import_module @@ -2189,10 +2187,9 @@ def import_module(name, package=None): return SimpleNamespace(is_available=lambda: True) return original(name, package) - monkeypatch.setattr(_swiglu_adapter.importlib, "import_module", import_module) + monkeypatch.setattr(_glu_adapter.importlib, "import_module", import_module) return module - @pytest.mark.parametrize("rubin", [False, True]) @pytest.mark.parametrize("operation", ["forward", "backward"]) @pytest.mark.parametrize( "change", @@ -2206,15 +2203,13 @@ def import_module(name, package=None): "opaque", ], ) - def test_api_signature(self, frontend, rubin, operation, change): + def test_api_signature(self, frontend, operation, change): name = ( - "grouped_gemm_glu_wrapper_sm100" if operation == "forward" else "grouped_gemm_dswiglu" - ) - args = ( - _swiglu_adapter._FORWARD_ARGS + "grouped_gemm_glu_wrapper_sm100" if operation == "forward" - else _swiglu_adapter._BACKWARD_ARGS + else "grouped_gemm_dglu_wrapper_sm100" ) + args = _glu_adapter._FORWARD_ARGS if operation == "forward" else _glu_adapter._BACKWARD_ARGS if change == "missing": delattr(frontend, name) elif change == "not_callable": @@ -2223,19 +2218,19 @@ def test_api_signature(self, frontend, rubin, operation, change): setattr(frontend, name, lambda **kwargs: None) else: setattr(frontend, name, self._api(args, change)) - available, reason = _swiglu_adapter.grouped_gemm_swiglu_dependencies_available(rubin) + available, reason = _glu_adapter.grouped_gemm_glu_dependencies_available() assert available == (change == "optional") assert bool(reason) == (change != "optional") @pytest.mark.parametrize("capability", [90, 100, 103, 107, 120]) @pytest.mark.parametrize( "missing", - [None, "grouped_gemm_glu_wrapper_sm100", "grouped_gemm_dswiglu", "all"], + [None, "grouped_gemm_glu_wrapper_sm100", "grouped_gemm_dglu_wrapper_sm100", "all"], ) def test_ordered_fallback(self, frontend, monkeypatch, capability, missing): if missing == "all": frontend.grouped_gemm_glu_wrapper_sm100 = None - frontend.grouped_gemm_dswiglu = None + frontend.grouped_gemm_dglu_wrapper_sm100 = None elif missing: delattr(frontend, missing) monkeypatch.setattr( @@ -2243,22 +2238,22 @@ def test_ordered_fallback(self, frontend, monkeypatch, capability, missing): ) expected = False if capability >= 100 and missing is None: - if capability == 107: - expected = "rubin" - else: - expected = "blackwell" + expected = True with warnings.catch_warnings(record=True) as caught: warnings.simplefilter("always") assert _moe_module._select_cudnn_jax_fusion([]) == expected - if expected == "rubin": + if expected: assert not caught else: assert len(caught) == 1 message = str(caught[0].message) - assert "falling back" in message and "shared grouped_gemm_glu API" in message - assert ("generic Blackwell+ fused" if expected else "unfused TE") in message + assert ( + "falling back" in message + and "shared grouped_gemm_glu/grouped_gemm_dglu JAX APIs" in message + ) + assert "unfused TE" in message - @pytest.mark.parametrize("name", ["cudnn", "cudnn.jax", "cutlass.jax"]) + @pytest.mark.parametrize("name", ["cudnn", "cutlass.jax"]) def test_missing_dependency(self, monkeypatch, name): original = importlib.import_module @@ -2267,8 +2262,8 @@ def import_module(module, package=None): raise ModuleNotFoundError(f"No module named {module}") return original(module, package) - monkeypatch.setattr(_swiglu_adapter.importlib, "import_module", import_module) - available, reason = _swiglu_adapter.grouped_gemm_swiglu_dependencies_available() + monkeypatch.setattr(_glu_adapter.importlib, "import_module", import_module) + available, reason = _glu_adapter.grouped_gemm_glu_dependencies_available() assert not available and name in reason def test_ineligible_call_and_device_query_failure(self, frontend, monkeypatch): @@ -2284,9 +2279,8 @@ def fail(_): assert _moe_module._select_cudnn_jax_fusion([]) is False def test_installed_frontend_contract(self): - for rubin in (False, True): - available, reason = _swiglu_adapter.grouped_gemm_swiglu_dependencies_available(rubin) - assert available or reason + available, reason = _glu_adapter.grouped_gemm_glu_dependencies_available() + assert available or reason @pytest.mark.parametrize("request_fusion", [None, True, False]) @pytest.mark.parametrize("native_layout", [False, True]) @@ -2294,9 +2288,9 @@ def test_installed_frontend_contract(self): @pytest.mark.parametrize( "missing,expected", [ - (None, "rubin"), + (None, True), ("grouped_gemm_glu_wrapper_sm100", False), - ("grouped_gemm_dswiglu", False), + ("grouped_gemm_dglu_wrapper_sm100", False), ], ) def test_public_moe_passes_selected_path( @@ -2347,14 +2341,11 @@ def execute(*args): ) assert received == [expected] assert output is x - assert len(caught) == (0 if request_fusion is False or expected == "rubin" else 1) + assert len(caught) == (0 if request_fusion is False or expected else 1) - @pytest.mark.parametrize("path", ["rubin", "blackwell"]) @pytest.mark.parametrize("native_layout", [False, True]) @pytest.mark.parametrize("gather_gated_dimension", [False, True]) - def test_forward_uses_selected_kernel( - self, monkeypatch, path, native_layout, gather_gated_dimension - ): + def test_forward_uses_selected_kernel(self, monkeypatch, native_layout, gather_gated_dimension): class SelectedKernel(Exception): pass @@ -2363,15 +2354,7 @@ def selected(*args, **kwargs): assert args[1].shape == (1, 256, 128) raise SelectedKernel - def rejected(*args, **kwargs): - raise AssertionError("Forward ignored the selected fallback") - - monkeypatch.setattr( - _moe_module.tex, "grouped_gemm_glu", selected if path == "rubin" else rejected - ) - monkeypatch.setattr( - _moe_module.tex, "grouped_gemm_swiglu", selected if path == "blackwell" else rejected - ) + monkeypatch.setattr(_moe_module.tex, "grouped_gemm_glu", selected) def quantize(data, *args, **kwargs): return SimpleNamespace( @@ -2406,7 +2389,7 @@ def reorder(tensor, *, interleave): wo=jnp.ones((1, 128, 128), jnp.bfloat16), quantizer_sets=(quantizers, quantizers), num_local_experts=1, - use_cudnn_jax_fusion=path, + use_cudnn_jax_fusion=True, cudnn_native_weight_layout=native_layout, quant_before_fsdp_ag=gather_gated_dimension, wi_fsdp_axis=(1 if native_layout else 2), diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index 9174e113887..5f6ab1b3047 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -960,26 +960,31 @@ def native_moe(*args, **kwargs): def test_mxfp8_forward_and_backward(self, mesh, apply_topk_weights_early, monkeypatch): if not _USE_CUDNN_FUSION: pytest.skip("run separately with --use-cudnn-fusion=1") - rubin_calls = [] - if get_device_compute_capability(0) == 107: - from transformer_engine.jax import cpp_extensions as tex + from transformer_engine.jax import cpp_extensions as tex + + shared_calls = [] + original_glu = tex.grouped_gemm_glu + original_dglu = tex.grouped_gemm_dglu - original_glu = tex.grouped_gemm_glu + def checked_glu(*args, **kwargs): + shared_calls.append("glu") + return original_glu(*args, **kwargs) - def checked_glu(*args, **kwargs): - rubin_calls.append(True) - return original_glu(*args, **kwargs) + def checked_dglu(*args, **kwargs): + shared_calls.append("dglu") + return original_dglu(*args, **kwargs) - monkeypatch.setattr(tex, "grouped_gemm_glu", checked_glu) + monkeypatch.setattr(tex, "grouped_gemm_glu", checked_glu) + monkeypatch.setattr(tex, "grouped_gemm_dglu", checked_dglu) block = _make_block( apply_topk_weights_early=apply_topk_weights_early, quantization_recipe=MXFP8BlockScaling(), ) x = _make_inputs(jax.random.PRNGKey(30)) variables, output, aux = _init_apply(block, mesh, x, jax.random.PRNGKey(31)) - if get_device_compute_capability(0) == 107: - assert rubin_calls, "Rubin MoE did not select the cuDNN GLU JAX API" + assert "glu" in shared_calls, "MoE did not select the shared cuDNN GLU JAX API" grads, grad_x = _grad_step(block, variables, mesh, x) + assert "dglu" in shared_calls, "MoE VJP did not select the shared cuDNN dGLU JAX API" assert output.shape == x.shape assert output.dtype == x.dtype @@ -996,63 +1001,6 @@ def checked_glu(*args, **kwargs): assert np.all(np.isfinite(grad_x_np)) assert np.any(grad_x_np != 0) - if get_device_compute_capability(0) == 107: - # The generic TE adapter is the reference for this adapter - # substitution. Its MXFP8 gradients can differ from pure JAX by - # more than the strict unfused test threshold on Rubin. - # Exercise the generic TE adapter using the shared cuDNN GLU API. - dependencies_available = tex.grouped_gemm_swiglu_dependencies_available - monkeypatch.setattr( - tex, - "grouped_gemm_swiglu_dependencies_available", - lambda rubin=False: ( - (False, "Rubin path disabled for test") - if rubin - else dependencies_available(rubin=False) - ), - ) - generic_calls = [] - original_swiglu = tex.grouped_gemm_swiglu - - def checked_swiglu(*args, **kwargs): - generic_calls.append(True) - return original_swiglu(*args, **kwargs) - - monkeypatch.setattr(tex, "grouped_gemm_swiglu", checked_swiglu) - with _ctx(mesh): - x_sh = _shard_inputs(x, mesh) - baseline_output, _, _ = jax.jit(block.apply)(variables, x_sh) - baseline_output.block_until_ready() - baseline_grads, baseline_grad_x = _grad_step(block, variables, mesh, x) - assert generic_calls, "Missing Rubin API did not select generic SwiGLU" - np.testing.assert_allclose( - output_np, - _to_global_numpy(baseline_output, mesh).astype(np.float32), - **FWD_TOLERANCE["mxfp8"], - err_msg="Rubin GLU forward differs from dedicated SwiGLU", - ) - for name in ("gate_kernel", "wi", "wo"): - tolerance = ( - GRAD_GATE_TOLERANCE["mxfp8"] - if name == "gate_kernel" - else GRAD_FFN_TOLERANCE["mxfp8"] - ) - np.testing.assert_allclose( - _to_global_numpy(_unwrap(grads["params"][name]), mesh).astype(np.float32), - _to_global_numpy(_unwrap(baseline_grads["params"][name]), mesh).astype( - np.float32 - ), - **tolerance, - err_msg=f"Rubin GLU {name} gradient differs from dedicated SwiGLU", - ) - np.testing.assert_allclose( - grad_x_np, - _to_global_numpy(baseline_grad_x, mesh).astype(np.float32), - **GRAD_FFN_TOLERANCE["mxfp8"], - err_msg="Rubin GLU input gradient differs from dedicated SwiGLU", - ) - return - params_np = _params_global_numpy(variables, mesh) x_np = np.asarray(jax.device_get(x)) @@ -1091,31 +1039,15 @@ def loss_fn(params, inputs): err_msg="d_x fused MXFP8 gradient parity breach", ) - @pytest.mark.parametrize("use_regular_swiglu", [False, True]) - def test_cudnn_fused_with_checkpoint_names(self, mesh, monkeypatch, use_regular_swiglu): + def test_cudnn_fused_with_checkpoint_names(self, mesh, monkeypatch): if not _USE_CUDNN_FUSION: pytest.skip("cuDNN grouped GEMM fusion is disabled") - if not use_regular_swiglu and get_device_compute_capability(0) != 107: - pytest.skip("Rubin grouped GLU requires SM107") - from transformer_engine.jax import cpp_extensions as tex moe_module = importlib.import_module("transformer_engine.jax.moe") flax_moe_module = importlib.import_module("transformer_engine.jax.flax.moe") - if use_regular_swiglu: - dependencies_available = tex.grouped_gemm_swiglu_dependencies_available - monkeypatch.setattr( - tex, - "grouped_gemm_swiglu_dependencies_available", - lambda rubin=False: ( - (False, "Rubin path disabled for test") - if rubin - else dependencies_available(rubin=False) - ), - ) - selected_calls = [] - fused_op_name = "grouped_gemm_swiglu" if use_regular_swiglu else "grouped_gemm_glu" + fused_op_name = "grouped_gemm_glu" original_fused_op = getattr(tex, fused_op_name) def checked_fused_op(*args, **kwargs): diff --git a/transformer_engine/jax/cpp_extensions/__init__.py b/transformer_engine/jax/cpp_extensions/__init__.py index a78916ff95a..c9637838577 100644 --- a/transformer_engine/jax/cpp_extensions/__init__.py +++ b/transformer_engine/jax/cpp_extensions/__init__.py @@ -10,7 +10,7 @@ from .quantization import * from .softmax import * from .gemm import * -from .grouped_gemm_swiglu import * +from .grouped_gemm_glu import * from .router import * from .ep import * from .topk import * diff --git a/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py b/transformer_engine/jax/cpp_extensions/grouped_gemm_glu.py similarity index 81% rename from transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py rename to transformer_engine/jax/cpp_extensions/grouped_gemm_glu.py index a65e05370a4..a7b1b563e2d 100644 --- a/transformer_engine/jax/cpp_extensions/grouped_gemm_swiglu.py +++ b/transformer_engine/jax/cpp_extensions/grouped_gemm_glu.py @@ -1,7 +1,7 @@ # Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. # # See LICENSE for license information. -"""cuDNN Frontend JAX API adapter for fused grouped GEMM + SwiGLU.""" +"""cuDNN Frontend JAX API adapters for shared grouped GEMM + GLU/dGLU.""" from __future__ import annotations @@ -12,10 +12,9 @@ import jax.numpy as jnp __all__ = [ - "grouped_gemm_dswiglu", + "grouped_gemm_dglu", "grouped_gemm_glu", - "grouped_gemm_swiglu", - "grouped_gemm_swiglu_dependencies_available", + "grouped_gemm_glu_dependencies_available", "pack_swiglu_pair", "unpack_swiglu_pair", ] @@ -90,21 +89,23 @@ def _compact_sf(scale: jax.Array, shape: tuple[int, ...], name: str) -> jax.Arra _BACKWARD_ARGS = tuple(arg for arg in _FORWARD_ARGS if arg != "c_dtype") + ( "c_tensor", "beta_tensor", + "dprob_tensor", + "sf_vec_size", + "act_func", ) _FORWARD_ARGS += ("sf_vec_size", "act_func", "generate_c") -def grouped_gemm_swiglu_dependencies_available(rubin: bool = False) -> tuple[bool, str]: +def grouped_gemm_glu_dependencies_available() -> tuple[bool, str]: """Check that the selected forward and shared backward accept TE's keywords.""" try: cutlass_jax = importlib.import_module("cutlass.jax") - cudnn_jax = importlib.import_module("cudnn.jax") if not cutlass_jax.is_available(): return False, "CuTeDSL JAX support is unavailable" cudnn = importlib.import_module("cudnn") for module, name, arguments in ( (cudnn, "grouped_gemm_glu_wrapper_sm100", _FORWARD_ARGS), - (cudnn_jax, "grouped_gemm_dswiglu", _BACKWARD_ARGS), + (cudnn, "grouped_gemm_dglu_wrapper_sm100", _BACKWARD_ARGS), ): api = getattr(module, name, None) if not callable(api): @@ -124,45 +125,6 @@ def grouped_gemm_swiglu_dependencies_available(rubin: bool = False) -> tuple[boo return True, "" -def grouped_gemm_swiglu( - a: jax.Array, - b: jax.Array, - sfa: jax.Array, - sfb: jax.Array, - padded_offsets: jax.Array, - prob: jax.Array, - *, - compute_dtype, - output_dtype, -) -> tuple[jax.Array, jax.Array, jax.Array, jax.Array, jax.Array]: - """Run cuDNN's shared GLU API in JAX MXFP8 SwiGLU mode. - - TE owns conservative flat scale allocations for ragged grouped operations. - cuDNN's JAX API accepts the compact physical atom layout, so this adapter - exposes the used prefix with a zero-copy reshape before making the call. - """ - if a.ndim != 3 or a.shape[-1] != 1: - raise ValueError(f"Expected A[M,K,1], got {a.shape}") - if b.ndim != 3: - raise ValueError(f"Expected physical B[E,N,K], got {b.shape}") - - cudnn_grouped_gemm_swiglu = getattr( - importlib.import_module("cudnn"), "grouped_gemm_glu_wrapper_sm100" - ) - - return _grouped_gemm_forward( - cudnn_grouped_gemm_swiglu, - a, - b, - sfa, - sfb, - padded_offsets, - prob, - compute_dtype=compute_dtype, - output_dtype=output_dtype, - ) - - def grouped_gemm_glu( a: jax.Array, b: jax.Array, @@ -174,7 +136,7 @@ def grouped_gemm_glu( compute_dtype, output_dtype, ) -> tuple[jax.Array, jax.Array, jax.Array, jax.Array, jax.Array]: - """Run the Rubin cuDNN grouped MXFP8 GEMM + SwiGLU kernel.""" + """Run cuDNN's shared grouped MXFP8 GLU API in SwiGLU mode.""" cudnn_grouped_gemm_glu = getattr( importlib.import_module("cudnn"), "grouped_gemm_glu_wrapper_sm100" ) @@ -235,7 +197,7 @@ def _grouped_gemm_forward( ) -def grouped_gemm_dswiglu( +def grouped_gemm_dglu( a: jax.Array, b: jax.Array, c: jax.Array, @@ -246,7 +208,7 @@ def grouped_gemm_dswiglu( *, output_dtype, ) -> tuple[jax.Array, jax.Array, jax.Array, jax.Array, jax.Array]: - """Run cuDNN's grouped MXFP8 GEMM + dSwiGLU + MXFP8 quantization API.""" + """Run cuDNN's shared grouped MXFP8 dGLU API in dSwiGLU mode.""" if a.ndim != 2: raise ValueError(f"Expected A[M,K], got {a.shape}") if b.ndim != 3: @@ -254,8 +216,8 @@ def grouped_gemm_dswiglu( if c.ndim != 2: raise ValueError(f"Expected C[M,2N], got {c.shape}") - cudnn_grouped_gemm_dswiglu = getattr( - importlib.import_module("cudnn.jax"), "grouped_gemm_dswiglu" + cudnn_grouped_gemm_dglu = getattr( + importlib.import_module("cudnn"), "grouped_gemm_dglu_wrapper_sm100" ) rows, hidden = a.shape @@ -268,7 +230,7 @@ def grouped_gemm_dswiglu( alpha = jnp.ones((experts,), dtype=jnp.float32) beta = jnp.ones((experts,), dtype=jnp.float32) norm_const = jnp.ones((1,), dtype=jnp.float32) - result = cudnn_grouped_gemm_dswiglu( + result = cudnn_grouped_gemm_dglu( a_tensor=a, b_tensor=b, c_tensor=c, @@ -281,6 +243,9 @@ def grouped_gemm_dswiglu( norm_const_tensor=norm_const, d_dtype=jnp.dtype(output_dtype), discrete_col_sfd=True, + dprob_tensor=None, + sf_vec_size=32, + act_func="dswiglu", ) return ( result["d_row_tensor"], diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index 1b2d4d162b3..877a0b7ba0d 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -169,47 +169,32 @@ def _resolve_moe_mesh_resource( _CUDNN_JAX_ALIGN_SIZE = 256 -def _select_cudnn_jax_fusion(rejection_reasons: list[str]) -> str | bool: - """Select Rubin GLU, generic Blackwell+ SwiGLU, or unfused TE, in order.""" +def _select_cudnn_jax_fusion(rejection_reasons: list[str]) -> bool: + """Select shared cuDNN GLU/dGLU fusion or the unfused TE path.""" from transformer_engine_jax import get_device_compute_capability reasons = list(rejection_reasons) - path = False try: capability = get_device_compute_capability(0) except RuntimeError as exc: reasons.append(f"could not query GPU compute capability: {exc}") else: + if capability < 100: + reasons.append(f"fused kernels require SM100+, got SM{capability}") if not reasons: - for candidate, supported in ( - ("rubin", capability == 107), - ("blackwell", capability >= 100), - ): - if not supported: - requirement = "SM107" if candidate == "rubin" else "SM100+" - reasons.append( - f"{candidate} fused kernel requires {requirement}, got SM{capability}" - ) - continue - available, error = tex.grouped_gemm_swiglu_dependencies_available( - rubin=candidate == "rubin" - ) - if available: - path = candidate - break - reasons.append(f"{candidate} fused API is incompatible: {error}") - if reasons: - destination = "generic Blackwell+ fused kernel" if path else "unfused TE grouped-GEMM path" - warnings.warn( - f"use_cudnn_fusion=True: falling back to the {destination}: " - + "; ".join(reasons) - + ". Install cuDNN Frontend with compatible JAX APIs (TE's signatures match " - "the shared grouped_gemm_glu API with JAX dispatch) and CuTeDSL JAX support; " - "use supported GPU hardware.", - UserWarning, - stacklevel=2, - ) - return path + available, error = tex.grouped_gemm_glu_dependencies_available() + if available: + return True + reasons.append(f"shared fused API is incompatible: {error}") + warnings.warn( + "use_cudnn_fusion=True: falling back to the unfused TE grouped-GEMM path: " + + "; ".join(reasons) + + ". Install cuDNN Frontend with shared grouped_gemm_glu/grouped_gemm_dglu " + "JAX APIs and CuTeDSL JAX support; use supported GPU hardware.", + UserWarning, + stacklevel=2, + ) + return False def _cudnn_jax_fusion_rejection_reasons( @@ -223,7 +208,7 @@ def _cudnn_jax_fusion_rejection_reasons( activation_type, ep_axis, ) -> list[str]: - """Return reasons this call cannot use cuDNN's grouped SwiGLU JAX API.""" + """Return reasons this call cannot use cuDNN's shared grouped GLU/dGLU JAX APIs.""" errors = [] if str(activation_type).lower() != "silu": errors.append("requires activation_type='silu'") @@ -803,7 +788,7 @@ def _ffn_fwd_per_shard( num_local_experts: int, activation_type: str, apply_topk_weights_early: bool, - use_cudnn_jax_fusion: str | bool, + use_cudnn_jax_fusion: bool, wi_0_checkpoint_name: Optional[str], wi_1_checkpoint_name: Optional[str], wo_checkpoint_name: Optional[str], @@ -864,7 +849,7 @@ def _ffn_fwd_per_shard( intermediate_col, intermediate_scale_row, intermediate_scale_col, - ) = (tex.grouped_gemm_glu if use_cudnn_jax_fusion == "rubin" else tex.grouped_gemm_swiglu)( + ) = tex.grouped_gemm_glu( casted_sorted_x_lhs.data.reshape(sorted_x.shape[0], hidden, 1), ( casted_wi_rhs.data.reshape(num_local_experts, combined, hidden) @@ -1016,7 +1001,7 @@ def _ffn_bwd_per_shard( activation_type: str, apply_topk_weights_early: bool, has_bias: bool, - use_cudnn_jax_fusion: str | bool, + use_cudnn_jax_fusion: bool, cudnn_native_weight_layout: bool, ): """Backward mirror of :func:`_ffn_fwd_per_shard`.""" @@ -1056,7 +1041,7 @@ def _ffn_bwd_per_shard( dprob, d_combined_scale_row, d_combined_scale_col, - ) = tex.grouped_gemm_dswiglu( + ) = tex.grouped_gemm_dglu( _casted_d_eo_lhs.data.reshape(rows, hidden), casted_wo_rhs_trans.data.reshape(num_local_experts, intermediate, hidden), combined_out, @@ -2089,8 +2074,8 @@ def moe( preserve compatibility with EP bootstrap buffer sizing. use_cudnn_fusion : bool - Defaults to ``True``: try cuDNN's JAX grouped MXFP8 APIs, Rubin GLU first, - then generic SM100+ SwiGLU. ``False`` uses unfused TE grouped GEMM. + Defaults to ``True``: try cuDNN's shared JAX grouped MXFP8 GLU/dGLU APIs + on SM100+ GPUs; cuDNN selects the architecture-specific kernels. ``False`` uses unfused TE grouped GEMM. Ineligible calls warn and fall back to TE's regular grouped-GEMM implementation. API signatures and GPU capability determine support. From ca11b8b8fad8d0710331411ec1b07ddbaeabb5d2 Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Fri, 9 Oct 2026 08:50:53 -0700 Subject: [PATCH 36/38] Configure JAX MoE parallelism exclusively through MeshResource Signed-off-by: Jeremy Berchtold --- docs/api/jax.rst | 7 +- tests/jax/moe/test_moe_resource_api.py | 73 +++--------------- tests/jax/test_te_ep_moe.py | 25 ++---- transformer_engine/jax/cpp_extensions/ep.py | 10 +-- transformer_engine/jax/flax/moe.py | 12 +-- transformer_engine/jax/moe.py | 84 ++------------------- 6 files changed, 34 insertions(+), 177 deletions(-) diff --git a/docs/api/jax.rst b/docs/api/jax.rst index 73b25b06087..dc6fc255780 100644 --- a/docs/api/jax.rst +++ b/docs/api/jax.rst @@ -98,9 +98,10 @@ gathering concatenates the quantized matrices and their existing scale blocks; matrix-shard gathering reconstructs each expert's scale blocks. These specs are resolved at trace time and do not require JAX Explicit sharding mode. -The old ``ep_axis`` and ``data_parallelism_axes`` arguments remain accepted -with a ``DeprecationWarning``. They are translated into a resource before -calling the new API. Conflicting old and new arguments raise ``ValueError``. +MoE parallelism is configured through ``MeshResource``. The former +``ep_axis`` and ``data_parallelism_axes`` arguments are no longer accepted. +Set ``dp_resource``, ``fsdp_resource``, and ``ep_resource`` to physical mesh +axis names. Outer batch axes follow DP then FSDP, with EP innermost. The cuDNN fused MoE path uses ``cudnn.grouped_gemm_glu_wrapper_sm100`` with JAX dispatch, ``act_func="swiglu"``, ``sf_vec_size=32``, and diff --git a/tests/jax/moe/test_moe_resource_api.py b/tests/jax/moe/test_moe_resource_api.py index 65a138054e4..23e3eca3fa7 100644 --- a/tests/jax/moe/test_moe_resource_api.py +++ b/tests/jax/moe/test_moe_resource_api.py @@ -2,11 +2,10 @@ # # See LICENSE for license information. -"""MoE resource resolution and deprecated API compatibility without EP kernels.""" +"""MoE resource resolution and MeshResource-only API without EP kernels.""" import importlib import inspect -import warnings import jax import jax.numpy as jnp @@ -86,47 +85,9 @@ def test_bool_and_resource_types(): _resolve_moe_mesh_resource("ep") -def test_legacy_preserves_arbitrary_outer_axis_order(): - with global_shard_guard(MeshResource(fsdp_resource="fsdp", ep_resource="ep")), pytest.warns( - DeprecationWarning - ): - resource, quantize = _resolve_moe_mesh_resource( - ep_axis="ep", - data_parallelism_axes=("outer", "fsdp", "replica"), - quant_before_fsdp_ag=True, - ) - assert _moe_mesh_axes(resource) == ("ep", ("outer", "fsdp", "replica")) - assert resource.fsdp_resource == "fsdp" - assert quantize is True - - -def test_legacy_defaults_to_no_outer_axes(): - with global_shard_guard(MeshResource(dp_resource="dp", fsdp_resource="fsdp", ep_resource="ep")): - with pytest.warns(DeprecationWarning): - resource, _ = _resolve_moe_mesh_resource(ep_axis="ep") - assert _moe_mesh_axes(resource) == ("ep", ()) - - -@pytest.mark.parametrize( - "kwargs,error", - [ - ({"ep_axis": "other"}, "ep_axis conflicts"), - ({"data_parallelism_axes": ("other",)}, "data_parallelism_axes conflicts"), - ], -) -def test_conflicting_old_and_new_args(kwargs, error): - with pytest.warns(DeprecationWarning): - with pytest.raises(ValueError, match=error): - _resolve_moe_mesh_resource( - MeshResource(fsdp_resource="fsdp", ep_resource="ep"), **kwargs - ) - - -@pytest.mark.parametrize("legacy", [False, True]) @pytest.mark.parametrize("quantize", [False, True]) -def test_public_api_delegates_with_selected_resource(monkeypatch, legacy, quantize): +def test_public_api_uses_selected_resource(monkeypatch, quantize): module = importlib.import_module("transformer_engine.jax.moe") - original_moe = module.moe signature = inspect.signature(module._moe) captured = {} @@ -135,36 +96,22 @@ def fake_vjp(*args): assert global_mesh_resource().ep_resource == "ep" return args[0], None, jnp.zeros((1,), jnp.int32) - def delegated_moe(*args, **kwargs): - assert "mesh_resource" in kwargs - assert not {"ep_axis", "data_parallelism_axes", "weight_gather"}.intersection(kwargs) - return original_moe(*args, **kwargs) - monkeypatch.setattr(module, "_moe", fake_vjp) - monkeypatch.setattr(module, "moe", delegated_moe) mesh = Mesh(np.asarray(jax.devices()[:1]).reshape(1, 1, 1), ("dp", "fsdp", "ep")) - kwargs = dict(num_experts=2, num_experts_per_tok=1) - if legacy: - kwargs.update( - ep_axis="ep", - data_parallelism_axes=("dp", "fsdp"), - quant_before_fsdp_ag=quantize, - ) - else: - kwargs.update( - mesh_resource=MeshResource(dp_resource="dp", fsdp_resource="fsdp", ep_resource="ep"), - quant_before_fsdp_ag=quantize, - ) - with jax.set_mesh(mesh), warnings.catch_warnings(record=True) as recorded: - warnings.simplefilter("always", DeprecationWarning) - original_moe( + kwargs = dict( + num_experts=2, + num_experts_per_tok=1, + mesh_resource=MeshResource(dp_resource="dp", fsdp_resource="fsdp", ep_resource="ep"), + quant_before_fsdp_ag=quantize, + ) + with jax.set_mesh(mesh): + module.moe( jnp.ones((1, 1, 4)), jnp.ones((4, 2)), jnp.ones((2, 4, 8)), jnp.ones((2, 4, 4)), **kwargs, ) - assert any("deprecated for TE MoE" in str(w.message) for w in recorded) == legacy assert _moe_mesh_axes(captured["mesh_resource"]) == ("ep", ("dp", "fsdp")) assert captured["quant_before_fsdp_ag"] is quantize assert "ep_axis" not in captured diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index 5f6ab1b3047..a53b08fa5a0 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -654,32 +654,21 @@ def layout_moe(*args, **kwargs): @pytest.mark.parametrize("quant_before_fsdp_ag", [False, True]) -@pytest.mark.parametrize("api", ["explicit", "legacy"]) -def test_mesh_resource_api_forward_and_backward(mesh, quant_before_fsdp_ag, api): - """Explicit resources need no global context; legacy calls retain numerical semantics.""" +def test_mesh_resource_api_forward_and_backward(mesh, quant_before_fsdp_ag): + """Explicit and global resources give identical forward and backward results.""" baseline = _make_block( quantization_recipe=MXFP8BlockScaling(), quant_before_fsdp_ag=quant_before_fsdp_ag ) - if api == "explicit": - candidate = baseline.clone( - mesh_resource=MeshResource(ep_resource=EP_AXIS, fsdp_resource=FSDP_AXIS) - ) - else: - candidate = baseline.clone( - data_parallelism_axes=(FSDP_AXIS,), - ) + candidate = baseline.clone( + mesh_resource=MeshResource(ep_resource=EP_AXIS, fsdp_resource=FSDP_AXIS) + ) x = _make_inputs(jax.random.PRNGKey(51)) variables, baseline_out, _ = _init_apply(baseline, mesh, x, jax.random.PRNGKey(52)) baseline_grads, baseline_dx = _grad_step(baseline, variables, mesh, x) with jax.set_mesh(mesh), nn_partitioning.axis_rules(LOGICAL_AXIS_RULES): - resource = None if api == "explicit" else MeshResource(ep_resource=EP_AXIS) - with global_shard_guard(resource): + with global_shard_guard(None): x_sh = _shard_inputs(x, mesh) - if api == "legacy": - with pytest.warns(DeprecationWarning, match="deprecated for TE MoE"): - candidate_out, _, _ = jax.jit(candidate.apply)(variables, x_sh) - else: - candidate_out, _, _ = jax.jit(candidate.apply)(variables, x_sh) + candidate_out, _, _ = jax.jit(candidate.apply)(variables, x_sh) def loss_fn(variables, inputs): output, _, _ = candidate.apply(variables, inputs) diff --git a/transformer_engine/jax/cpp_extensions/ep.py b/transformer_engine/jax/cpp_extensions/ep.py index 5556d0b665f..60c8451fae9 100644 --- a/transformer_engine/jax/cpp_extensions/ep.py +++ b/transformer_engine/jax/cpp_extensions/ep.py @@ -199,13 +199,11 @@ def ep_handle_mem_size(cfg: EpLayerConfig) -> int: def _capture_ep_resource_axes(): """Capture physical axes at bind time for asynchronous SPMD partitioning.""" resource = global_mesh_resource() - outer_axes = getattr(resource, "_legacy_data_parallelism_axes", None) - if outer_axes is None: - outer_axes = tuple( - dict.fromkeys( - axis for axis in (resource.dp_resource, resource.fsdp_resource) if axis is not None - ) + outer_axes = tuple( + dict.fromkeys( + axis for axis in (resource.dp_resource, resource.fsdp_resource) if axis is not None ) + ) return resource.ep_resource, outer_axes diff --git a/transformer_engine/jax/flax/moe.py b/transformer_engine/jax/flax/moe.py index 5528185bbd6..6c4d5a2c8cf 100644 --- a/transformer_engine/jax/flax/moe.py +++ b/transformer_engine/jax/flax/moe.py @@ -100,12 +100,9 @@ class _MoEBlock(TransformerEngineBase): quant_before_fsdp_ag : bool Quantize MXFP8 weight shards before their FSDP all-gather. Defaults to ``False``; ``True`` requires ``mesh_resource.fsdp_resource``. - ep_axis, data_parallelism_axes : deprecated - Compatibility axis arguments converted into MeshResource - with a DeprecationWarning. use_cudnn_fusion : bool - Defaults to ``True``: try Rubin fused GLU, then generic Blackwell+ fused - SwiGLU, then unfused TE grouped GEMM. Unsupported fused paths warn. + Defaults to ``True``: try shared cuDNN GLU/dGLU fusion on SM100+ GPUs. + Ineligible calls warn and use unfused TE grouped GEMM. ``False`` selects unfused execution. Use the same value when calculating EP bootstrap receive capacity with ``get_moe_recv_capacity_per_rank``. apply_topk_weights_early : bool @@ -159,9 +156,6 @@ class _MoEBlock(TransformerEngineBase): # Parallelism mesh_resource: Optional[MeshResource] = None quant_before_fsdp_ag: bool = False - # Deprecated compatibility arguments. - ep_axis: Optional[str] = None - data_parallelism_axes: Optional[Tuple[str, ...]] = None # MoE knobs forwarded to ``moe()`` use_cudnn_fusion: bool = True @@ -211,8 +205,6 @@ def __call__(self, inputs: Array) -> Tuple[Array, Optional[Array], Array]: mesh_resource, quant_before_fsdp_ag = _resolve_moe_mesh_resource( self.mesh_resource, self.quant_before_fsdp_ag, - self.ep_axis, - self.data_parallelism_axes, ) _, data_parallelism_axes = _moe_mesh_axes(mesh_resource) assert ( diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index 877a0b7ba0d..751b7a28bf7 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -32,7 +32,7 @@ import math import warnings -from dataclasses import dataclass, fields, replace +from dataclasses import replace from functools import partial from typing import Any, Optional, Tuple, Union @@ -64,28 +64,18 @@ __all__ = ["get_moe_recv_capacity_per_rank", "moe"] -@dataclass -class _LegacyMoEMeshResource(MeshResource): - """Preserve arbitrary ordered outer axes accepted by the deprecated API.""" - - _legacy_data_parallelism_axes: Tuple[str, ...] = () - - def _moe_mesh_axes(resource: MeshResource): """Resolve physical MoE axes, keeping EP innermost in the batch shard.""" if not isinstance(resource.ep_resource, str) or not resource.ep_resource: raise ValueError("TE MoE requires MeshResource.ep_resource to name a physical mesh axis.") - if isinstance(resource, _LegacyMoEMeshResource): - outer_axes = resource._legacy_data_parallelism_axes - else: - outer_axes = tuple( - dict.fromkeys( - axis for axis in (resource.dp_resource, resource.fsdp_resource) if axis is not None - ) + outer_axes = tuple( + dict.fromkeys( + axis for axis in (resource.dp_resource, resource.fsdp_resource) if axis is not None ) + ) if any(not isinstance(axis, str) or not axis for axis in outer_axes): raise ValueError("TE MoE DP and FSDP resources must name physical mesh axes.") - if resource.ep_resource in outer_axes or len(set(outer_axes)) != len(outer_axes): + if resource.ep_resource in outer_axes: raise ValueError("TE MoE EP and outer data-parallel axes must be distinct.") return resource.ep_resource, outer_axes @@ -93,60 +83,18 @@ def _moe_mesh_axes(resource: MeshResource): def _resolve_moe_mesh_resource( mesh_resource=None, quant_before_fsdp_ag=False, - ep_axis=None, - data_parallelism_axes=None, ): - """Resolve the canonical API and adapt deprecated axis arguments.""" + """Resolve and snapshot the explicit or active mesh resource.""" if not isinstance(quant_before_fsdp_ag, bool): raise TypeError("quant_before_fsdp_ag must be a bool.") if mesh_resource is not None and not isinstance(mesh_resource, MeshResource): raise TypeError("mesh_resource must be a MeshResource or None.") - explicit_resource = mesh_resource is not None if mesh_resource is None: try: mesh_resource = global_mesh_resource() except AssertionError: mesh_resource = None - legacy = ep_axis is not None or data_parallelism_axes is not None - if legacy: - warnings.warn( - "ep_axis and data_parallelism_axes are deprecated for TE MoE; " - "pass mesh_resource=MeshResource(...) and quant_before_fsdp_ag instead.", - DeprecationWarning, - stacklevel=3, - ) - if explicit_resource: - if ep_axis is not None and ep_axis != mesh_resource.ep_resource: - raise ValueError("ep_axis conflicts with mesh_resource.ep_resource.") - if ( - data_parallelism_axes is not None - and tuple(data_parallelism_axes) != _moe_mesh_axes(mesh_resource)[1] - ): - raise ValueError("data_parallelism_axes conflicts with mesh_resource.") - else: - # The old functional API defaulted to no outer axes even in a global context. - axes = tuple(data_parallelism_axes or ()) - fsdp_axis = getattr(mesh_resource, "fsdp_resource", None) - if fsdp_axis not in axes: - fsdp_axis = axes[-1] if axes else None - dp_axes = tuple(axis for axis in axes if axis != fsdp_axis) - resources = { - field.name: getattr(mesh_resource, field.name, None) - for field in fields(MeshResource) - } - resources.update( - ep_resource=( - ep_axis if ep_axis is not None else getattr(mesh_resource, "ep_resource", None) - ), - dp_resource=dp_axes[0] if dp_axes else None, - fsdp_resource=fsdp_axis, - ) - mesh_resource = _LegacyMoEMeshResource( - **resources, - _legacy_data_parallelism_axes=axes, - ) - if mesh_resource is None: raise ValueError( "TE MoE requires mesh_resource=MeshResource(...) or an active global_shard_guard" @@ -1988,8 +1936,6 @@ def moe( ), mesh_resource: Optional[MeshResource] = None, quant_before_fsdp_ag: bool = False, - ep_axis: Optional[str] = None, - data_parallelism_axes: Optional[Tuple[str, ...]] = None, input_axes: Tuple[Optional[str], ...] = (), gate_kernel_axes: Tuple[Optional[str], ...] = (), wi_kernel_axes: Tuple[Optional[str], ...] = ("exp", "embed", "mlp"), @@ -2066,9 +2012,6 @@ def moe( FSDP shards the gated dimension, layout conversion permutes gathered FP8 data and inverse scales, preserving global gate/up pairing without gathering or requantizing full-precision weights. - ep_axis, data_parallelism_axes : deprecated - Compatibility axis arguments converted into a MeshResource, - with a DeprecationWarning. Conflicting old and new arguments raise. Per-expert dispatch-slot alignment defaults to 128 tokens (``_ALIGN_SIZE``). Requesting cuDNN fusion reserves 256 tokens, also when falling back, to preserve compatibility with EP bootstrap buffer sizing. @@ -2097,19 +2040,6 @@ def moe( """ if not isinstance(use_cudnn_fusion, bool): raise TypeError("use_cudnn_fusion must be a bool") - if ep_axis is not None or data_parallelism_axes is not None: - call_args = locals().copy() - resource, quantize = _resolve_moe_mesh_resource( - mesh_resource, - quant_before_fsdp_ag, - ep_axis, - data_parallelism_axes, - ) - for name in ("ep_axis", "data_parallelism_axes"): - call_args.pop(name) - call_args.update(mesh_resource=resource, quant_before_fsdp_ag=quantize) - return moe(**call_args) - mesh_resource, quant_before_fsdp_ag = _resolve_moe_mesh_resource( mesh_resource, quant_before_fsdp_ag ) From 79f0ad81929f74f514a5a1b4a1e4ac98abbf3abc Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Fri, 9 Oct 2026 09:13:02 -0700 Subject: [PATCH 37/38] Group MoE tests by fusion mode with module-scoped EP setup Signed-off-by: Jeremy Berchtold --- tests/jax/conftest.py | 8 +-- tests/jax/run_te_ep_moe.sh | 114 +++++++++++++----------------------- tests/jax/test_te_ep_moe.py | 113 +++++++++++++++++++---------------- 3 files changed, 105 insertions(+), 130 deletions(-) diff --git a/tests/jax/conftest.py b/tests/jax/conftest.py index c7cebc7fdb1..13ae6ff0aa8 100644 --- a/tests/jax/conftest.py +++ b/tests/jax/conftest.py @@ -2,6 +2,7 @@ # # See LICENSE for license information. """conftest for tests/jax""" + import os import jax import pytest @@ -98,12 +99,7 @@ def pytest_addoption(parser): """ parser.addoption("--num-process", action="store", default=0) parser.addoption("--process-id", action="store", default=0) - parser.addoption( - "--use-cudnn-fusion", - choices=("0", "1"), - default="1", - help="Pass use_cudnn_fusion to MoE blocks and EP capacity calculation", - ) + def pytest_configure(config): config.addinivalue_line( diff --git a/tests/jax/run_te_ep_moe.sh b/tests/jax/run_te_ep_moe.sh index c8233605ec9..78622a85804 100755 --- a/tests/jax/run_te_ep_moe.sh +++ b/tests/jax/run_te_ep_moe.sh @@ -48,7 +48,7 @@ echo "Per-process logs: $LOG_DIR" PIDS=() EXITS=() -PHASE_FAILED=0 +FAILED=0 cleanup() { for pid in "${PIDS[@]:-}"; do @@ -65,90 +65,56 @@ cleanup() { } trap cleanup EXIT INT TERM -run_phase() { - local phase_name="$1" - local use_cudnn_fusion="$2" - shift 2 - local -a phase_args=("$@") - local phase_log_dir="$LOG_DIR/$phase_name" - mkdir -p "$phase_log_dir" - PIDS=() - EXITS=() - - echo - echo "============================================================" - echo "Phase: $phase_name" - echo " use_cudnn_fusion=$use_cudnn_fusion" - echo " phase pytest args : ${phase_args[*]:-}" - echo " logs : $phase_log_dir" - echo "============================================================" - - for i in $(seq 0 $((NUM_GPUS - 1))); do - local log_file="$phase_log_dir/proc_${i}.log" - local -a pytest_cmd=( - python3 -m pytest -c "$PYTEST_INI" - "$TEST_FILE" - -p no:typeguard - -v -s - --num-process="$NUM_GPUS" - --process-id="$i" - --use-cudnn-fusion="$use_cudnn_fusion" - "${phase_args[@]}" - ) - if [ "$i" -eq 0 ]; then - echo "=== Live output from process 0 ($phase_name) ===" - "${pytest_cmd[@]}" 2>&1 | tee "$log_file" & - else - "${pytest_cmd[@]}" > "$log_file" 2>&1 & - fi - PIDS+=("$!") - done - - for pid in "${PIDS[@]}"; do - if wait "$pid"; then - EXITS+=("0") - else - EXITS+=("$?") - fi - done - - echo - echo "Per-process exit codes for $phase_name:" - for i in "${!EXITS[@]}"; do - echo " proc $i -> ${EXITS[$i]}" - done +for i in $(seq 0 $((NUM_GPUS - 1))); do + log_file="$LOG_DIR/proc_${i}.log" + pytest_cmd=( + python3 -m pytest -c "$PYTEST_INI" + "$TEST_FILE" + -p no:typeguard + -v -s + --num-process="$NUM_GPUS" + --process-id="$i" + "$@" + ) + if [ "$i" -eq 0 ]; then + echo "=== Live output from process 0 ===" + "${pytest_cmd[@]}" 2>&1 | tee "$log_file" & + else + "${pytest_cmd[@]}" > "$log_file" 2>&1 & + fi + PIDS+=("$!") +done - local failed=0 - for e in "${EXITS[@]}"; do - if [ "$e" != "0" ] && [ "$e" != "5" ]; then - failed=1 - break - fi - done - if [ "$failed" -ne 0 ]; then - PHASE_FAILED=1 - echo "[run_te_ep_moe.sh] phase $phase_name FAILED" - echo " process 0 tail:" - tail -20 "$phase_log_dir/proc_0.log" 2>/dev/null || true +for pid in "${PIDS[@]}"; do + if wait "$pid"; then + EXITS+=("0") else - echo "[run_te_ep_moe.sh] phase $phase_name PASSED" + EXITS+=("$?") fi -} +done -# Keep ordinary CUDA C++ and cuDNN JAX coverage in separate process groups. -# TE EP/NCCL caches alignment process-wide; these use 128 and 256 respectively. -run_phase "ordinary" "0" -k "not TestTeEpMoeCudnnCutedslFusion" "$@" -run_phase "cutedsl" "1" -k "TestTeEpMoeCudnnCutedslFusion" "$@" +echo +echo "Per-process exit codes:" +for i in "${!EXITS[@]}"; do + echo " proc $i -> ${EXITS[$i]}" +done + +for e in "${EXITS[@]}"; do + if [ "$e" != "0" ] && [ "$e" != "5" ]; then + FAILED=1 + break + fi +done echo -if [ "$PHASE_FAILED" -eq 0 ]; then - echo "[run_te_ep_moe.sh] all phases PASSED" +if [ "$FAILED" -eq 0 ]; then + echo "[run_te_ep_moe.sh] all processes PASSED" if [ -z "${TE_EP_MOE_MP_LOG_DIR:-}" ]; then rm -rf "$LOG_DIR" fi exit 0 fi -echo "[run_te_ep_moe.sh] at least one phase FAILED" +echo "[run_te_ep_moe.sh] at least one process FAILED" echo " retaining logs at $LOG_DIR for diagnosis" exit 1 diff --git a/tests/jax/test_te_ep_moe.py b/tests/jax/test_te_ep_moe.py index a53b08fa5a0..7bc766dbe15 100644 --- a/tests/jax/test_te_ep_moe.py +++ b/tests/jax/test_te_ep_moe.py @@ -29,7 +29,8 @@ single run supports — shape, dtype, finiteness AND numerical parity vs a pure-JAX reference. Variations on the block are pytest parametrize values rather than separate test -classes: +classes. The module-scoped fusion parameter runs all unfused cases first, +then tears down EP and runs all cuDNN cases in the same processes: * ``test_forward`` covers BF16 and MXFP8 forward execution across a curated set of configurations (softmax/sigmoid scoring, optional @@ -56,7 +57,7 @@ import numpy as np import pytest -from jax.experimental import mesh_utils +from jax.experimental import mesh_utils, multihost_utils from jax.sharding import Mesh, NamedSharding, PartitionSpec as P from flax.linen import partitioning as nn_partitioning @@ -99,16 +100,6 @@ def _read_mp_options(): return num, pid -def _read_cudnn_fusion_option() -> bool: - for index, argument in enumerate(sys.argv): - if argument.startswith("--use-cudnn-fusion="): - return argument.split("=", 1)[1] == "1" - if argument == "--use-cudnn-fusion" and index + 1 < len(sys.argv): - return sys.argv[index + 1] == "1" - return True - - -_USE_CUDNN_FUSION = _read_cudnn_fusion_option() _MP_NUM_PROCESS, _MP_PROCESS_ID = _read_mp_options() _MP_ACTIVE = _init_distributed(_MP_NUM_PROCESS, _MP_PROCESS_ID) @@ -131,13 +122,11 @@ def _read_cudnn_fusion_option() -> bool: from transformer_engine.jax.flax import _MoEBlock as MoEBlock from transformer_engine.jax.moe import ( - _ALIGN_SIZE, - _CUDNN_JAX_ALIGN_SIZE, get_moe_recv_capacity_per_rank, moe, record_ep_bootstrap_signature_for_moe, ) -from transformer_engine.jax.ep import ep_bootstrap +from transformer_engine.jax.ep import ep_bootstrap, ep_finalize from transformer_engine.common.recipe import MXFP8BlockScaling from transformer_engine.jax.sharding import MeshResource, global_shard_guard @@ -200,8 +189,14 @@ def _read_cudnn_fusion_option() -> bool: # ----------------------------------------------------------------------------- +@pytest.fixture(scope="module", params=[False, True], ids=["unfused", "cudnn"], autouse=True) +def use_cudnn_fusion(request): + """Group every test by fusion mode before changing EP's cached alignment.""" + return request.param + + @pytest.fixture(scope="module") -def mesh(): +def mesh(use_cudnn_fusion): if jax.device_count() < NUM_DEVICES_REQUIRED: pytest.skip( f"Need >={NUM_DEVICES_REQUIRED} devices for ep={EP_SIZE} x fsdp={FSDP_SIZE};" @@ -218,13 +213,12 @@ def mesh(): # Worst-case recv capacity per rank # TODO(jberchtold) support configurations other than worst-case by refactoring tests # but if possible avoid bootstrap/teardown for each test - alignment = _CUDNN_JAX_ALIGN_SIZE if _USE_CUDNN_FUSION else _ALIGN_SIZE recv_capacity_per_rank = get_moe_recv_capacity_per_rank( num_experts=NUM_EXPERTS, num_experts_per_tok=TOPK, max_tokens_per_rank=max_tokens_per_rank, ep_size=EP_SIZE, - alignment=alignment, + use_cudnn_fusion=use_cudnn_fusion, ) # Eager bootstrap: ep_bootstrap does a host-side NCCL UID allgather @@ -247,7 +241,15 @@ def mesh(): hidden_dim=HIDDEN, ep_size=EP_SIZE, ) - return mesh_obj + mode = "cudnn" if use_cudnn_fusion else "unfused" + try: + yield mesh_obj + finally: + # Every rank finishes the mode before releasing cached EP executables. + jax.effects_barrier() + multihost_utils.sync_global_devices(f"te_moe_{mode}_finished") + ep_finalize() + multihost_utils.sync_global_devices(f"te_moe_{mode}_finalized") # ----------------------------------------------------------------------------- @@ -377,7 +379,7 @@ def _make_block( dispatch_checkpoint_name=None, quant_before_fsdp_ag=False, mesh_resource=None, - use_cudnn_fusion=_USE_CUDNN_FUSION, + use_cudnn_fusion, ): kwargs = dict( num_experts=NUM_EXPERTS, @@ -402,6 +404,11 @@ def _make_block( return MoEBlock(**kwargs) +@pytest.fixture +def make_block(use_cudnn_fusion): + return partial(_make_block, use_cudnn_fusion=use_cudnn_fusion) + + def _strong_expert_bias_init(key, shape, dtype): """Half +5, half -5 — large enough to force topk onto the +ve half.""" del key @@ -588,7 +595,7 @@ def _quantization_recipe(quantization): @pytest.mark.parametrize("native_weight_layout", [False, True]) @pytest.mark.parametrize("fsdp_dimension", ["k", "gated", "expert"]) def test_quantized_weight_gather_matches_full_precision_gather( - mesh, monkeypatch, native_weight_layout, fsdp_dimension + mesh, monkeypatch, native_weight_layout, fsdp_dimension, make_block ): """The FP8 weight gather retains forward and backward MoE semantics.""" if native_weight_layout or fsdp_dimension != "k": @@ -626,8 +633,8 @@ def layout_moe(*args, **kwargs): monkeypatch.setattr(flax_moe_module, "moe", layout_moe) x = _make_inputs(jax.random.PRNGKey(41)) - baseline = _make_block(quantization_recipe=MXFP8BlockScaling()) - quantized_ag = _make_block(quantization_recipe=MXFP8BlockScaling(), quant_before_fsdp_ag=True) + baseline = make_block(quantization_recipe=MXFP8BlockScaling()) + quantized_ag = make_block(quantization_recipe=MXFP8BlockScaling(), quant_before_fsdp_ag=True) variables, baseline_out, _ = _init_apply(baseline, mesh, x, jax.random.PRNGKey(42)) with _ctx(mesh): x_sh = _shard_inputs(x, mesh) @@ -654,9 +661,9 @@ def layout_moe(*args, **kwargs): @pytest.mark.parametrize("quant_before_fsdp_ag", [False, True]) -def test_mesh_resource_api_forward_and_backward(mesh, quant_before_fsdp_ag): +def test_mesh_resource_api_forward_and_backward(mesh, quant_before_fsdp_ag, make_block): """Explicit and global resources give identical forward and backward results.""" - baseline = _make_block( + baseline = make_block( quantization_recipe=MXFP8BlockScaling(), quant_before_fsdp_ag=quant_before_fsdp_ag ) candidate = baseline.clone( @@ -718,8 +725,8 @@ class TestTeEpMoeForward: @pytest.mark.parametrize("config", _CONFIGS) @pytest.mark.parametrize("quantization", _QUANTIZATION_CASES) - def test_forward(self, mesh, config, quantization): - block = _make_block(**config, quantization_recipe=_quantization_recipe(quantization)) + def test_forward(self, mesh, config, quantization, make_block): + block = make_block(**config, quantization_recipe=_quantization_recipe(quantization)) x = _make_inputs(jax.random.PRNGKey(0)) variables, output, aux = _init_apply(block, mesh, x, jax.random.PRNGKey(1)) @@ -758,8 +765,8 @@ class TestTeEpMoeBackward: @pytest.mark.parametrize("config", _CONFIGS) @pytest.mark.parametrize("quantization", _QUANTIZATION_CASES) - def test_backward(self, mesh, config, quantization): - block = _make_block(**config, quantization_recipe=_quantization_recipe(quantization)) + def test_backward(self, mesh, config, quantization, make_block): + block = make_block(**config, quantization_recipe=_quantization_recipe(quantization)) x = _make_inputs(jax.random.PRNGKey(2)) variables, _, _ = _init_apply(block, mesh, x, jax.random.PRNGKey(3)) grads_te, grad_x_te = _grad_step(block, variables, mesh, x) @@ -829,9 +836,7 @@ def loss_fn(params, x): ) -def test_ep_checkpoint_names(mesh, monkeypatch): - if _USE_CUDNN_FUSION: - pytest.skip("BF16 fallback uses a different EP alignment than the cuDNN bootstrap") +def test_ep_checkpoint_names(mesh, monkeypatch, make_block): moe_module = importlib.import_module("transformer_engine.jax.moe") named_values = {} original_checkpoint_name = moe_module.checkpoint_name @@ -841,7 +846,7 @@ def recorded_checkpoint_name(value, name): return original_checkpoint_name(value, name) monkeypatch.setattr(moe_module, "checkpoint_name", recorded_checkpoint_name) - block = _make_block( + block = make_block( dispatch_checkpoint_name="saved_dispatch", ) inputs = _make_inputs(jax.random.PRNGKey(34)) @@ -856,9 +861,11 @@ def recorded_checkpoint_name(value, name): assert np.all(np.isfinite(_to_global_numpy(_unwrap(grads["params"][name]), mesh))) -def test_explicitly_disabled_fusion_with_native_weights(mesh, monkeypatch): +def test_explicitly_disabled_fusion_with_native_weights( + mesh, monkeypatch, make_block, use_cudnn_fusion +): """Disabling fusion bypasses dependency probing and preserves native gradients.""" - if _USE_CUDNN_FUSION: + if use_cudnn_fusion: pytest.skip("Requires the ordinary 128-token EP bootstrap") from transformer_engine.jax import cpp_extensions as tex @@ -868,7 +875,7 @@ def unexpected_selection(*args, **kwargs): raise AssertionError("Explicitly disabled fusion must not probe cuDNN") monkeypatch.setattr(moe_module, "_select_cudnn_jax_fusion", unexpected_selection) - block = _make_block(quantization_recipe=MXFP8BlockScaling(), use_cudnn_fusion=False) + block = make_block(quantization_recipe=MXFP8BlockScaling(), use_cudnn_fusion=False) x = _make_inputs(jax.random.PRNGKey(53)) variables, baseline_output, _ = _init_apply(block, mesh, x, jax.random.PRNGKey(54)) baseline_grads, baseline_dx = _grad_step(block, variables, mesh, x) @@ -903,15 +910,17 @@ class TestTeEpMoeCudnnCutedslFusion: """End-to-end MXFP8 coverage for cuDNN's grouped GLU JAX APIs.""" @pytest.mark.parametrize("native_layout", [False, True]) - def test_missing_apis_unfused_forward_and_backward(self, mesh, monkeypatch, native_layout): + def test_missing_apis_unfused_forward_and_backward( + self, mesh, monkeypatch, native_layout, make_block, use_cudnn_fusion + ): """Unfused fallback retains bootstrap capacity and native-layout gradients.""" - if not _USE_CUDNN_FUSION: + if not use_cudnn_fusion: pytest.skip("Requires fusion requested at bootstrap") import cudnn from transformer_engine.jax import cpp_extensions as tex monkeypatch.setattr(cudnn, "grouped_gemm_glu_wrapper_sm100", None) - block = _make_block(quantization_recipe=MXFP8BlockScaling()) + block = make_block(quantization_recipe=MXFP8BlockScaling()) x = _make_inputs(jax.random.PRNGKey(51)) variables, baseline_output, _ = _init_apply(block, mesh, x, jax.random.PRNGKey(52)) baseline_grads, baseline_dx = _grad_step(block, variables, mesh, x) @@ -946,9 +955,11 @@ def native_moe(*args, **kwargs): ) @pytest.mark.parametrize("apply_topk_weights_early", [False, True]) - def test_mxfp8_forward_and_backward(self, mesh, apply_topk_weights_early, monkeypatch): - if not _USE_CUDNN_FUSION: - pytest.skip("run separately with --use-cudnn-fusion=1") + def test_mxfp8_forward_and_backward( + self, mesh, apply_topk_weights_early, monkeypatch, make_block, use_cudnn_fusion + ): + if not use_cudnn_fusion: + pytest.skip("requires the cuDNN fusion fixture parameter") from transformer_engine.jax import cpp_extensions as tex shared_calls = [] @@ -965,7 +976,7 @@ def checked_dglu(*args, **kwargs): monkeypatch.setattr(tex, "grouped_gemm_glu", checked_glu) monkeypatch.setattr(tex, "grouped_gemm_dglu", checked_dglu) - block = _make_block( + block = make_block( apply_topk_weights_early=apply_topk_weights_early, quantization_recipe=MXFP8BlockScaling(), ) @@ -1028,8 +1039,10 @@ def loss_fn(params, inputs): err_msg="d_x fused MXFP8 gradient parity breach", ) - def test_cudnn_fused_with_checkpoint_names(self, mesh, monkeypatch): - if not _USE_CUDNN_FUSION: + def test_cudnn_fused_with_checkpoint_names( + self, mesh, monkeypatch, make_block, use_cudnn_fusion + ): + if not use_cudnn_fusion: pytest.skip("cuDNN grouped GEMM fusion is disabled") from transformer_engine.jax import cpp_extensions as tex @@ -1063,7 +1076,7 @@ def checkpointed_moe(*args, **kwargs): return original_moe(*args, **kwargs) monkeypatch.setattr(flax_moe_module, "moe", checkpointed_moe) - block = _make_block( + block = make_block( quantization_recipe=MXFP8BlockScaling(), dispatch_checkpoint_name="saved_dispatch", ) @@ -1093,9 +1106,9 @@ class TestTeEpMoeAuxLoss: finite + non-zero per tensor. """ - def test_aux_loss(self, mesh): + def test_aux_loss(self, mesh, make_block): coeff = 1e-2 - block = _make_block(aux_loss_coeff=coeff) + block = make_block(aux_loss_coeff=coeff) x = _make_inputs(jax.random.PRNGKey(20)) variables, _, aux = _init_apply(block, mesh, x, jax.random.PRNGKey(21)) @@ -1135,10 +1148,10 @@ def test_aux_loss(self, mesh): assert np.all(np.isfinite(g_gate)), "gate grad NaN/Inf under aux-only loss" assert np.any(g_gate != 0.0), "aux bwd should propagate to gate_kernel" - def test_combined_loss_grads(self, mesh): + def test_combined_loss_grads(self, mesh, make_block): """Joint main + aux loss bwd: per-tensor finite + non-zero in one pass.""" - block = _make_block(aux_loss_coeff=1e-2) + block = make_block(aux_loss_coeff=1e-2) x = _make_inputs(jax.random.PRNGKey(22)) variables, _, _ = _init_apply(block, mesh, x, jax.random.PRNGKey(23)) grads, _ = _grad_step(block, variables, mesh, x, include_aux=True) From 48208161799a92b355bba389d18d689a8d02eaeb Mon Sep 17 00:00:00 2001 From: Jeremy Berchtold Date: Fri, 9 Oct 2026 09:48:27 -0700 Subject: [PATCH 38/38] Simplify unified GLU and dGLU fusion fallback tests Signed-off-by: Jeremy Berchtold --- tests/jax/test_custom_call_compute.py | 55 ++++++++++++++------------- 1 file changed, 28 insertions(+), 27 deletions(-) diff --git a/tests/jax/test_custom_call_compute.py b/tests/jax/test_custom_call_compute.py index 0a46740e1db..cf33a092ae7 100644 --- a/tests/jax/test_custom_call_compute.py +++ b/tests/jax/test_custom_call_compute.py @@ -2181,7 +2181,7 @@ def frontend(self, monkeypatch): original = importlib.import_module def import_module(name, package=None): - if name in ("cudnn", "cudnn.jax"): + if name == "cudnn": return module if name == "cutlass.jax": return SimpleNamespace(is_available=lambda: True) @@ -2222,23 +2222,24 @@ def test_api_signature(self, frontend, operation, change): assert available == (change == "optional") assert bool(reason) == (change != "optional") - @pytest.mark.parametrize("capability", [90, 100, 103, 107, 120]) @pytest.mark.parametrize( - "missing", - [None, "grouped_gemm_glu_wrapper_sm100", "grouped_gemm_dglu_wrapper_sm100", "all"], + "capability,missing,expected", + [ + (90, None, False), + (100, None, True), + (103, None, True), + (107, None, True), + (120, None, True), + (100, "grouped_gemm_glu_wrapper_sm100", False), + (100, "grouped_gemm_dglu_wrapper_sm100", False), + ], ) - def test_ordered_fallback(self, frontend, monkeypatch, capability, missing): - if missing == "all": - frontend.grouped_gemm_glu_wrapper_sm100 = None - frontend.grouped_gemm_dglu_wrapper_sm100 = None - elif missing: + def test_fused_or_unfused_selection(self, frontend, monkeypatch, capability, missing, expected): + if missing: delattr(frontend, missing) monkeypatch.setattr( transformer_engine_jax, "get_device_compute_capability", lambda _: capability ) - expected = False - if capability >= 100 and missing is None: - expected = True with warnings.catch_warnings(record=True) as caught: warnings.simplefilter("always") assert _moe_module._select_cudnn_jax_fusion([]) == expected @@ -2278,22 +2279,22 @@ def fail(_): with pytest.warns(UserWarning, match="device query failed"): assert _moe_module._select_cudnn_jax_fusion([]) is False - def test_installed_frontend_contract(self): - available, reason = _glu_adapter.grouped_gemm_glu_dependencies_available() - assert available or reason - - @pytest.mark.parametrize("request_fusion", [None, True, False]) - @pytest.mark.parametrize("native_layout", [False, True]) - @pytest.mark.parametrize("square", [False, True]) @pytest.mark.parametrize( - "missing,expected", + "request_fusion,missing,native_layout,square,expected", [ - (None, True), - ("grouped_gemm_glu_wrapper_sm100", False), - ("grouped_gemm_dglu_wrapper_sm100", False), + (None, None, False, False, True), + (True, None, False, False, True), + (False, None, False, False, False), + (True, "grouped_gemm_glu_wrapper_sm100", False, False, False), + (True, "grouped_gemm_dglu_wrapper_sm100", False, False, False), + (False, "grouped_gemm_glu_wrapper_sm100", False, False, False), + (False, "grouped_gemm_dglu_wrapper_sm100", False, False, False), + (True, None, True, False, True), + (True, None, False, True, True), + (True, None, True, True, True), ], ) - def test_public_moe_passes_selected_path( + def test_public_moe_passes_fusion_bool( self, frontend, monkeypatch, native_layout, square, missing, expected, request_fusion ): from jax.sharding import Mesh @@ -2325,8 +2326,6 @@ def execute(*args): if square: wi = jnp.ones((2, 128, 128), jnp.bfloat16) kwargs["cudnn_native_weight_layout"] = native_layout - if request_fusion is False: - expected = False with warnings.catch_warnings(record=True) as caught: warnings.simplefilter("always") output, _, _ = _moe_module.moe( @@ -2345,7 +2344,9 @@ def execute(*args): @pytest.mark.parametrize("native_layout", [False, True]) @pytest.mark.parametrize("gather_gated_dimension", [False, True]) - def test_forward_uses_selected_kernel(self, monkeypatch, native_layout, gather_gated_dimension): + def test_fused_forward_uses_shared_glu( + self, monkeypatch, native_layout, gather_gated_dimension + ): class SelectedKernel(Exception): pass