[JAX] Optimize MoE block - #3354
Conversation
Greptile SummaryThe 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.
Confidence Score: 4/5The PR is not yet safe to merge because existing The current module requests one Files Needing Attention: transformer_engine/jax/flax/moe.py Important Files Changed
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]
Reviews (9): Last reviewed commit: "Merge branch 'main' into jberchtold/moeb..." | Re-trigger Greptile |
072422b to
0645751
Compare
|
/te-ci L1 jax |
d9fa28c to
bc28b24
Compare
|
/te-ci L1 jax |
ca32c8d to
ec09e5a
Compare
|
/te-ci L1 jax |
79d7cda to
a5ecb6c
Compare
|
LGTM pending CI |
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
c387772 to
b830367
Compare
|
/te-ci L1 jax |
Signed-off-by: Jeremy Berchtold <jberchtold@nvidia.com>
|
/te-ci L1 jax |
* 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)
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
Changes
jnp.wheremasking that wasn't required as TE EP and grouped GEMM are group-awareChecklist: