NVIDIA / NVIDIA/TransformerEngine
[PyTorch] Reduce CPU overhead in grouped MLP block
@vthumbe1503 is already working on this.
Since Apr 16, 2026.
- Dominant language
- Python
- Stars
- 3.5k
- Forks
- 831
- Avg merge
- 3d 11h
- Merged PRs (30d)
- 65
Description
Is your feature request related to a problem? Please describe.
The fused operations for the grouped MLP block (see ForwardGroupedMLP_CuTeGEMMSwiGLU_MXFP8 and BackwardGroupedMLP_CuTeGEMMDSwiGLU_MXFP8, first added in https://github.com/NVIDIA/TransformerEngine/pull/2769) has significant CPU overhead. When I run a basic benchmark (on GB200 with 64 experts and 128 hidden size), I find the forward pass takes ~1.2 ms and the backward pass takes ~2.1 ms.
Describe the solution you'd like
Based on profiling, here are some rough estimates for some slow sections:
- wgrad tensor allocations: 350 us
-
tex.get_device_pointer_for_data_and_scales: >100 us, with ~50 us before enteringnvte_multi_tensor_swizzle_scaling_factors -
clear_tensor_data: 200 us - Initializing weight quantizers in the forward pass: 90 us
- Tensor reshapes before and after cuDNN kernels: ~50 us
- cuDNN group GEMM kernels: ~150 us
As we make optimizations, we should also adapt them to the unfused grouped linear op, and generally consider cleanups and refactors.
Describe alternatives you've considered
Additional context
Contributor guide
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Assessment
This issue has not been assessed yet.