Skip to content

[Common][PyTorch] EP dispatch with unfused MXFP8 quantization - #3270

Open
phu0ngng wants to merge 11 commits into
NVIDIA:mainfrom
phu0ngng:ep_mxfp8
Open

[Common][PyTorch] EP dispatch with unfused MXFP8 quantization#3270
phu0ngng wants to merge 11 commits into
NVIDIA:mainfrom
phu0ngng:ep_mxfp8

Conversation

@phu0ngng

Copy link
Copy Markdown
Collaborator

Description

This PR adds MXFP8 support to the dispatch op of the NCCL EP path. The dispatch op is used in two places, and MXFP8 applies to both:

  • Dispatch forward bfloat16 tokens are quantized to MXFP8 internally and dispatched to the target experts; recv is returned as a per-expert GroupedTensor.
  • Combine backward the result-grad is scattered back to expert positions through the same (reverse) dispatch op, quantized to MXFP8, returning the expert-output grad as a per-expert GroupedTensor.

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

PyTorch frontend (transformer_engine/pytorch/ep.py, distributed.py, csrc/extensions/ep.cpp)**

  • The dispatch op quantizes bfloat16 tokens to MXFP8 internally when the buffer's dispatch_quant_recipe is set (MXFP8BlockScaling only for now); dispatch-forward recv is returned as a per-expert GroupedTensor. A pre-quantized input is rejected.
  • Combine backward reuses the dispatch op to scatter the result-grad: it quantizes the grad to MXFP8 and returns the expert-output grad as a per-expert GroupedTensor. Combine forward is unchanged (high-precision).
  • Recv data and block scales share a single caller-supplied (optionally symm-mem-backed) buffer, sliced into data-then-scale regions; the same convention is used for the combine backward grad buffer.

Common backend (common/ep/ep_backend.cpp, include/.../ep.h, comm_window.h)**

  • Backend and public headers extended to carry block-scale buffers/windows through the dispatch primitive.

NCCL EP submodule**

  • Bumped 3rdparty/nccl-extensions to the revision providing block-scaled dispatch.

Tests (tests/cpp_distributed/test_ep.cu, tests/pytorch/distributed/run_ep.py, run_test_ep.sh)**

  • Added C++ distributed coverage for the MXFP8 dispatch path.
  • Added PyTorch MXFP8 test passes for dispatch forward (normal, zero-copy, eager IO modes) and combine backward, gated behind a dedicated NVTE_EP_MXFP8_PASS run since the grouped path pins the per-expert alignment process-wide.

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

@greptile-apps

greptile-apps Bot commented Jul 28, 2026

Copy link
Copy Markdown
Contributor

Greptile Summary

Adds unfused MXFP8 quantization to NCCL expert-parallel dispatch and combine backward.

  • Extends the common EP API and PyTorch bindings to communicate MXFP8 data with block-scale buffers.
  • Returns quantized dispatch and combine-backward outputs as per-expert grouped tensors.
  • Adds caller-managed and symmetric-memory-backed data/scale buffer handling.
  • Adds dedicated Python and C++ distributed coverage for MXFP8 routing, eager allocation, and zero-copy operation.

Confidence Score: 5/5

The PR appears safe to merge.

No blocking failure remains.

Important Files Changed

Filename Overview
transformer_engine/pytorch/ep.py Adds recipe-controlled MXFP8 quantization, grouped dispatch outputs, and quantized combine-backward handling.
transformer_engine/pytorch/csrc/extensions/ep.cpp Extends native EP bindings to validate and forward MXFP8 scale tensors and communication windows.
transformer_engine/common/ep/ep_backend.cpp Routes block-scale descriptors alongside token data through forward and reverse NCCL dispatch.
transformer_engine/pytorch/distributed.py Adds explicit lifecycle management for the process-wide symmetric-memory pool.
tests/pytorch/distributed/run_ep.py Adds isolated MXFP8 dispatch and combine-backward coverage across fixed, eager, and zero-copy modes.
tests/cpp_distributed/test_ep.cu Adds native validation that MXFP8 data and scale rows follow the same expert routing.

Sequence Diagram

sequenceDiagram
  participant Caller
  participant PyTorch as PyTorch EP
  participant Quantizer as MXFP8 Quantizer
  participant Backend as Common EP Backend
  participant NCCL as NCCL EP
  Caller->>PyTorch: ep_dispatch(BF16 tokens)
  PyTorch->>Quantizer: Quantize data and block scales
  Quantizer-->>PyTorch: E4M3 data + E8M0 scales
  PyTorch->>Backend: Dispatch data and scale windows
  Backend->>NCCL: Route data and scales together
  NCCL-->>PyTorch: Per-expert routed buffers
  PyTorch-->>Caller: MXFP8 GroupedTensor
  Caller->>PyTorch: ep_combine backward(result gradient)
  PyTorch->>Quantizer: Quantize gradient
  PyTorch->>Backend: Reverse dispatch data and scales
  Backend->>NCCL: Scatter gradient to expert positions
  PyTorch-->>Caller: Per-expert MXFP8 gradient
Loading

Reviews (5): Last reviewed commit: "[PyTorch] Drop unused device param from ..." | Re-trigger Greptile

@phu0ngng
phu0ngng requested a review from zhongbozhu July 28, 2026 23:33
Comment thread transformer_engine/pytorch/ep.py Outdated
@phu0ngng

Copy link
Copy Markdown
Collaborator Author

/te-ci L1 pytorch

Comment thread .gitignore
Comment thread transformer_engine/pytorch/ep.py Outdated
Comment thread transformer_engine/pytorch/ep.py
Comment thread transformer_engine/pytorch/ep.py Outdated
Comment thread transformer_engine/pytorch/ep.py Outdated
Comment thread transformer_engine/pytorch/ep.py
Comment thread transformer_engine/pytorch/ep.py Outdated
return _SYMM_MEM_POOL


def release_symm_mem_pool(device: Optional[torch.device] = None) -> None:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is great, shall we also call it in ep_finalize?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

release_symm_mem_pool is only required before destroy_process_group().
User may call ep_finalize without destroy_process_group().

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Hmm why is that? I did not have it in mcore but it seems working (could be luck)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Also just see input device is never used, shall we remove?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

On why

release_symm_mem_pool is only required before destroy_process_group().

I think destroy_process_group() deletes the NCCL communicator that is used to create the symmem, so the cached symmem-s become limbo.

Are you calling destroy_process_group() in MCore? Why?

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Mcore shutdown will call destroy_process_group

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Good to know.
Anyway, I found out that release_symm_mem_pool is idempotent, so I added it to ep_finalize as you suggested.

Comment thread transformer_engine/pytorch/distributed.py
Comment thread transformer_engine/pytorch/ep.py
Comment thread transformer_engine/pytorch/ep.py Outdated
phu0ngng added 11 commits August 5, 2026 17:18
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
…der CUDA graph capture

Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
…CUDA-graph capture

Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
@phu0ngng

phu0ngng commented Aug 6, 2026

Copy link
Copy Markdown
Collaborator Author

/te-ci L1

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants