flashinfer.comm¶
This module provides communication primitives and utilities for distributed computing, including CUDA IPC, AllReduce operations, and memory management utilities.
CUDA IPC Utilities¶
|
|
|
Allocate a buffer and share it across the process group via CUDA IPC. |
|
Free a shared buffer previously created by |
DLPack Utilities¶
|
Pack a strided device allocation as a PyTorch tensor view. |
Mapping Utilities¶
|
A node with 8 GPUs, tp_size = 4, cp_size = 1, pp_size = 2 |
TensorRT-LLM AllReduce¶
Types and Enums¶
Core Operations¶
|
Parameters: - allreduce_in: the input tensor. [token_num, hidden_dim] - world_size: the size of the process group. - world_rank: the rank of the current process. - token_num: the number of tokens in the sequence. - hidden_dim: the dimension of the hidden states. - workspace_ptrs: the workspace pointers. - launch_with_pdl: whether to launch with pdl. - use_oneshot: whether to use oneshot. If None, internal heuristics will be used. - trigger_completion_at_end: whether to trigger completion at the end. - fp32_acc: whether to use fp32 accumulation. - pattern_code: the pattern code. - allreduce_out: the output tensor. [token_num, hidden_dim] - residual_in: the residual input tensor. [token_num, hidden_dim] - residual_out: the residual output tensor. [token_num, hidden_dim] - norm_out: the norm output tensor. [token_num, hidden_dim] - quant_out: the quant output tensor. [token_num, hidden_dim] - scale_out: the scale output tensor. Initialization referece: tests/comm/test_trtllm_allreduce_fusion.py - rms_gamma: the rms gamma tensor. [hidden_dim] - rms_eps: the rms epsilon value. - scale_factor: the scale factor. For cudaGraphs safety, it should be a tensor. - layout_code: the layout code. - metadata: optional workspace metadata dict from create_ipc_workspace_for_all_reduce_fusion. If provided, validates that token_num <= max_token_num, world_size == tp_size, and hidden_dim == workspace hidden_dim. Raises ValueError if validation fails. - block_quant_group_size: group size (in elements along hidden_dim) for per-token-group block-wise FP8 quantization patterns (e.g. |
|
Parameters: - inp: the input tensor. |
|
Parameters: - world_size: the size of the process group. - world_rank: the rank of the current process. - token_num: the number of tokens in the sequence. - hidden_dim: the dimension of the hidden states. - workspace_ptrs: the workspace pointers. - launch_with_pdl: whether to launch with pdl. - residual_in: the residual input tensor. [token_num, hidden_dim] - rms_gamma: the rms gamma tensor. [hidden_dim] - rms_eps: the rms epsilon value. - scale_factor: the scale factor. - moe_reduction_device_num_experts: the number of experts. - moe_reduction_scale_input: the scale input tensor. [token_num, hidden_dim] - moe_reduction_active_experts_token_input: the active experts token input tensor. [token_num, hidden_dim] - moe_reduction_token_input: the token input tensor. [token_num, hidden_dim] - layout_code: the layout code. - moe_allreduce_out: the moe allreduce output tensor. [token_num, hidden_dim] - residual_out: the residual output tensor. [token_num, hidden_dim] - norm_out: the norm output tensor. [token_num, hidden_dim] - quant_out: the quant output tensor. [token_num // 4, hidden_dim], fp16/bf16 -> fp4 - scale_out: the scale output tensor. Initialization referece: tests/comm/test_trtllm_moe_allreduce_fusion.py - weight_bias: bias added to rms_gamma before scaling. None or 0.0 -> standard RMSNorm (out = gamma * x * rsqrt(...)). 1.0 -> Gemma / Qwen3.5 RMSNorm (out = (1 + gamma) * x * rsqrt(...)). |
|
Parameters: - allreduce_in: the permuted/padded MoE expert output tensor. Shape [num_permuted_rows, hidden_dim]. Rows are referenced by expanded_idx_to_permuted_idx; num_permuted_rows may be larger than token_num * top_k due to expert padding. - residual_in: the residual input tensor. [token_num, hidden_dim] - norm_weight: the norm weight tensor. [hidden_dim] - expanded_idx_to_permuted_idx: the expanded index to permuted index tensor. [token_num, top_k] - norm_out: the norm output tensor. [token_num, hidden_dim] - residual_out: the residual output tensor. [token_num, hidden_dim] - quant_out: the quant output tensor. [token_num // 4, hidden_dim], fp16/bf16 -> fp4 - scale_out: the scale output tensor. [token_num // SF_VEC_SIZE, hidden_dim], fp16/bf16 -> fp4 - workspace_ptrs: the workspace pointers. - launch_with_pdl: whether to launch with pdl. - world_rank: the rank of the current process. - world_size: the size of the process group. - eps: the epsilon value. - shared_expert_output: the shared expert output tensor. [token_num, hidden_dim] - expert_scale_factor: the expert scale factor tensor. [token_num, top_k] - routed_scaling_factor: the routed scaling factor. - weight_bias: bias added to rms_gamma before scaling. None or 0.0 -> standard RMSNorm (out = gamma * x * rsqrt(...)). 1.0 -> Gemma / Qwen3.5 RMSNorm (out = (1 + gamma) * x * rsqrt(...)). |
Workspace Management¶
Parameters: - rank: the rank of the current process. |
|
Parameters: - tp_rank: the rank of the current process. |
|
Destroy a workspace created by trtllm_create_ipc_workspace_for_all_reduce. |
|
Destroy a workspace created by trtllm_create_ipc_workspace_for_all_reduce_fusion. |
Initialization and Utilities¶
|
Initialize a single Lamport-style buffer to negative zero. |
|
Initialize three Lamport buffers to negative zero. |
Compute the padded size (rows times columns) of the FP4 swizzled layout. |
Unified AllReduce Fusion API¶
|
AllReduce + RMSNorm fusion operation, with optional FP8/NVFP4 quantization for supported backends. |
|
Create workspace for AllReduce fusion operations. |
|
Base class for AllReduce fusion workspaces. |
All-reduce workspaces backed by SymmDeviceMemory preserve their CUDA
virtual addresses across process checkpoint/restore. After quiescing all
work, release the physical handles and restore them with a fresh communication
backend before replaying a captured CUDA graph:
workspace.checkpoint_prepare()
workspace.checkpoint_restore(comm_backend)
Both methods are collective. Every rank must call them in the same order, and
comm_backend must reproduce the original rank and world size. Repeated
calls are no-ops after the workspace reaches the requested state. If an
exception occurs after detach or reattach begins, do not retry or reuse the
workspace; restart the affected rank. Workspaces backed by torch symmetric
memory do not support this lifecycle.
- class flashinfer.comm.TRTLLMAllReduceFusionWorkspace(tp_size: int, tp_rank: int, max_token_num: int, hidden_dim: int, dtype: dtype = torch.float16, comm_backend: CommBackend | None = None, group: ProcessGroup | None = None)¶
Bases:
AllReduceFusionWorkspaceTensorRT-LLM workspace for AllReduce fusion.
- __init__(tp_size: int, tp_rank: int, max_token_num: int, hidden_dim: int, dtype: dtype = torch.float16, comm_backend: CommBackend | None = None, group: ProcessGroup | None = None)¶
Create TensorRT-LLM AllReduce fusion workspace.
- Parameters:
tp_size – Tensor parallel size (world size)
tp_rank – Tensor parallel rank
max_token_num – Maximum number of tokens
hidden_dim – Hidden dimension size
dtype – Data type
comm_backend – Communication backend
group – Process group for symmetric memory rendezvous. Defaults to torch.distributed.group.WORLD.
- property backend: str¶
Return backend name.
- checkpoint_prepare() None¶
Detach physical backing; repeated successful calls are no-ops.
- checkpoint_restore(comm_backend: CommBackend) None¶
Restore physical backing; repeated successful calls are no-ops.
- Parameters:
comm_backend (CommBackend) – Communication backend used to recreate and exchange workspace memory handles. It must have the same rank and world size as the original allocation.
- destroy() None¶
Destroy workspace and free resources.
- class flashinfer.comm.MNNVLAllReduceFusionWorkspace(mapping: Mapping, max_num_tokens: int | None = None, hidden_dim: int | None = None, dtype: dtype | None = None, buffer_size_in_bytes: int | None = None, comm_backend: CommBackend | None = None)¶
Bases:
AllReduceFusionWorkspace- __init__(mapping: Mapping, max_num_tokens: int | None = None, hidden_dim: int | None = None, dtype: dtype | None = None, buffer_size_in_bytes: int | None = None, comm_backend: CommBackend | None = None)¶
Initialize the MNNVL Allreduce Fusion Workspace. The workspace will be allocated and initialized based on the provided problem size. If max_num_tokens is larger than the one-shot threshold, the workspace will be created according to the max of required one-shot size at threshold, or the required two-shot size. Note that the workspace is not bind to the given problem size. It can be reused for different problem size without reinitialization given the allocated size is sufficient.
If the buffer_size_in_bytes is provided, the workspace will be created according to the provided size. The user is expected to use the utility function get_required_buffer_size_bytes to calculate the required size. The actual allocation size may be larger due to alignment requirements. This covers the advanced used case, for example, the user may want to enforce oneshot strategy and ignore the heuristics.
Either max_num_tokens or buffer_size_in_bytes must be provided.
comm_backend will be used for creating the workspace and synchronization. If not provided, MPIBackend will be used which will use COMM_WORLD for synchronization.
- Parameters:
mapping – Mapping configuration containing rank info
max_num_tokens – The maximum number of tokens in the input tensor.
hidden_dim – The hidden dimension of the tensors to be reduced.
dtype – The data type of the tensors to be reduced.
buffer_size_in_bytes – The requested size in bytes for each lamport buffer. The actual allocation size may be larger due to alignment requirements. The actual usable size will be NUM_LAMPORT_BUFFERS * actual_buffer_size_per_lamport_buffer.
- property backend: str¶
Return backend name.
- checkpoint_prepare() None¶
Detach physical backing; repeated successful calls are no-ops.
- checkpoint_restore(comm_backend: CommBackend) None¶
Restore physical backing; repeated successful calls are no-ops.
- Parameters:
comm_backend (CommBackend) – Communication backend used to recreate and exchange MNNVL memory handles. It must have the same rank and world size as the original allocation.
- destroy() None¶
Destroy workspace and free resources.
- static get_required_buffer_size_bytes(tp_size: int, num_tokens: int, hidden_dim: int, dtype: dtype, strategy: MNNVLAllreduceFusionStrategy = MNNVLAllreduceFusionStrategy.AUTO) int¶
Calculate the required buffer size for a given problem size.
- is_buffer_size_sufficient(tp_size: int, num_tokens: int, hidden_dim: int, dtype: dtype, strategy: MNNVLAllreduceFusionStrategy = MNNVLAllreduceFusionStrategy.AUTO) bool¶
Calculate the required buffer size for a given problem size.
vLLM AllReduce¶
|
Perform an out-of-place all-reduce via the vLLM custom kernel. |
|
Release the resources held by a vLLM custom all-reduce handle. |
|
Initialize the vLLM custom all-reduce backend. |
|
Register a peer's IPC-shared buffer with the local all-reduce handle. |
|
Register graph-capture buffers across the all-reduce world. |
Return IPC metadata for graph-capture buffers. |
|
Return the size of the vLLM all-reduce metadata structure in bytes. |
MNNVL (Multi-Node NVLink)¶
Core Classes¶
|
|
|
Wrapper class for SymmDeviceMemory to facilitate PyTorch tensor creation. |
TensorRT-LLM MNNVL AllReduce¶
|
Deprecated pointer-based MNNVL all-reduce API. |
|
Perform an MNNVL all-reduce sum across tensor-parallel ranks. |
Performs MNNVL Allreduce + Residual + RMSNorm. |
|
Perform MNNVL AllReduce + Residual + RMSNorm + FP8/NVFP4 quantization. |
|
Performs MNNVL TwoShot Allreduce + RMSNorm. |
|
MNNVL A2A (Throughput Backend)¶
|
Initialize the MoE all-to-all workspace and return a metainfo tensor. |
|
Dispatch tokens and payloads to their target expert ranks. |
|
Combine per-expert outputs back to the originating ranks. |
|
Sanitize invalid slots that contain no token routed to this rank by setting their expert IDs to |
|
Compute the per-rank workspace size for the MoE all-to-all primitive. |
Wrap a slice of the shared workspace as a typed tensor view. |
- class flashinfer.comm.MoeAlltoAll(mapping: Mapping, max_num_tokens: int, top_k: int, num_experts: int, workspace_size_per_rank: int = None, hidden_size: int = None, mnnvl_config: MnnvlConfig | None = None, eplb_stats_num_experts: int = 0, enable_rank_mask: bool = False)¶
Bases:
objectManages MoE All-to-All operations with proper workspace allocation and synchronization.
This class provides the throughput-optimized backend that supports multiple payloads per collective operation, explicit dispatch/combine phases, and workspace-backed tensors.
Example
>>> moe_a2a = MoeAlltoAll(mapping, max_num_tokens=2048, top_k=2, num_experts=8) >>> recv = moe_a2a.dispatch(experts, [hidden, ids, scales], batch_size) >>> output = moe_a2a.combine(processed, batch_size)
- __init__(mapping: Mapping, max_num_tokens: int, top_k: int, num_experts: int, workspace_size_per_rank: int = None, hidden_size: int = None, mnnvl_config: MnnvlConfig | None = None, eplb_stats_num_experts: int = 0, enable_rank_mask: bool = False)¶
Initialize
MoeAlltoAlland allocate the shared workspace.- Parameters:
mapping (Mapping) – Mapping object describing the parallel layout (must expose
moe_ep_rankandmoe_ep_size).max_num_tokens (int) – Maximum number of tokens this rank will dispatch in any single call.
top_k (int) – Number of experts assigned per token.
num_experts (int) – Total number of experts (across all ranks).
workspace_size_per_rank (int, optional) – Pre-computed workspace size in bytes per rank. When
None,hidden_sizemust be provided and the workspace is sized viaget_moe_workspace_size_per_rank().hidden_size (int, optional) – Hidden dimension size, used to derive
workspace_size_per_rankwhen the latter is omitted.mnnvl_config (MnnvlConfig, optional) – Optional configuration for the underlying MNNVL communication backend.
eplb_stats_num_experts (int) – Number of experts to reserve for the EPLB gathered-stats region.
0(default) disables EPLB; when non-zero, pass aneplb_local_statstensor of this length todispatch().enable_rank_mask (bool) – Whether
dispatch()/combine()may be called with anactive_rank_mask. Fixed for the lifetime of this instance (mirrors the underlying kernel’s compile-time specialization):False(default) compiles out every rank-mask check and forbids passingactive_rank_mask.
- checkpoint_prepare() None¶
Unmap MNNVL handles for checkpointing; repeated calls are no-ops.
- checkpoint_restore(comm_backend: CommBackend) None¶
Remap MNNVL handles after restore; repeated calls are no-ops.
- Parameters:
comm_backend (CommBackend) – Communication backend used to recreate and exchange MNNVL memory handles. It must have the same rank and world size as the original allocation.
- combine(payload: Tensor, runtime_max_tokens_per_rank: int, payload_in_workspace: bool = False, output_dtype: dtype | None = None, output_scales: Tensor | None = None, output_scalar_scale: float = 1.0, sf_layout: SfLayout = SfLayout.layout_linear, output: Tensor | None = None, *, use_low_precision: bool = False, active_rank_mask: Tensor | None = None) Tensor¶
Run the MoE all-to-all combine phase.
- Parameters:
payload (torch.Tensor) –
[ep_size, runtime_max_tokens_per_rank, elements_per_token]output payload to scatter back to source ranks.runtime_max_tokens_per_rank (int) – Maximum tokens per rank in this batch (same value passed to
dispatch()).payload_in_workspace (bool) –
Trueifpayloadis already a workspace-backed view (skips the staging copy). Defaults toFalse.output_dtype (Optional[torch.dtype]) – Optional output data type. Currently supports
torch.bfloat16,torch.float8_e4m3fn, andtorch.uint8(packed fp4).output_scales (Optional[torch.Tensor]) – Optional output scale tensor for quantized outputs. Currently supports UE8M0 (packed in
torch.uint8) with vector size 32.output_scalar_scale (float) – Per-tensor global scale applied before FP4 block scaling (NVFP4 SFScaleVal). Defaults to
1.0; ignored by MXFP8/MXFP4 paths.sf_layout (SfLayout) – Output swizzle layout. Defaults to
SfLayout.layout_linear.output (Optional[torch.Tensor]) – Caller-provided contiguous output tensor. Its shape and dtype must match the requested combine output, and it must be on the same device as
payload.use_low_precision (bool) – If
True, quantize the recv-buffer payload to FP8 (e4m3) before accumulating; the combine upcasts to a bf16 output.active_rank_mask (torch.Tensor, optional) – CPU
uint64tensor of shape[MOE_A2A_RANK_MASK_WORDS]. Should match the mask passed to the precedingdispatch()call (or be omitted from both). Requires the instance to have been constructed withenable_rank_mask=True.
- Returns:
[local_num_tokens, elements_per_token]combined tensor.- Return type:
torch.Tensor
- dispatch(token_selected_experts: Tensor, input_payloads: list[Tensor], runtime_max_tokens_per_rank: int, invalid_token_expert_id: int | None = None, expert_id_payload_index: int | None = None, eplb_local_stats: Tensor | None = None, active_rank_mask: Tensor | None = None) list[Tensor]¶
Run the MoE all-to-all dispatch phase.
- Parameters:
token_selected_experts (torch.Tensor) –
[local_num_tokens, top_k]int32tensor of expert assignments.input_payloads (list[torch.Tensor]) – Per-token payload tensors, each shaped
[local_num_tokens, *].runtime_max_tokens_per_rank (int) – Maximum tokens per rank in this batch. Must be
<=max_num_tokensused at construction.invalid_token_expert_id (int, optional) – If supplied, expert IDs not owned by the current rank are rewritten to this value. Requires
expert_id_payload_index.expert_id_payload_index (int, optional) – Index into
input_payloadsthat holds the expert IDs to sanitize. Required wheninvalid_token_expert_idis set.eplb_local_stats (torch.Tensor, optional) –
[eplb_stats_num_experts]int32tensor of this rank’s local EPLB statistics to all-gather during dispatch. Requires the instance to have been constructed witheplb_stats_num_expertsset. The gathered result is available afterwards viaeplb_gathered_stats.active_rank_mask (torch.Tensor, optional) – CPU
uint64tensor of shape[MOE_A2A_RANK_MASK_WORDS](seemoe_a2a_active_rank_mask()). Tokens routed to a masked-off rank are dropped instead of hanging the collective. Requires the instance to have been constructed withenable_rank_mask=True.
- Returns:
Workspace-backed receive tensors, one per
input_payloadsentry, each shaped[ep_size, runtime_max_tokens_per_rank, *].- Return type:
list[torch.Tensor]
- property eplb_gathered_stats: Tensor | None¶
Gathered EPLB stats from the most recent
dispatch().Workspace-backed
[ep_size, eplb_stats_num_experts]int32view (rowrholds rankr’seplb_local_stats), orNonewhen EPLB is disabled or noeplb_local_statswas passed. Valid only betweendispatch()andcombine()(combine resets state).
- get_combine_payload_tensor_in_workspace(runtime_max_tokens_per_rank: int, hidden_size: int, dtype: dtype) Tensor¶
Return a workspace-backed view to use as the combine payload.
Zero-copy variant of
combine(): experts can write directly into the returned tensor and callcombine()withpayload_in_workspace=True. Must be called after a successfuldispatch()and beforecombine().- Parameters:
runtime_max_tokens_per_rank (int) – Maximum tokens per rank in this batch.
hidden_size (int) – Hidden dimension size.
dtype (torch.dtype) – Element dtype of the resulting view.
- Returns:
[ep_size, runtime_max_tokens_per_rank, hidden_size]workspace-backed tensor.- Return type:
torch.Tensor
- Raises:
RuntimeError – If called before a successful
dispatch().
- static get_moe_workspace_size_per_rank(ep_size: int, top_k: int, max_num_tokens: int, hidden_size: int, extra_payload_bytes_per_token: int = 0, eplb_stats_num_experts: int = 0) int¶
Compute the per-rank workspace size for the MoE all-to-all primitive.
Convenience wrapper around
moe_a2a_get_workspace_size_per_rank()that derives the dispatch / combine payload sizes fromhidden_sizeandtop_kassuming 16-bit hidden states. For a tighter bound on quantized models usemoe_a2a_get_workspace_size_per_rank()directly.- Parameters:
ep_size (int) – Total expert-parallel world size.
top_k (int) – Number of experts assigned per token.
max_num_tokens (int) – Maximum number of tokens across all ranks.
hidden_size (int) – Hidden dimension size.
extra_payload_bytes_per_token (int) – Extra payload bytes per token to reserve (e.g. for quantization scales). Defaults to
0.eplb_stats_num_experts (int) – Number of experts reserved for the EPLB gathered-stats region (
0disables it).
- Returns:
Required workspace size per rank, in bytes.
- Return type:
int
MoeAlltoAll preserves its CUDA virtual addresses across process
checkpoint/restore. After quiescing all work, call checkpoint_prepare to
release the non-checkpointable physical MNNVL handles. Then call
checkpoint_restore with a fresh communication backend before replaying a
captured CUDA graph:
moe_alltoall.checkpoint_prepare()
moe_alltoall.checkpoint_restore(comm_backend)
Both methods are collective. Every rank must call them in the same order, and
comm_backend must reproduce the original rank and world size.
Repeated calls are no-ops after the workspace reaches the requested state.
If an exception occurs after physical handle unmapping or remapping begins,
do not retry or reuse the workspace; restart the affected rank.
Unmap MNNVL handles for checkpointing; repeated calls are no-ops. |
|
|
Remap MNNVL handles after restore; repeated calls are no-ops. |
DCP All-to-All (Context-Parallel Attention Reduction)¶
|
Return the workspace size (in bytes) per rank for the given CP group size. |
Allocate an MNNVL-backed workspace of shape |
|
|
Initialize the workspace FIFO buffers (call once before the first alltoall). |
|
Perform the DCP all-to-all exchange. |
Mixed Communication¶
|
Enumeration of mixed communication operation types. |
|
Enumeration of mixed communication execution modes. |
|
An implementation for the combinations of all-reduce + all-gather and reduce-scatter + all-reduce. |
|
Execute a mixed communication operation. |