Repository navigation
[Draft][JAX] MoE block optimization - #3638
jberchtold-nvidia wants to merge 36 commits into
Conversation
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>
|
| dp_size = 1 | ||
| for ax in data_parallelism_axes: | ||
| dp_size *= mesh.shape[ax] | ||
| num_procs = num_ep * dp_size |
There was a problem hiding this comment.
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.
| 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) |
There was a problem hiding this comment.
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" "$@" |
There was a problem hiding this comment.
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.
| .. 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. |
There was a problem hiding this comment.
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>
Description
Please include a brief summary of the changes, relevant motivation and context.
Fixes # (issue)
Type of change
Changes
Please list the changes introduced in this PR:
Checklist: