Skip to content

[Draft][JAX] MoE block optimization - #3638

Draft
jberchtold-nvidia wants to merge 36 commits into
NVIDIA:mainfrom
jberchtold-nvidia:jberchtold/rubin-gmm-fusion
Draft

jberchtold-nvidia wants to merge 36 commits into
NVIDIA:mainfrom
jberchtold-nvidia:jberchtold/rubin-gmm-fusion

Conversation

@jberchtold-nvidia

Copy link
Copy Markdown
Collaborator

Description

Please include a brief summary of the changes, relevant motivation and context.

Fixes # (issue)

Type of change

  • Documentation change (change only to the documentation, either a fix or a new content)
  • Bug fix (non-breaking change which fixes an issue)
  • New feature (non-breaking change which adds functionality)
  • Breaking change (fix or feature that would cause existing functionality to not work as expected)
  • Infra/Build change
  • Code refactoring

Changes

Please list the changes introduced in this PR:

  • Change A
  • Change B

Checklist:

  • I have read and followed the contributing guidelines
  • The functionality is complete
  • I have commented my code, particularly in hard-to-understand areas
  • I have made corresponding changes to the documentation
  • My changes generate no new warnings
  • I have added tests that prove my fix is effective or that my feature works
  • New and existing unit tests pass locally with my changes

jberchtold-nvidia and others added 29 commits July 23, 2026 14:37
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.
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.
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
…jberchtold/rubin-gmm-fusion

Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>

# Conflicts:
#	transformer_engine/jax/moe.py
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
…moe' into jberchtold/rubin-gmm-fusion

Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>

# Conflicts:
#	tests/jax/test_te_ep_moe.py
#	transformer_engine/jax/flax/moe.py
#	transformer_engine/jax/moe.py
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
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 <jberchtold@nvidia.com>
@jberchtold-nvidia
jberchtold-nvidia marked this pull request as draft October 6, 2026 20:05
@greptile-apps

greptile-apps Bot commented Oct 6, 2026 •

Copy link
Copy Markdown
Contributor

RetriggerConfidence Score: 2/5

[High risk] Adds MoE optimization paths with new kernel dispatch and weight handling.

Fix the bootstrap axis count, expert-weight ordering, and optional-dependency test failure before merging.

Findings

  1. P1 Separate DP and FSDP fail ▶
  2. P1 Tokens use wrong experts ▶
  3. P1 Optional package breaks tests ▶
  4. P2 Fusion switch does nothing ▶

Summary

Adds cuDNN grouped GEMM/SwiGLU execution, quantized FSDP weight gathering, resource-based MoE configuration, and checkpoint names.

  • Separate DP and FSDP axes still conflict with bootstrap output sizes.
  • FSDP-first expert sharding can select the wrong expert weights.
  • The default test launcher now fails without optional cuDNN JAX support.
  • The new environment-variable documentation does not match the implementation.

Diagram

%%{init: {'theme': 'neutral'}}%%
flowchart TD
  A[MoE call and mesh resources] --> B[Route tokens]
  B --> C[EP dispatch]
  W[Expert weight shards] --> Q{Quantize before gather?}
  Q -->|Yes| G[Quantize local shards and gather data and scales]
  Q -->|No| H[Gather full-precision weights and quantize]
  G --> F{Selected FFN path}
  H --> F
  C --> F
  F -->|cuDNN| D[Fused grouped GEMM and SwiGLU]
  F -->|Fallback| E[TE grouped GEMM and activation]
  D --> I[Down projection]
  E --> I
  I --> J[EP combine]
Loading

Reviews (1) · Last reviewed commit: "Preserve gate/up pairing in quantized JA..."

Comment on lines +1331 to +1334
dp_size = 1
for ax in data_parallelism_axes:
dp_size *= mesh.shape[ax]
num_procs = num_ep * dp_size

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Separate DP and FSDP fail

With distinct DP and FSDP axes larger than one, moe() now shards batches and EP outputs across both axes. But ep_bootstrap() still calculates num_ep_groups from just one outer axis. On a DP2/FSDP2/EP2 mesh, EP outputs have leading size 4 but are partitioned across 8 devices, so compilation fails. Update bootstrap group counting alongside the new resource API.

Comment on lines +1351 to +1353
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)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Tokens use wrong experts

When logical rules place ("fsdp", "ep") on the expert dimension, gathering only FSDP does not recover the contiguous expert slice that EP routing expects. With 8 experts and EP2/FSDP2, EP rank 0 gathers experts [0, 1, 4, 5] instead of [0, 1, 2, 3]. The code accepts this order and uses those weights as EP-local experts, producing wrong outputs. Require EP before FSDP on this dimension, or restore the expected expert ownership before running the FFN.

# 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" "$@"

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Optional package breaks tests

The launcher now always runs the cuDNN phase, but its tests import cudnn.jax without a skip. If that optional module is absent on a supported multi-GPU machine, the previously usable launcher fails with ModuleNotFoundError even when ordinary TE execution works. Skip the cuDNN-only tests when the required APIs are unavailable, or make this phase explicitly opt-in.

Comment thread docs/envvars.rst Outdated
Comment on lines +559 to +568
.. 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.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Fusion switch does nothing

This entry documents an environment variable that the implementation never reads. It also says fusion defaults off and backward uses ordinary TE execution, while use_cudnn_fusion defaults to True and fused backward calls grouped_gemm_dswiglu. Replace this entry with the actual argument and execution behavior so users do not rely on an ineffective switch.

Note: If this suggestion doesn't match your team's coding style, reply to this and let me know. I'll remember it for next time!

Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants