Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
36 commits
Select commit Hold shift + click to select a range
f6ee566
Add quantized shard-map MoE support
jberchtold-nvidia Jul 8, 2026
1ae5c3d
JAX: integrate cuDNN CuTeDSL fused MoE
tdophung Jul 23, 2026
8a9f704
JAX: add TVM-FFI grouped GEMM SwiGLU bridge
tdophung Jul 23, 2026
f7abdef
Use TVM-FFI for fused dSwiGLU backward
tdophung Jul 31, 2026
5c0a4b2
Merge upstream main and use cuDNN grouped SwiGLU JAX API
jberchtold-nvidia Sep 15, 2026
d3199ae
Align external MoE bootstrap with cuDNN fusion
jberchtold-nvidia Sep 15, 2026
4f278d8
Add temporary cuDNN grouped GEMM fusion flag
jberchtold-nvidia Sep 15, 2026
4f71efd
Add cuDNN grouped dSwiGLU backward fusion
jberchtold-nvidia Sep 15, 2026
98b88a2
Fix SM100 version check
jberchtold-nvidia Sep 22, 2026
a4e2c7c
Update cuDNN grouped SwiGLU JAX API
jberchtold-nvidia Sep 22, 2026
abd0405
Route Rubin MoE grouped GLU through cuDNN JAX
jberchtold-nvidia Sep 28, 2026
bffeab0
Add MoE activation checkpoint names
jberchtold-nvidia Sep 29, 2026
effe0a0
[JAX] Add quantized MoE weight gather policy
jberchtold-nvidia Sep 30, 2026
e2830fd
Merge remote-tracking branch 'github/jberchtold/moe-checkpoint' into …
jberchtold-nvidia Oct 1, 2026
358908e
Checkpoint fused JAX MoE projections
jberchtold-nvidia Oct 1, 2026
00fe8b7
Checkpoint fused MoE combined projection without repacking
jberchtold-nvidia Oct 1, 2026
e254e04
Use expert-packed scales for fused JAX MoE GEMMs
jberchtold-nvidia Oct 1, 2026
89b19c4
Checkpoint fused cuDNN MoE quantized outputs
jberchtold-nvidia Oct 1, 2026
dfde5f3
Add EP dispatch and combine checkpoint names to JAX MoE
jberchtold-nvidia Oct 1, 2026
1410d52
Add native cuDNN grouped GEMM weight layout
jberchtold-nvidia Oct 1, 2026
68d6ef0
Reject cuDNN-native MoE weights without fusion
jberchtold-nvidia Oct 1, 2026
0d2326b
Remove redundant MoE combine checkpoint name
jberchtold-nvidia Oct 1, 2026
7e9cc14
Merge remote-tracking branch 'github/jberchtold/quant-before-fsdp-ag-…
jberchtold-nvidia Oct 5, 2026
791e8ae
Use MeshResource for JAX MoE parallelism and quantized weight gathering
jberchtold-nvidia Oct 5, 2026
f63ebbb
Support expert-axis quantized FSDP gathers in JAX MoE
jberchtold-nvidia Oct 5, 2026
ccf8a9d
Add capability-based fallbacks for cuDNN JAX MoE fusion
jberchtold-nvidia Oct 6, 2026
ae3ec09
Replace the temporary MoE fusion flag with an explicit boolean
jberchtold-nvidia Oct 6, 2026
fdc1e3c
Disambiguate JAX MoE weight layouts and fix FSDP fallback axes
jberchtold-nvidia Oct 6, 2026
600beda
Preserve gate/up pairing in quantized JAX MoE FSDP gathers
jberchtold-nvidia Oct 6, 2026
a76e2f9
Replace JAX MoE WeightGather with a boolean flag
jberchtold-nvidia Oct 6, 2026
c7c2651
Remove documentation for unused JAX MoE fusion environment variable
jberchtold-nvidia Oct 6, 2026
c31ddb0
Remove obsolete CuTeDSL MoE Shardy reproducer
jberchtold-nvidia Oct 6, 2026
fd250a2
Remove unused CuTeDSL JAX requirements file
jberchtold-nvidia Oct 6, 2026
1949887
Consolidate SwiGLU fallback tests into custom call compute suite
jberchtold-nvidia Oct 6, 2026
d739211
Move MoE FSDP layout tests into dedicated subdirectory
jberchtold-nvidia Oct 6, 2026
f91fc49
Move MoE resource API tests into dedicated subdirectory
jberchtold-nvidia Oct 6, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 42 additions & 0 deletions docs/api/jax.rst
Original file line number Diff line number Diff line change
Expand Up @@ -59,3 +59,45 @@ 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.

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`` 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``.
7 changes: 6 additions & 1 deletion tests/jax/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Loading
Loading