custom all gather kernel

Signed-off-by: Woosuk Kwon <woosuk@inferact.ai>
This commit is contained in:
Woosuk Kwon
2026-04-14 19:08:10 +00:00
parent 4fcf47661a
commit d00bdaee51
14 changed files with 897 additions and 30 deletions
+47
View File
@@ -104,6 +104,53 @@ void all_reduce(fptr_t _fa, torch::Tensor& inp, torch::Tensor& out,
}
}
void reduce_scatter(fptr_t _fa, torch::Tensor& inp, torch::Tensor& out,
fptr_t _reg_buffer, int64_t reg_buffer_sz_bytes) {
auto fa = reinterpret_cast<vllm::CustomAllreduce*>(_fa);
const at::cuda::OptionalCUDAGuard device_guard(device_of(inp));
auto stream = c10::cuda::getCurrentCUDAStream().stream();
TORCH_CHECK_EQ(inp.scalar_type(), out.scalar_type());
TORCH_CHECK(inp.numel() == out.numel() * fa->world_size_);
TORCH_CHECK(_is_weak_contiguous(inp));
TORCH_CHECK(_is_weak_contiguous(out));
auto input_size = inp.numel() * inp.element_size();
auto reg_buffer = reinterpret_cast<void*>(_reg_buffer);
if (reg_buffer) {
TORCH_CHECK_LE(input_size, reg_buffer_sz_bytes);
AT_CUDA_CHECK(cudaMemcpyAsync(reg_buffer, inp.data_ptr(), input_size,
cudaMemcpyDeviceToDevice, stream));
} else {
reg_buffer = inp.data_ptr();
}
switch (out.scalar_type()) {
case at::ScalarType::Float: {
fa->reduce_scatter<float>(stream, reinterpret_cast<float*>(reg_buffer),
reinterpret_cast<float*>(out.data_ptr()),
inp.numel());
break;
}
case at::ScalarType::Half: {
fa->reduce_scatter<half>(stream, reinterpret_cast<half*>(reg_buffer),
reinterpret_cast<half*>(out.data_ptr()),
inp.numel());
break;
}
#if (__CUDA_ARCH__ >= 800 || !defined(__CUDA_ARCH__))
case at::ScalarType::BFloat16: {
fa->reduce_scatter<nv_bfloat16>(
stream, reinterpret_cast<nv_bfloat16*>(reg_buffer),
reinterpret_cast<nv_bfloat16*>(out.data_ptr()), inp.numel());
break;
}
#endif
default:
throw std::runtime_error(
"custom reduce_scatter only supports float32, float16 and "
"bfloat16");
}
}
void dispose(fptr_t _fa) {
delete reinterpret_cast<vllm::CustomAllreduce*>(_fa);
}
+82 -1
View File
@@ -313,6 +313,24 @@ __global__ void __launch_bounds__(512, 1)
barrier_at_end<ngpus, true>(sg, self_sg, rank);
}
template <typename T, int ngpus>
__global__ void __launch_bounds__(512, 1)
cross_device_reduce_scatter(RankData* _dp, RankSignals sg, Signal* self_sg,
T* __restrict__ result, int rank,
int chunk_size) {
using P = typename packed_t<T>::P;
using A = typename packed_t<T>::A;
auto dp = *_dp;
int offset = rank * chunk_size;
barrier_at_start<ngpus>(sg, self_sg, rank);
for (int idx = blockIdx.x * blockDim.x + threadIdx.x; idx < chunk_size;
idx += gridDim.x * blockDim.x) {
((P*)result)[idx] =
packed_reduce<P, ngpus, A>((const P**)&dp.ptrs[0], offset + idx);
}
barrier_at_end<ngpus, true>(sg, self_sg, rank);
}
template <typename P>
DINLINE P* get_tmp_buf(Signal* sg) {
return (P*)(((Signal*)sg) + 1);
@@ -616,6 +634,69 @@ class CustomAllreduce {
#undef KL
}
template <typename T>
void reduce_scatter(cudaStream_t stream, T* input, T* output, int size,
int threads = 512, int block_limit = defaultBlockLimit) {
auto d = packed_t<T>::P::size;
if (size % d != 0)
throw std::runtime_error(
"custom reduce_scatter currently requires input length to be "
"multiple of " +
std::to_string(d));
if (size % world_size_ != 0)
throw std::runtime_error(
"custom reduce_scatter requires input length to be divisible by "
"world_size");
if (block_limit > kMaxBlocks)
throw std::runtime_error("max supported block limit is " +
std::to_string(kMaxBlocks) + ". Got " +
std::to_string(block_limit));
RankData* ptrs;
cudaStreamCaptureStatus status;
CUDACHECK(cudaStreamIsCapturing(stream, &status));
if (status == cudaStreamCaptureStatusActive) {
ptrs = d_rank_data_base_ + graph_unreg_buffers_.size();
graph_unreg_buffers_.push_back(input);
} else {
auto it = buffers_.find(input);
if (it == buffers_.end())
throw std::runtime_error(
"buffer address " +
std::to_string(reinterpret_cast<uint64_t>(input)) +
" is not registered!");
ptrs = it->second;
}
int chunk_size = size / world_size_ / d;
int blocks = std::min(block_limit, (chunk_size + threads - 1) / threads);
#define KL(ngpus) \
cross_device_reduce_scatter<T, ngpus><<<blocks, threads, 0, stream>>>( \
ptrs, sg_, self_sg_, output, rank_, chunk_size);
switch (world_size_) {
case 2:
KL(2);
break;
case 4:
KL(4);
break;
case 6:
KL(6);
break;
case 8:
KL(8);
break;
default:
throw std::runtime_error(
"custom reduce_scatter only supports num gpus in (2,4,6,8). "
"Actual num gpus = " +
std::to_string(world_size_));
}
#undef KL
}
~CustomAllreduce() {
for (auto [_, ptr] : ipc_handles_) {
CUDACHECK(cudaIpcCloseMemHandle(ptr));
@@ -629,4 +710,4 @@ class CustomAllreduce {
* template void vllm::CustomAllreduce::allreduce<half>(cudaStream_t, half *,
half *, int, int, int);
*/
} // namespace vllm
} // namespace vllm
+2
View File
@@ -280,6 +280,8 @@ fptr_t init_custom_ar(const std::vector<int64_t>& fake_ipc_ptrs,
bool fully_connected);
void all_reduce(fptr_t _fa, torch::Tensor& inp, torch::Tensor& out,
fptr_t reg_buffer, int64_t reg_buffer_sz_bytes);
void reduce_scatter(fptr_t _fa, torch::Tensor& inp, torch::Tensor& out,
fptr_t reg_buffer, int64_t reg_buffer_sz_bytes);
void dispose(fptr_t _fa);
int64_t meta_size();
void register_buffer(fptr_t _fa, const std::vector<int64_t>& fake_ipc_ptrs);
+318
View File
@@ -0,0 +1,318 @@
// Standalone fused MoE all-gather kernel for EP dispatch.
//
// JIT-compilable via torch.utils.cpp_extension — no vLLM build required.
//
// Gathers the MoE dispatch tensors from all EP ranks in a single kernel:
// - topk_ids [N, topk] int32
// - topk_weights [N, topk] float32 / bfloat16
// - hidden_states [N, D_h] uint8 (NVFP4) / bfloat16
// - quant_scales [N, D_s] (optional) any dtype
//
// Each tensor is packed into a pre-registered IPC buffer at 16-byte-aligned
// offsets. The kernel:
// 1. Copies inputs into the IPC buffer (local SM write).
// 2. Barrier — all ranks' writes become visible via NVLink.
// 3. Gathers from all peers' buffers into separate output tensors.
// 4. Barrier — done.
//
// Under CUDA graphs, the input tensor addresses are fixed and the buffer
// copies are captured. The total overhead is dominated by the two barrier
// round-trips (~5µs each on NVLink).
//
// All data movement uses 128-bit (int4) loads/stores.
#include <cuda.h>
#include <cuda_runtime.h>
#include <torch/extension.h>
#include <c10/cuda/CUDAGuard.h>
#include <c10/cuda/CUDAStream.h>
// ---------------------------------------------------------------------------
// Flag-based barrier (from custom_all_reduce.cuh, standalone)
// ---------------------------------------------------------------------------
constexpr int kMaxBlocks = 36;
using FlagType = uint32_t;
struct Signal {
alignas(128) FlagType start[kMaxBlocks][8];
alignas(128) FlagType end[kMaxBlocks][8];
alignas(128) FlagType _flag[kMaxBlocks];
};
struct __align__(16) RankData {
const void* ptrs[8];
};
struct __align__(16) RankSignals {
Signal* signals[8];
};
#define DINLINE __device__ __forceinline__
static DINLINE void st_flag_volatile(FlagType* addr, FlagType val) {
asm volatile("st.volatile.global.u32 [%1], %0;" ::"r"(val), "l"(addr));
}
static DINLINE FlagType ld_flag_volatile(FlagType* addr) {
FlagType v;
asm volatile("ld.volatile.global.u32 %0, [%1];" : "=r"(v) : "l"(addr));
return v;
}
template <int ngpus>
DINLINE void barrier_at_start(const RankSignals& sg, Signal* self_sg,
int rank) {
FlagType flag = self_sg->_flag[blockIdx.x] + 1;
if (threadIdx.x < ngpus) {
st_flag_volatile(&sg.signals[threadIdx.x]->start[blockIdx.x][rank], flag);
while (ld_flag_volatile(&self_sg->start[blockIdx.x][threadIdx.x]) != flag);
}
__syncthreads();
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
}
template <int ngpus>
DINLINE void barrier_at_end(const RankSignals& sg, Signal* self_sg, int rank) {
__syncthreads();
FlagType flag = self_sg->_flag[blockIdx.x] + 1;
if (threadIdx.x < ngpus) {
st_flag_volatile(&sg.signals[threadIdx.x]->end[blockIdx.x][rank], flag);
while (ld_flag_volatile(&self_sg->end[blockIdx.x][threadIdx.x]) != flag);
}
if (threadIdx.x == 0) self_sg->_flag[blockIdx.x] = flag;
}
// ---------------------------------------------------------------------------
// Fused MoE dispatch all-gather kernel
// ---------------------------------------------------------------------------
// The kernel has 4 phases:
// 1. Pack local inputs into IPC buffer.
// 2. Barrier (all ranks' data visible).
// 3. Gather: one tight contiguous read from each peer → flat staging buffer.
// 4. Barrier (gather done, safe to overwrite IPC buffer next iteration).
// 5. Scatter: redistribute the flat staging into per-tensor outputs (local
// L2).
//
// Phase 3 is NVLink-critical: one loop, no conditionals, all peers pipelined.
// Phase 5 is local-memory only (L2 speed), runs AFTER the end barrier.
// has_scales: compile-time flag for the optional 4th tensor (quant_scales).
//
// Future optimization (TODO): register hidden_states via IPC (like custom
// allreduce does in graph mode) to skip its Phase 1 copy. This would save
// ~3µs by eliminating the hidden_states copy + reducing the scatter.
// Requires CUDA graph integration to register the hidden_states tensor
// address during capture.
template <int ngpus, int nbufs>
__global__ void __launch_bounds__(512, 1)
moe_allgather_kernel(RankData* _dp, RankSignals sg, Signal* self_sg,
int rank,
// up to 4 inputs
const void* inp0, const void* inp1, const void* inp2,
const void* inp3, int off0, int sz0, int off1, int sz1,
int off2, int sz2, int off3, int sz3, void* out0,
void* out1, void* out2, void* out3, void* staging,
int total_sz) {
using V = int4;
int tid = blockIdx.x * blockDim.x + threadIdx.x;
int stride = gridDim.x * blockDim.x;
auto dp = *_dp;
char* my_buf = (char*)dp.ptrs[rank];
// Phase 1: pack local inputs into IPC buffer.
if constexpr (nbufs > 0) {
const V* s = (const V*)inp0;
V* d = (V*)(my_buf + off0);
for (int i = tid; i < sz0; i += stride) d[i] = s[i];
}
if constexpr (nbufs > 1) {
const V* s = (const V*)inp1;
V* d = (V*)(my_buf + off1);
for (int i = tid; i < sz1; i += stride) d[i] = s[i];
}
if constexpr (nbufs > 2) {
const V* s = (const V*)inp2;
V* d = (V*)(my_buf + off2);
for (int i = tid; i < sz2; i += stride) d[i] = s[i];
}
if constexpr (nbufs > 3) {
const V* s = (const V*)inp3;
V* d = (V*)(my_buf + off3);
for (int i = tid; i < sz3; i += stride) d[i] = s[i];
}
__threadfence_system();
// Phase 2: barrier.
barrier_at_start<ngpus>(sg, self_sg, rank);
// Phase 3: single contiguous gather into staging buffer.
{
const V* peers[ngpus];
#pragma unroll
for (int s = 0; s < ngpus; s++) peers[s] = (const V*)dp.ptrs[s];
for (int i = tid; i < total_sz; i += stride) {
#pragma unroll
for (int s = 0; s < ngpus; s++)
((V*)staging)[s * total_sz + i] = peers[s][i];
}
}
// Phase 4: end barrier.
barrier_at_end<ngpus>(sg, self_sg, rank);
// Phase 5: scatter from staging to per-tensor outputs (local L2).
{
const V* stg = (const V*)staging;
if constexpr (nbufs > 0) {
const int b0 = off0 / (int)sizeof(V);
for (int i = tid; i < sz0; i += stride) {
#pragma unroll
for (int s = 0; s < ngpus; s++)
((V*)out0)[s * sz0 + i] = stg[s * total_sz + b0 + i];
}
}
if constexpr (nbufs > 1) {
const int b1 = off1 / (int)sizeof(V);
for (int i = tid; i < sz1; i += stride) {
#pragma unroll
for (int s = 0; s < ngpus; s++)
((V*)out1)[s * sz1 + i] = stg[s * total_sz + b1 + i];
}
}
if constexpr (nbufs > 2) {
const int b2 = off2 / (int)sizeof(V);
for (int i = tid; i < sz2; i += stride) {
#pragma unroll
for (int s = 0; s < ngpus; s++)
((V*)out2)[s * sz2 + i] = stg[s * total_sz + b2 + i];
}
}
if constexpr (nbufs > 3) {
const int b3 = off3 / (int)sizeof(V);
for (int i = tid; i < sz3; i += stride) {
#pragma unroll
for (int s = 0; s < ngpus; s++)
((V*)out3)[s * sz3 + i] = stg[s * total_sz + b3 + i];
}
}
}
}
// ---------------------------------------------------------------------------
// Host launcher
// ---------------------------------------------------------------------------
// Compute 16-byte-aligned offset and int4-unit size for one tensor.
struct TensorDesc {
void* inp;
int off; // byte offset in IPC buffer
int sz; // size in int4 (16-byte) units
int64_t nbytes;
};
static TensorDesc make_desc(torch::Tensor& inp, int64_t& cursor) {
TORCH_CHECK(inp.is_contiguous(), "input must be contiguous");
int64_t nbytes = inp.numel() * inp.element_size();
TORCH_CHECK(nbytes % 16 == 0, "tensor byte size must be multiple of 16, got ",
nbytes);
cursor = (cursor + 15) & ~15;
TensorDesc d;
d.inp = inp.data_ptr();
d.off = static_cast<int>(cursor);
d.sz = static_cast<int>(nbytes / 16);
d.nbytes = nbytes;
cursor += nbytes;
return d;
}
void moe_all_gather(int64_t rank_data_ptr, int64_t signals_ptr,
int64_t self_signal_ptr, int64_t rank, int64_t world_size,
std::vector<torch::Tensor>& inputs,
std::vector<torch::Tensor>& outputs) {
auto stream = c10::cuda::getCurrentCUDAStream().stream();
int n = static_cast<int>(inputs.size());
TORCH_CHECK(n >= 2 && n <= 4, "2-4 input tensors required");
TORCH_CHECK(inputs.size() == outputs.size());
int64_t cursor = 0;
TensorDesc descs[4] = {};
for (int i = 0; i < n; i++) descs[i] = make_desc(inputs[i], cursor);
TORCH_CHECK(cursor % 16 == 0);
int total_sz = static_cast<int>(cursor / 16);
for (int i = 0; i < n; i++) {
TORCH_CHECK(outputs[i].is_contiguous());
TORCH_CHECK(outputs[i].numel() == inputs[i].numel() * world_size);
}
auto* ptrs = reinterpret_cast<RankData*>(rank_data_ptr);
RankSignals sg = *reinterpret_cast<RankSignals*>(signals_ptr);
auto* self_sg = reinterpret_cast<Signal*>(self_signal_ptr);
int r = static_cast<int>(rank);
int threads = 512;
int blocks =
std::max(1, std::min(kMaxBlocks, (total_sz + threads - 1) / threads));
void *inps[4] = {}, *outs[4] = {};
int offs[4] = {}, szs[4] = {};
for (int i = 0; i < n; i++) {
inps[i] = descs[i].inp;
offs[i] = descs[i].off;
szs[i] = descs[i].sz;
outs[i] = outputs[i].data_ptr();
}
auto staging = torch::empty(
{(int64_t)world_size * total_sz * (int64_t)sizeof(int4)},
torch::TensorOptions().dtype(torch::kUInt8).device(inputs[0].device()));
#define KL(ngpus, nb) \
moe_allgather_kernel<ngpus, nb><<<blocks, threads, 0, stream>>>( \
ptrs, sg, self_sg, r, inps[0], inps[1], inps[2], inps[3], offs[0], \
szs[0], offs[1], szs[1], offs[2], szs[2], offs[3], szs[3], outs[0], \
outs[1], outs[2], outs[3], staging.data_ptr(), total_sz);
#define GPU_CASE(ngpus) \
case ngpus: \
switch (n) { \
case 2: \
KL(ngpus, 2); \
break; \
case 3: \
KL(ngpus, 3); \
break; \
case 4: \
KL(ngpus, 4); \
break; \
} \
break;
switch (world_size) {
GPU_CASE(2)
GPU_CASE(4)
GPU_CASE(6)
GPU_CASE(8)
default:
TORCH_CHECK(false, "world_size must be 2, 4, 6, or 8");
}
#undef GPU_CASE
#undef KL
}
// ---------------------------------------------------------------------------
// Python binding
// ---------------------------------------------------------------------------
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("moe_all_gather", &moe_all_gather,
"Fused MoE dispatch all-gather with in-kernel scatter");
}
@@ -0,0 +1,98 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""
Standalone fused MoE all-gather for EP dispatch.
JIT-compiles the CUDA kernel on first use (cached afterwards).
The kernel gathers contiguously from all peers (NVLink-optimal),
then scatters to per-tensor outputs inside the kernel (local L2).
Zero Python-side post-processing.
Usage:
ag = MoeAllGather(custom_allreduce)
ids_g, wt_g, hs_g, sc_g = ag.gather(topk_ids, topk_weights, hidden, scales)
"""
from __future__ import annotations
import os
from pathlib import Path
import torch
_lib = None
def _load_lib():
global _lib
if _lib is not None:
return _lib
from torch.utils.cpp_extension import load
src = str(Path(__file__).with_name("moe_allgather.cu"))
_lib = load(
name="moe_allgather_kernel",
sources=[src],
extra_cuda_cflags=["-O3", "--use_fast_math"],
verbose=os.environ.get("MOE_AG_VERBOSE", "") == "1",
)
return _lib
class MoeAllGather:
"""Fused MoE dispatch all-gather backed by a one-shot flag-barrier kernel."""
def __init__(self, ca_comm):
self.rank = ca_comm.rank
self.world_size = ca_comm.world_size
self.device = ca_comm.device
self.meta_ptrs = ca_comm.meta_ptrs
self.buffer_ptrs = ca_comm.buffer_ptrs
self.max_size = ca_comm.max_size
self._rank_signals = torch.zeros(8, dtype=torch.int64)
for i in range(self.world_size):
self._rank_signals[i] = self.meta_ptrs[i]
self._rank_signals_ptr = self._rank_signals.data_ptr()
self._self_signal_ptr = self.meta_ptrs[self.rank]
self._rank_data = torch.zeros(
8, dtype=torch.int64, device=f"cuda:{self.device.index}"
)
for i in range(self.world_size):
self._rank_data[i] = self.buffer_ptrs[i]
self._rank_data_ptr = self._rank_data.data_ptr()
def gather(
self,
topk_ids: torch.Tensor,
topk_weights: torch.Tensor,
hidden_states: torch.Tensor,
scales: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]:
lib = _load_lib()
ws = self.world_size
inputs = [topk_ids, topk_weights, hidden_states]
if scales is not None:
inputs.append(scales)
outputs = [
torch.empty((t.shape[0] * ws, *t.shape[1:]), dtype=t.dtype, device=t.device)
for t in inputs
]
lib.moe_all_gather(
self._rank_data_ptr,
self._rank_signals_ptr,
self._self_signal_ptr,
self.rank,
self.world_size,
inputs,
outputs,
)
if scales is not None:
return outputs[0], outputs[1], outputs[2], outputs[3]
return outputs[0], outputs[1], outputs[2], None
+4
View File
@@ -655,6 +655,10 @@ TORCH_LIBRARY_EXPAND(CONCAT(TORCH_EXTENSION_NAME, _custom_ar), custom_ar) {
"all_reduce(int fa, Tensor inp, Tensor! out, int reg_buffer, "
"int reg_buffer_sz_bytes) -> ()");
custom_ar.impl("all_reduce", torch::kCUDA, &all_reduce);
custom_ar.def(
"reduce_scatter(int fa, Tensor inp, Tensor! out, int reg_buffer, "
"int reg_buffer_sz_bytes) -> ()");
custom_ar.impl("reduce_scatter", torch::kCUDA, &reduce_scatter);
custom_ar.def("dispose", &dispose);
custom_ar.def("meta_size", &meta_size);
+10
View File
@@ -2797,6 +2797,16 @@ def all_reduce(
torch.ops._C_custom_ar.all_reduce(fa, inp, out, reg_buffer, reg_buffer_sz_bytes)
def reduce_scatter(
fa: int,
inp: torch.Tensor,
out: torch.Tensor,
reg_buffer: int,
reg_buffer_sz_bytes: int,
) -> None:
torch.ops._C_custom_ar.reduce_scatter(fa, inp, out, reg_buffer, reg_buffer_sz_bytes)
def dispose(fa: int) -> None:
torch.ops._C_custom_ar.dispose(fa)
+290 -23
View File
@@ -16,9 +16,147 @@ from vllm.utils.flashinfer import (
has_flashinfer_nvlink_two_sided,
)
from vllm.utils.import_utils import has_deep_ep, has_mori
from vllm.utils.torch_utils import direct_register_custom_op
from .base_device_communicator import All2AllManagerBase, Cache
logger = init_logger(__name__)
# ---- MoE dispatch custom op (torch.compile-safe) -----------------------
# Registered as torch.ops.vllm.moe_dispatch so torch.compile treats it as
# opaque. The real implementation decides custom kernel vs NCCL at runtime
# based on tensor sizes.
# Global reference to the MoeAllGather instance per group, set lazily.
_moe_ag_instances: dict[str, Any] = {}
def _get_or_create_moe_ag(group_name: str):
"""Get or lazily create the MoeAllGather for this group."""
if group_name in _moe_ag_instances:
return _moe_ag_instances[group_name]
from vllm.distributed.parallel_state import _groups
group_ref = _groups.get(group_name)
if group_ref is None:
return None
group = group_ref()
if group is None:
return None
dc = getattr(group, "device_communicator", None)
if dc is None:
return None
ca = getattr(dc, "ca_comm", None)
if ca is None or ca.disabled or not getattr(ca, "fully_connected", False):
_moe_ag_instances[group_name] = None
return None
try:
from vllm.jit_kernels.moe_all_gather.moe_allgather import (
MoeAllGather,
)
ag = MoeAllGather(ca)
_moe_ag_instances[group_name] = ag
logger.info("MoE custom all-gather initialized for group %s", group_name)
return ag
except ImportError:
_moe_ag_instances[group_name] = None
return None
def moe_dispatch(
tensors: list[torch.Tensor],
group_name: str,
) -> list[torch.Tensor]:
"""All-gather a list of tensors across the EP/DP group.
Tries the custom one-shot P2P kernel first; falls back to NCCL
all_gatherv if the custom kernel is unavailable or tensors are too large.
"""
from vllm.distributed.parallel_state import _groups
group_ref = _groups.get(group_name)
assert group_ref is not None
group = group_ref()
assert group is not None
# Try custom kernel for small, uniform-sized batches.
n = len(tensors)
if 2 <= n <= 4:
moe_ag = _get_or_create_moe_ag(group_name)
if moe_ag is not None:
ws = moe_ag.world_size
# Compute packed size.
total_bytes = 0
ok = True
for t in tensors:
if not t.is_contiguous():
ok = False
break
nbytes = t.numel() * t.element_size()
if nbytes % 16 != 0:
ok = False
break
total_bytes = (total_bytes + 15) & ~15
total_bytes += nbytes
max_bytes = min(1 * 1024 * 1024, moe_ag.max_size)
if ok and total_bytes <= max_bytes:
outputs = [
torch.empty(
(t.shape[0] * ws, *t.shape[1:]), dtype=t.dtype, device=t.device
)
for t in tensors
]
from vllm.jit_kernels.moe_all_gather.moe_allgather import (
_load_lib,
)
lib = _load_lib()
lib.moe_all_gather(
moe_ag._rank_data_ptr,
moe_ag._rank_signals_ptr,
moe_ag._self_signal_ptr,
moe_ag.rank,
ws,
tensors,
outputs,
)
return outputs
# Fallback: NCCL all_gatherv.
return group.device_communicator.all_gatherv(tensors, dim=0)
def moe_dispatch_fake(
tensors: list[torch.Tensor],
group_name: str,
) -> list[torch.Tensor]:
"""Fake impl for torch.compile tracing — returns empty tensors of the
correct gathered shape."""
from vllm.distributed.parallel_state import _groups
group_ref = _groups.get(group_name)
assert group_ref is not None
group = group_ref()
ws = group.world_size if group is not None else 1
return [
torch.empty((t.shape[0] * ws, *t.shape[1:]), dtype=t.dtype, device=t.device)
for t in tensors
]
direct_register_custom_op(
op_name="moe_dispatch",
op_func=moe_dispatch,
fake_impl=moe_dispatch_fake,
mutates_args=[],
)
if has_flashinfer_nvlink_two_sided():
from flashinfer.comm import Mapping # type: ignore[import-not-found]
from flashinfer.comm.mnnvl import MnnvlConfig # type: ignore[import-not-found]
@@ -155,6 +293,131 @@ class AgRsAll2AllManager(All2AllManagerBase):
def __init__(self, cpu_group, tcp_store_group=None):
super().__init__(cpu_group, tcp_store_group)
self._moe_ag = None # lazy: MoeAllGather instance
def _get_moe_allgather(self, dist_group):
"""Lazily build the standalone MoE all-gather kernel wrapper."""
if self._moe_ag is not None:
return self._moe_ag
import logging
log = logging.getLogger(__name__)
# Check prerequisites: intra-node, NVLink, ca_comm available.
if self.internode:
log.warning("MoE AG: skipped (internode)")
return None
dc = getattr(dist_group, "device_communicator", None)
if dc is None:
log.warning("MoE AG: skipped (no device_communicator)")
return None
ca = getattr(dc, "ca_comm", None)
if ca is None:
log.warning("MoE AG: skipped (ca_comm is None)")
return None
if ca.disabled:
log.warning("MoE AG: skipped (ca_comm disabled)")
return None
if not getattr(ca, "fully_connected", False):
log.warning("MoE AG: skipped (not fully_connected)")
return None
try:
from vllm.jit_kernels.moe_all_gather.moe_allgather import ( # type: ignore
MoeAllGather,
)
except ImportError as e:
log.warning("MoE AG: skipped (import failed: %s)", e)
return None
self._moe_ag = MoeAllGather(ca)
log.info("MoE AG: custom kernel initialized successfully")
return self._moe_ag
# ---- helpers --------------------------------------------------------
def _custom_all_gather(
self,
dist_group,
tensors: list[torch.Tensor],
sizes: list[int],
) -> list[torch.Tensor] | None:
"""Try the one-shot P2P MoE all-gather kernel. Returns None on
fallback (kernel unavailable, too many tensors, etc.)."""
import logging
_log = logging.getLogger(__name__)
_dbg = getattr(self, "_ag_dbg_count", 0)
n = len(tensors)
if n < 2 or n > 4:
if _dbg < 10:
_log.warning(f"MoE AG: n={n} (need 2-4)")
self._ag_dbg_count = _dbg + 1
return None
if any(s != sizes[0] for s in sizes[1:]):
if _dbg < 10:
_log.warning(f"MoE AG: non-uniform sizes {sizes}")
self._ag_dbg_count = _dbg + 1
return None
moe_ag = self._get_moe_allgather(dist_group)
if moe_ag is None:
return None
ws = moe_ag.world_size
total_bytes = 0
for t in tensors:
if not t.is_contiguous():
if _dbg < 10:
_log.warning(f"MoE AG: non-contiguous {t.shape} {t.stride()}")
self._ag_dbg_count = _dbg + 1
return None
nbytes = t.numel() * t.element_size()
if nbytes % 16 != 0:
if _dbg < 10:
_log.warning(
f"MoE AG: unaligned {nbytes}B shape={t.shape} dtype={t.dtype}"
)
self._ag_dbg_count = _dbg + 1
return None
total_bytes = (total_bytes + 15) & ~15
total_bytes += nbytes
if total_bytes > moe_ag.max_size:
if _dbg < 10:
_log.warning(
f"MoE AG: too large {total_bytes} > {moe_ag.max_size} "
f"n={n} shapes={[t.shape for t in tensors]}"
)
self._ag_dbg_count = _dbg + 1
return None
# Allocate per-tensor outputs.
outputs = [
torch.empty((t.shape[0] * ws, *t.shape[1:]), dtype=t.dtype, device=t.device)
for t in tensors
]
from vllm.jit_kernels.moe_all_gather.moe_allgather import ( # type: ignore
_load_lib,
)
lib = _load_lib()
lib.moe_all_gather(
moe_ag._rank_data_ptr,
moe_ag._rank_signals_ptr,
moe_ag._self_signal_ptr,
moe_ag.rank,
ws,
tensors,
outputs,
)
return outputs
# ---- public API -----------------------------------------------------
def dispatch_router_logits(
self,
@@ -176,19 +439,22 @@ class AgRsAll2AllManager(All2AllManagerBase):
dist_group = get_ep_group() if is_sequence_parallel else get_dp_group()
assert sizes[dist_group.rank_in_group] == hidden_states.shape[0]
tensors_to_gather = [hidden_states, router_logits]
tensors = [hidden_states, router_logits]
if extra_tensors is not None:
tensors_to_gather.extend(extra_tensors)
tensors.extend(extra_tensors)
gathered_tensors = dist_group.all_gatherv(
tensors_to_gather,
dim=0,
sizes=sizes,
)
# Use custom op so torch.compile treats this as opaque.
# Runtime decides custom kernel vs NCCL based on tensor sizes.
if all(s == sizes[0] for s in sizes):
gathered = torch.ops.vllm.moe_dispatch(
tensors, group_name=dist_group.unique_name
)
else:
gathered = dist_group.all_gatherv(tensors, dim=0, sizes=sizes)
if extra_tensors is not None:
return (gathered_tensors[0], gathered_tensors[1], gathered_tensors[2:])
return gathered_tensors[0], gathered_tensors[1]
return (gathered[0], gathered[1], gathered[2:])
return gathered[0], gathered[1]
def dispatch(
self,
@@ -202,7 +468,7 @@ class AgRsAll2AllManager(All2AllManagerBase):
| tuple[torch.Tensor, torch.Tensor, torch.Tensor, list[torch.Tensor]]
):
"""
Gather hidden_states and router_logits from all dp ranks.
Gather hidden_states, topk_weights, topk_ids from all dp ranks.
"""
dp_metadata = get_forward_context().dp_metadata
assert dp_metadata is not None
@@ -211,24 +477,24 @@ class AgRsAll2AllManager(All2AllManagerBase):
dist_group = get_ep_group() if is_sequence_parallel else get_dp_group()
assert sizes[dist_group.rank_in_group] == hidden_states.shape[0]
tensors_to_gather = [hidden_states, topk_weights, topk_ids]
tensors = [hidden_states, topk_weights, topk_ids]
if extra_tensors is not None:
tensors_to_gather.extend(extra_tensors)
tensors.extend(extra_tensors)
gathered_tensors = dist_group.all_gatherv(
tensors_to_gather,
dim=0,
sizes=sizes,
)
if all(s == sizes[0] for s in sizes):
gathered = torch.ops.vllm.moe_dispatch(
tensors, group_name=dist_group.unique_name
)
else:
gathered = dist_group.all_gatherv(tensors, dim=0, sizes=sizes)
hidden_states = gathered_tensors[0]
topk_weights = gathered_tensors[1]
topk_ids = gathered_tensors[2]
hidden_states = gathered[0]
topk_weights = gathered[1]
topk_ids = gathered[2]
if extra_tensors is None:
return hidden_states, topk_weights, topk_ids
return hidden_states, topk_weights, topk_ids, gathered_tensors[3:]
return hidden_states, topk_weights, topk_ids, gathered[3:]
def combine(
self, hidden_states: torch.Tensor, is_sequence_parallel: bool = False
@@ -246,7 +512,8 @@ class AgRsAll2AllManager(All2AllManagerBase):
return hidden_states
def destroy(self):
pass
if self._moe_ag is not None:
self._moe_ag = None
class DeepEPAll2AllManagerBase(All2AllManagerBase):
@@ -41,8 +41,7 @@ class CudaCommunicator(DeviceCommunicatorBase):
global_ranks,
global_world_size,
)
if "tp" not in unique_name:
# custom allreduce or torch symm mem can be used only by tp
if not any(tag in unique_name for tag in ("tp", "dp", "ep")):
use_custom_allreduce = False
use_torch_symm_mem = False
use_flashinfer_allreduce = False
@@ -50,8 +49,12 @@ class CudaCommunicator(DeviceCommunicatorBase):
from vllm.distributed.parallel_state import _ENABLE_CUSTOM_ALL_REDUCE
use_custom_allreduce = _ENABLE_CUSTOM_ALL_REDUCE
use_torch_symm_mem = envs.VLLM_ALLREDUCE_USE_SYMM_MEM
use_flashinfer_allreduce = envs.VLLM_ALLREDUCE_USE_FLASHINFER
use_torch_symm_mem = (
"tp" in unique_name and envs.VLLM_ALLREDUCE_USE_SYMM_MEM
)
use_flashinfer_allreduce = (
"tp" in unique_name and envs.VLLM_ALLREDUCE_USE_FLASHINFER
)
self.use_custom_allreduce = use_custom_allreduce
self.use_torch_symm_mem = use_torch_symm_mem
@@ -238,8 +241,6 @@ class CudaCommunicator(DeviceCommunicatorBase):
def reduce_scatter(self, input_: torch.Tensor, dim: int = -1):
world_size = self.world_size
pynccl_comm = self.pynccl_comm
assert pynccl_comm is not None
if dim < 0:
# Convert negative dim to positive.
dim += input_.dim()
@@ -249,6 +250,10 @@ class CudaCommunicator(DeviceCommunicatorBase):
input_tensor = input_.movedim(0, dim).contiguous()
assert input_tensor.shape[0] % world_size == 0
# TODO: custom reduce_scatter disabled for now — needs testing
# with torch.compile + CUDA graph pipeline before enabling.
chunk_size = input_tensor.shape[0] // world_size
output_shape = (chunk_size,) + input_tensor.shape[1:]
@@ -256,6 +261,8 @@ class CudaCommunicator(DeviceCommunicatorBase):
output_shape, dtype=input_tensor.dtype, device=input_tensor.device
)
pynccl_comm = self.pynccl_comm
assert pynccl_comm is not None
pynccl_comm.reduce_scatter(output, input_tensor)
# Reshape before returning
@@ -281,6 +281,37 @@ class CustomAllreduce:
# latency) compared to the performance gain of using custom kernels
return self.all_reduce(input, registered=False)
def reduce_scatter(
self,
inp: torch.Tensor,
*,
out: torch.Tensor | None = None,
registered: bool = False,
) -> torch.Tensor:
if out is None:
out_shape = (inp.shape[0] // self.world_size, *inp.shape[1:])
out = torch.empty(out_shape, dtype=inp.dtype, device=inp.device)
if registered:
ops.reduce_scatter(self._ptr, inp, out, 0, 0)
else:
ops.reduce_scatter(
self._ptr, inp, out, self.buffer_ptrs[self.rank], self.max_size
)
return out
def custom_reduce_scatter(self, input: torch.Tensor) -> torch.Tensor | None:
"""The main reduce_scatter API that provides support for cuda graph."""
if self.disabled or not self.should_custom_ar(input):
return None
if self._IS_CAPTURING:
if torch.cuda.is_current_stream_capturing():
return self.reduce_scatter(input, registered=True)
else:
out_shape = (input.shape[0] // self.world_size, *input.shape[1:])
return torch.empty(out_shape, dtype=input.dtype, device=input.device)
else:
return self.reduce_scatter(input, registered=False)
def close(self):
if not self.disabled and self._ptr:
if ops is not None:
View File
+1
View File
@@ -0,0 +1 @@
/home/woosuk/workspace/vllm/csrc/standalone/custom_ar/moe_allgather.cu
+1
View File
@@ -0,0 +1 @@
/home/woosuk/workspace/vllm/csrc/standalone/custom_ar/moe_allgather.py