Repository navigation
Conversation
|
| ) | ||
| if aux_loss_coeff <= 0.0: | ||
| aux_loss = None | ||
| assert output.dtype == x.dtype, f"moe() output dtype {output.dtype} != input dtype {x.dtype}" |
There was a problem hiding this comment.
This comment has been minimized.
This comment has been minimized.
|
@dinodeep Please sign every commit in this PR. See https://github.com/NVIDIA/TransformerEngine/blob/main/CONTRIBUTING.rst#sign-your-work |
| d_router_info = RoutingMapInfo( | ||
| routing_indices=None, | ||
| routing_weights=d_topk_w.astype(ctx.routing_weights_dtype), |
There was a problem hiding this comment.
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.
ae0d2ee to
a1bbdda
Compare
a1bbdda to
d525cb1
Compare
|
/te-ci L1 jax |
1 similar comment
|
/te-ci L1 jax |
|
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>
d04a1d4 to
283eb34
Compare
|
/te-ci L1 jax |
Description
This PR adds support for DSv4 features in TE MoE block which includes the following
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
Changes
Please list the changes introduced in this PR:
moeto passcollect_expert_counts=Truewhich will compute the expert counts and return them to the caller as a part of the metadata.moe_from_routing(...)which allows the user to provide their expert assignments and weights. Given this routing information,moe_from_routingwill re-use functionality ofmoeexcept the gate computation. Backwards gradients will be computed through weight.Checklist: