Skip to content

TE MoE Return Expert Counts and Custom Routing for DSv4 Support - #3596

Open
dinodeep wants to merge 3 commits into
NVIDIA:mainfrom
dinodeep:dinodeep/te-moe-block-dsv4-support
Open

dinodeep wants to merge 3 commits into
NVIDIA:mainfrom
dinodeep:dinodeep/te-moe-block-dsv4-support

Conversation

@dinodeep

Copy link
Copy Markdown
Contributor

Description

This PR adds support for DSv4 features in TE MoE block which includes the following

  • Providing expert count metadata (this is useful for users to perform routing bias updates, as DSv4 does)
  • Adding support for custom routing policies with moe_from_routing (allows for hash routing)

These changes will be useful to make TE MoE block more flexible and configurable for new, upcoming models such as DSv4.

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:

  • If desired by the user, allow the caller of moe to pass collect_expert_counts=True which will compute the expert counts and return them to the caller as a part of the metadata.
    • Tested by validating expert counts in MoE block tests
  • A new API function moe_from_routing(...) which allows the user to provide their expert assignments and weights. Given this routing information, moe_from_routing will re-use functionality of moe except the gate computation. Backwards gradients will be computed through weight.
    • Tested by implementing a simple hash routing in test and validating the TE MoE block with the pure JAX implementation
    • Had to slightly increase tolerance to adjust for accumulation differences

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

@github-actions github-actions Bot added the community-contribution PRs from external contributor outside the core maintainers, representing community-driven work. label Sep 30, 2026
@greptile-apps

greptile-apps Bot commented Sep 30, 2026 •

Copy link
Copy Markdown
Contributor

RetriggerConfidence Score: 3/5

[Medium risk] Mixture-of-Experts layer gains expert counting and custom routing modes.

The PR does not appear safe to merge while existing callers can fail on the changed return shape and externally routed backward remains unresolved.

Findings

  1. P1 Existing callers cannot unpack results ▶
  2. P1 Externally routed backward fails ▶

Summary

The PR adds per-expert assignment counts and supports caller-supplied routing alongside TE-computed routing.

  • Refactors the Flax MoE block to pass router configuration through RouterComputationInfo.
  • Adds tests for expert counts, hash routing, and gradients.

Diagram

%%{init: {'theme': 'neutral'}}%%
flowchart LR
  A[MoE inputs] --> B{Router info}
  B -->|RouterComputationInfo| C[Gate and top-k]
  B -->|RoutingMapInfo| D[Caller-supplied indices and weights]
  C --> E[Expert dispatch]
  D --> E
  E --> F[Expert FFN]
  F --> G[Combine and return output, loss, receive total, counts]
Loading

Reviews (11) · Last reviewed commit: "Fixing tolerance assignment" · Reviewed by Greptile

)
if aux_loss_coeff <= 0.0:
aux_loss = None
assert output.dtype == x.dtype, f"moe() output dtype {output.dtype} != input dtype {x.dtype}"

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 Existing callers cannot unpack results

moe() now returns four values instead of three, and _MoEBlock passes that result through. Existing callers that unpack the previous three-value result will raise ValueError: too many values to unpack, making this a breaking change to the public API.

Comment thread transformer_engine/jax/moe.py Outdated
@greptile-apps

This comment has been minimized.

@ptrendx

ptrendx commented Sep 30, 2026

Copy link
Copy Markdown
Member

@dinodeep Please sign every commit in this PR. See https://github.com/NVIDIA/TransformerEngine/blob/main/CONTRIBUTING.rst#sign-your-work

Comment thread transformer_engine/jax/moe.py
Comment thread transformer_engine/jax/moe.py Outdated
Comment thread transformer_engine/jax/moe.py Outdated
Comment on lines +1214 to +1216
d_router_info = RoutingMapInfo(
routing_indices=None,
routing_weights=d_topk_w.astype(ctx.routing_weights_dtype),

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 Externally routed backward fails

When a caller differentiates moe() with RoutingMapInfo, the backward rule returns routing_indices=None even though the forward call supplied an integer array. That changes the JAX pytree structure of the router argument, so the custom VJP cannot match the gradient to its input and the backward pass fails. Return a structurally matching cotangent for the indices.

@dinodeep
dinodeep force-pushed the dinodeep/te-moe-block-dsv4-support branch 2 times, most recently from ae0d2ee to a1bbdda Compare October 2, 2026 16:27
Comment thread tests/jax/test_te_ep_moe.py Outdated
Comment thread transformer_engine/jax/moe.py
Comment thread transformer_engine/jax/moe.py
@dinodeep
dinodeep force-pushed the dinodeep/te-moe-block-dsv4-support branch from a1bbdda to d525cb1 Compare October 4, 2026 16:52
@jberchtold-nvidia

Copy link
Copy Markdown
Collaborator

/te-ci L1 jax

1 similar comment
@jberchtold-nvidia

Copy link
Copy Markdown
Collaborator

/te-ci L1 jax

@greptile-apps

greptile-apps Bot commented Oct 7, 2026

Copy link
Copy Markdown
Contributor

Want your agent to iterate on Greptile's feedback? Try greploops.

Signed-off-by: Deep Patel <deepatel@nvidia.com>
Signed-off-by: Deep Patel <deepatel@nvidia.com>
Signed-off-by: Deep Patel <deepatel@nvidia.com>
@dinodeep
dinodeep force-pushed the dinodeep/te-moe-block-dsv4-support branch from d04a1d4 to 283eb34 Compare October 7, 2026 16:09
@jberchtold-nvidia

Copy link
Copy Markdown
Collaborator

/te-ci L1 jax

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

community-contribution PRs from external contributor outside the core maintainers, representing community-driven work.

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants