Skip to content

[JAX] Optimize MoE block - #3354

Merged
jberchtold-nvidia merged 6 commits into
NVIDIA:mainfrom
jberchtold-nvidia:jberchtold/moeblock-debug
Aug 27, 2026
Merged

jberchtold-nvidia merged 6 commits into
NVIDIA:mainfrom
jberchtold-nvidia:jberchtold/moeblock-debug

Conversation

@jberchtold-nvidia

@jberchtold-nvidia jberchtold-nvidia commented Aug 12, 2026 •

Copy link
Copy Markdown
Collaborator

Description

Improves performance of the MoE block by exposing support for quantization, removal of unnecessary masking overheads, and support for less memory usage via a reduced receive capacity in TE EP

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

  • Direct support for MXFP8 quantization in the MoE block along with corresponding tests
  • Removal of additional overheads like jnp.where masking that wasn't required as TE EP and grouped GEMM are group-aware
  • Support for a reduced receive capacity and integration with TE EP's overflow detection

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
jberchtold-nvidia marked this pull request as draft August 12, 2026 14:47
@greptile-apps

greptile-apps Bot commented Aug 12, 2026 •

Copy link
Copy Markdown
Contributor

Greptile Summary

The PR optimizes the JAX expert-parallel MoE path by adding MXFP8 grouped quantization, reducing receive-capacity memory usage, removing redundant padding masks, and storing the gated FC1 weights contiguously.

  • Adds configurable receive capacity and overflow reporting to the functional and Flax MoE interfaces.
  • Threads grouped quantizer sets through the MoE forward and backward paths.
  • Expands distributed tests to cover BF16 and MXFP8 numerical behavior.

Confidence Score: 4/5

The PR is not yet safe to merge because existing _MoEBlock checkpoints and optimizer states cannot be restored after the FC1 parameter-tree replacement.

The current module requests one wi leaf where prior checkpoints contain wi_0 and wi_1, and current HEAD provides no migration or compatibility path.

Files Needing Attention: transformer_engine/jax/flax/moe.py

Important Files Changed

Filename Overview
transformer_engine/jax/flax/moe.py Adds recipe-driven grouped quantizers, receive-capacity configuration, and a contiguous FC1 parameter layout to the Flax MoE adapter.
transformer_engine/jax/moe.py Reworks the expert FFN forward and backward paths for grouped quantization, explicit capacity sizing, overflow reporting, and global sharding.
transformer_engine/jax/cpp_extensions/quantization.py Allows global stateless MXFP8 grouped-quantizer descriptors to cover shard-local group counts.
transformer_engine/jax/quantize/tensor.py Updates quantized tensor checkpoint handling used by the MoE custom differentiation path.
tests/jax/test_te_ep_moe.py Extends distributed MoE numerical tests to BF16 and MXFP8 configurations and the contiguous FC1 layout.

Flowchart

%%{init: {'theme': 'neutral'}}%%
flowchart LR
  A[Sharded hidden states] --> B[Gate and top-k routing]
  B --> C[EP prepare and dispatch]
  C --> D[Grouped FC1 quantization and GEMM]
  D --> E[Activation]
  E --> F[Grouped FC2 quantization and GEMM]
  F --> G[EP combine]
  C -. receive demand .-> H[Overflow detection]
Loading

Reviews (9): Last reviewed commit: "Merge branch 'main' into jberchtold/moeb..." | Re-trigger Greptile

Comment thread transformer_engine/jax/flax/moe.py
@jberchtold-nvidia
jberchtold-nvidia force-pushed the jberchtold/moeblock-debug branch 4 times, most recently from 072422b to 0645751 Compare August 13, 2026 15:49
@nvMelissa nvMelissa added the 2.19 label Aug 13, 2026
@jberchtold-nvidia

Copy link
Copy Markdown
Collaborator Author

/te-ci L1 jax

@jberchtold-nvidia
jberchtold-nvidia force-pushed the jberchtold/moeblock-debug branch from d9fa28c to bc28b24 Compare August 13, 2026 23:47
@jberchtold-nvidia

Copy link
Copy Markdown
Collaborator Author

/te-ci L1 jax

@jberchtold-nvidia
jberchtold-nvidia force-pushed the jberchtold/moeblock-debug branch from ca32c8d to ec09e5a Compare August 14, 2026 14:13
@jberchtold-nvidia

Copy link
Copy Markdown
Collaborator Author

/te-ci L1 jax

@jberchtold-nvidia
jberchtold-nvidia marked this pull request as ready for review August 14, 2026 17:55
@jberchtold-nvidia jberchtold-nvidia changed the title [DRAFT][JAX] Optimize MoE block [JAX] Optimize MoE block Aug 14, 2026
Comment thread tests/jax/test_te_ep_moe.py Outdated
Comment thread transformer_engine/jax/moe.py Outdated
Comment thread transformer_engine/jax/flax/moe.py
Comment thread transformer_engine/jax/moe.py Outdated
Comment thread transformer_engine/jax/moe.py Outdated
Comment thread transformer_engine/jax/moe.py Outdated
Comment thread transformer_engine/jax/moe.py
Comment thread tests/jax/test_te_ep_moe.py Outdated
@jberchtold-nvidia
jberchtold-nvidia force-pushed the jberchtold/moeblock-debug branch from 79d7cda to a5ecb6c Compare August 24, 2026 21:42
@tdophung

Copy link
Copy Markdown
Collaborator

LGTM pending CI

tdophung
tdophung previously approved these changes Aug 24, 2026
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
@jberchtold-nvidia
jberchtold-nvidia force-pushed the jberchtold/moeblock-debug branch from c387772 to b830367 Compare August 24, 2026 22:55
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
@jberchtold-nvidia

Copy link
Copy Markdown
Collaborator Author

/te-ci L1 jax

@jberchtold-nvidia

Copy link
Copy Markdown
Collaborator Author

/te-ci L1 jax

@jberchtold-nvidia
jberchtold-nvidia merged commit 2fd4604 into NVIDIA:main Aug 27, 2026
13 of 17 checks passed
fheinecke pushed a commit that referenced this pull request Aug 27, 2026
* TE/JAX MoEBlock optimizations

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

* Fix lint

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

* Fix lint

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

* Keep arch guard consistent in EP tests

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

---------

Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
Co-authored-by: Teddy Do <tdophung@nvidia.com>
(cherry picked from commit 2fd4604)
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants