大模型推理引擎vLLM(29): 参考sglang代码,重构vllm021中EP高吞吐代码,消除空泡问题:400us减小到25us
目录
2.1 代码第 1 步 —— dispatch 只交作业,不拷名单
2.2 代码:第 2 步 —— 同一个函数里:拷名单 + scatter
2.2.4 下一行就是 scatter(里面先 scan 再 scatter)
3.1 代码:第 1 步 —— _receiver 里就把名单拷了
3.1.2 改 topk(sglang 这条 DeepGEMM 路径基本不做)
3.3.2 两个 native fill(你 profile 里看到的)
3.3.3 再数一遍 counts(不用刚才 Memcpy 上去的那份做 scatter)
abstract:
其实这里消除空泡的核心方法就是:看dispatch和scatter之间有哪些cpu调用消耗了时间,然后看看这些cpu调用能不能替换成更省时间的,或者直接删掉,最终效果就是空泡从400us减小成了25us,效果显著。
1 问题描述

上面的这个是vllm的prof图,

这个是sglang的prof,可以看到sglang是没有空泡的,那么把 sglang和vllm的这块代码看懂,然后借鉴sglang的代码,消除下vllm的空泡问题。
2 sglang的这个过程:一件事做完再干下一件
假设 DeepEP 通信刚结束,本 rank 手里有:
- GPU 上的 token 数据
hidden - GPU 上的
topk_ids - CPU 上的一份名单
counts = [3, 5, 2, ...](每个 expert 分到几个 token)
后面要做的事本质一样:把这份名单拷到 GPU,再按名单做 scan + scatter。
差别只在于:这两步中间夹了没有别的事。
sglang的大体过程如下
时间 →
[1] DeepEP dispatch 结束
手里有 counts(还在 CPU 的 List)[2] 马上进 pre_permute 这一个函数
CPU: sum(counts) → 算要开多大 buffer
GPU: empty 开几块内存
GPU: 把 counts 拷上去 ← profile 里的 Memcpy
GPU: 立刻 ep_scatter ← 紧接着 scan + scatter[3] 去做 grouped gemm
Memcpy 和 scatter 写在同一个函数里,前后两行,所以中间几乎没空泡。
2.1 代码第 1 步 —— dispatch 只交作业,不拷名单
sglang/python/sglang/srt/layers/moe/token_dispatcher/deepep.py
def dispatch_b(self, hidden_states, topk_ids, topk_weights, previous_event):
(
hidden_states,
topk_ids,
topk_weights,
num_recv_tokens_per_expert,
event,
) = self._dispatch_core(hidden_states, topk_ids, topk_weights, previous_event)
event.current_stream_wait() if self.async_finish else ()
if isinstance(hidden_states, tuple):
hidden_states, hidden_states_scale = hidden_states
else:
hidden_states_scale = None
return DeepEPNormalDispatchOutput(
hidden_states,
hidden_states_scale,
topk_ids,
topk_weights,
num_recv_tokens_per_expert,
)
- DeepEP 跑完了,通信结束。
num_recv_tokens_per_expert仍然是 CPU 上的List[int]。- 这里没有
.cuda(),所以 这里不会出现你盯的那次 Memcpy。
输出是这样的
DeepEPNormalDispatchOutput(
hidden_states=..., # GPU
hidden_states_scale=..., # GPU
topk_ids=..., # GPU
topk_weights=..., # GPU
num_recv_tokens_per_expert=[3, 5, 2, ...], # CPU list
)
2.2 代码:第 2 步 —— 同一个函数里:拷名单 + scatter
sglang/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py
@register_pre_permute("deepep_normal", "deep_gemm")
def pre_permute_deepep_normal_to_deep_gemm(
dispatch_output: DeepEPNormalDispatchOutput,
quant_info: DeepGemmMoeQuantInfo,
runner_config: MoeRunnerConfig,
running_state: dict,
) -> DeepGemmRunnerInput:
from sglang.srt.layers.moe.ep_moe.kernels import ep_scatter
(
hidden_states,
hidden_states_scale,
topk_ids,
topk_weights,
num_recv_tokens_per_expert,
) = dispatch_output
assert runner_config.activation == "silu"
all_tokens = sum(num_recv_tokens_per_expert)
running_state["all_tokens"] = all_tokens
K = hidden_states.shape[1]
hidden_states_shape = hidden_states.shape
hidden_states_device = hidden_states.device
hidden_states_dtype = hidden_states.dtype
running_state["hidden_states_shape"] = hidden_states_shape
running_state["hidden_states_device"] = hidden_states_device
running_state["hidden_states_dtype"] = hidden_states_dtype
running_state["topk_ids"] = topk_ids
running_state["topk_weights"] = topk_weights
input_tensor = torch.empty(
(all_tokens, K),
device=hidden_states.device,
dtype=hidden_states.dtype,
)
if deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0:
# TODO check whether need `zeros`
input_tensor_scale = torch.zeros(
(ceil_div(K // 128, 4), all_tokens),
device=hidden_states.device,
dtype=torch.int,
).transpose(0, 1)
else:
input_tensor_scale = torch.empty(
(all_tokens, K // 128),
device=hidden_states.device,
dtype=torch.float32,
)
m_indices = torch.empty(all_tokens, device=hidden_states.device, dtype=torch.int32)
output_index = torch.empty_like(topk_ids)
if get_offloader().forbid_copy_engine_usage:
num_recv_tokens_per_expert_gpu = copy_list_to_gpu_no_ce(
num_recv_tokens_per_expert
)
else:
num_recv_tokens_per_expert_gpu = torch.tensor(
num_recv_tokens_per_expert,
dtype=torch.int32,
pin_memory=True,
device="cpu",
).cuda(non_blocking=True)
expert_start_loc = torch.empty_like(num_recv_tokens_per_expert_gpu)
ep_scatter(
hidden_states,
hidden_states_scale,
topk_ids,
num_recv_tokens_per_expert_gpu,
expert_start_loc,
input_tensor,
input_tensor_scale,
m_indices,
output_index,
scale_ue8m0=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
)
dispose_tensor(hidden_states)
dispose_tensor(hidden_states_scale)
running_state["output_index"] = output_index
return DeepGemmRunnerInput(
hidden_states=input_tensor,
hidden_states_scale=input_tensor_scale,
use_masked_gemm=False,
m_indices=m_indices,
)
把上面的代码逐段读一下
2.2.1 拆开 dispatch 的结果
(
hidden_states,
hidden_states_scale,
topk_ids,
topk_weights,
num_recv_tokens_per_expert, # 还是 list
) = dispatch_output
2.2.2 纯 CPU:算总长度,开输出 buffer
all_tokens = sum(num_recv_tokens_per_expert) # CPU 加法,不上 GPU
input_tensor = torch.empty((all_tokens, K), ...) # 开输出
input_tensor_scale = torch.empty(...)
m_indices = torch.empty(...) # 注意:empty,不是 full(-1)
output_index = torch.empty_like(topk_ids)
这些是在准备 scatter 要用的空盒子。还没拷 counts。
2.2.3 立刻 HtoD
num_recv_tokens_per_expert_gpu = torch.tensor(
num_recv_tokens_per_expert, # CPU list
dtype=torch.int32,
pin_memory=True,
device="cpu",
).cuda(non_blocking=True) # ← 这里出现 Memcpy
2.2.4 下一行就是 scatter(里面先 scan 再 scatter)
ep_scatter(
hidden_states,
hidden_states_scale,
topk_ids,
num_recv_tokens_per_expert_gpu, # 刚拷上去的 counts
expert_start_loc,
input_tensor,
...
)
所以 sglang 的 GPU 时间线就是:
... dispatch 通信 ... | Memcpy(counts) | scan | scatter | gemm ...
↑________________↑
几乎贴在一起
3 vLLM:同样两件事,但拆成两个房间做
时间 →
[1] DeepEP dispatch 结束(和 sglang 一样)
手里也有 counts: List[int]
[2] 进 _receiver(prepare 收尾) ← 「第一个房间」
torch.where 改 topk_ids
立刻 make_from_list:把 counts 拷到 GPU ← Memcpy 出现在这里!
return,带着 meta 离开这个房间
[3] 回到 modular_kernel ← 「走廊」
_prepare 结束
再调 _fused_experts
再进 DeepGemmExperts.apply
再算 workspace / M_sum ...
(这段 GPU 往往没事干 → 空泡)
[4] 终于进 deepgemm_moe_permute ← 「第二个房间」
torch.full(-1) × 2
count_expert(再数一遍)
才 ep_scatter(scan + scatter)
3.1 代码:第 1 步 —— _receiver 里就把名单拷了
vllm021/vllm/model_executor/layers/fused_moe/prepare_finalize/deepep_ht.py
3.1.1 等通信
if event.event is not None:
event.current_stream_wait()
和 sglang dispatch_b 里 wait 一样,通信结束。
3.1.2 改 topk(sglang 这条 DeepGEMM 路径基本不做)
expert_topk_ids = torch.where(
expert_topk_ids == -1,
...,
expert_topk_ids + self.rank_expert_offset, # local → global
)
3.1.3 立刻 HtoD,也就是memcpy
expert_tokens_meta = mk.ExpertTokensMetadata.make_from_list(
expert_num_tokens_per_expert_list, device=expert_x.device
)
make_from_list 实际干的事:
expert_num_tokens_cpu = torch.tensor(list, device="cpu", pin_memory=True)
return ExpertTokensMetadata(
expert_num_tokens=expert_num_tokens_cpu.to(device, non_blocking=True),
# ↑ 这里就是 Memcpy
expert_num_tokens_cpu=expert_num_tokens_cpu,
)
注意:到这里 还没有 调用 ep_scatter。
函数直接 return 了 token、scale、meta、topk。
3.2 代码:中间走廊 —— modular_kernel
a1q, a1q_scale, expert_tokens_meta, topk_ids, topk_weights = self._prepare(...)
# ↑ 里面已经跑完 _receiver → Memcpy 已经发生
fused_out = self._fused_experts(..., expert_tokens_meta=expert_tokens_meta, ...)
# ↑ 这里面很晚才调到 deepgemm_moe_permute → 才 scatter
_prepare 和 _fused_experts 之间,CPU 还在调 Python、进 experts、算 workspace。
GPU 上 counts 已经拷完了,但 scan/scatter 还没 enqueue → profile 里就是白的。
3.3 代码:第 2 步 —— 很晚才 scatter
vllmhcu021/vllm_hcu/model_executor/layers/fused_moe/deep_gemm_utils.py
def deepgemm_moe_permute(
aq: torch.Tensor,
aq_scale: torch.Tensor,
topk_ids: torch.Tensor,
local_num_experts: int,
expert_map: torch.Tensor | None,
expert_tokens_meta: mk.ExpertTokensMetadata | None,
aq_out: torch.Tensor | None = None,
):
assert aq.ndim == 2
assert topk_ids.dtype.is_signed, "The kernel uses -1 to represent invalid topk_ids"
H = aq.size(1)
device = aq.device
# block_m, block_k = get_mk_alignment_for_contiguous_layout()
block_m = 256
M_sum = compute_aligned_M(
M=topk_ids.size(0),
num_topk=topk_ids.size(1),
local_num_experts=local_num_experts,
alignment=block_m,
expert_tokens_meta=expert_tokens_meta,
)
expert_start_loc = torch.empty(
(local_num_experts), device=device, dtype=torch.int32
)
assert aq_out is None or aq_out.shape == (M_sum, H)
if aq_out is None:
aq_out = torch.empty((M_sum, H), device=device, dtype=aq.dtype)
# aq_scale_out = torch.empty(
# (M_sum, H // block_k), device=device, dtype=torch.float32
# )
aq_scale_out = torch.empty(
(M_sum, aq_scale.shape[-1]), device=device, dtype=torch.float32
)
# DeepGEMM uses negative values in m_indices (here expert_ids) to mark
# completely invalid / padded blocks that should be skipped. We always
# initialize expert_ids to -1 so any row that is not explicitly written
# by the scatter kernel will be treated as invalid and skipped by
# DeepGEMM's scheduler.
expert_ids = torch.full(
(M_sum,),
fill_value=-1,
device=device,
dtype=torch.int32,
)
inv_perm = torch.full(
topk_ids.shape, fill_value=-1, device=device, dtype=torch.int32
)
# Derive per-expert counts from topk_ids so ep_scatter layout matches the
# indices written into inv_perm (dispatch meta can diverge after remap).
expert_num_tokens = count_expert_num_tokens(
topk_ids, local_num_experts, expert_map
)
ep_scatter(
recv_x=aq,
recv_x_scale=aq_scale,
recv_topk=topk_ids,
num_recv_tokens_per_expert=expert_num_tokens,
expert_start_loc=expert_start_loc,
expert_map=expert_map,
output_tensor=aq_out,
output_tensor_scale=aq_scale_out,
m_indices=expert_ids,
output_index=inv_perm,
)
return aq_out, aq_scale_out, expert_ids, inv_perm
3.3.1 用之前拷上来的 meta 算对齐长度(CPU)
M_sum = compute_aligned_M(..., expert_tokens_meta=expert_tokens_meta)
3.3.2 两个 native fill(你 profile 里看到的)
expert_ids = torch.full((M_sum,), fill_value=-1, ...)
inv_perm = torch.full(topk_ids.shape, fill_value=-1, ...)
3.3.3 再数一遍 counts(不用刚才 Memcpy 上去的那份做 scatter)
expert_num_tokens = count_expert_num_tokens(topk_ids, local_num_experts, expert_map)
3.3.4 这才 scan + scatter
ep_scatter(..., num_recv_tokens_per_expert=expert_num_tokens, ...)
所以 vLLM 的 GPU 时间线是:
... dispatch ... | Memcpy | ........空白........ | fill | fill | count | scan | scatter | gemm
↑ ↑
_receiver 里 permute 里才到
4 消除空泡方法1
通过prof发现,在memcpy之后,还有很多cpu调用,于是要想办法减少这些cpu调用,
def compute_aligned_M(
M: int,
num_topk: int,
local_num_experts: int,
alignment: int,
expert_tokens_meta: mk.ExpertTokensMetadata | None,
):
# Conservative upper bound on permuted rows (M_sum). Safe even when
# dispatch meta under-counts vs post-dispatch topk_ids after DeepEP remap.
M_sum_upper = (M * num_topk) + local_num_experts * (alignment - 1)
M_sum_upper = round_up(M_sum_upper, alignment)
# Fast path: reuse cached sum(list) from make_from_list (no aten round_up),
# but still take max with upper bound for safety.
if expert_tokens_meta is not None and expert_tokens_meta.m_sum is not None:
return max(expert_tokens_meta.m_sum, M_sum_upper)
if (expert_tokens_meta is not None) and (
expert_tokens_meta.expert_num_tokens_cpu is not None
):
M_sum_meta = expert_num_tokens_round_up_and_sum(
expert_tokens_meta.expert_num_tokens_cpu, alignment=alignment
)
return max(M_sum_meta, M_sum_upper)
return M_sum_upper
通过分析prof发现,其中一个函数被调用了很多次,而通过sglang代码以及添加打印发现,其实这里不需要这么复杂,因为dispatch接口已经传入了256对齐了,所以之类计算的时候,只需要简单的一个sum函数就可以解决
@dataclass
class ExpertTokensMetadata:
"""
Metadata regarding expert-token routing.
"""
expert_num_tokens: torch.Tensor
expert_num_tokens_cpu: torch.Tensor | None
m_sum: int | None = None
@staticmethod
def make_from_list(
expert_num_tokens_list: list[int], device: str
) -> "ExpertTokensMetadata":
expert_num_tokens_cpu = torch.tensor(
expert_num_tokens_list, device="cpu", dtype=torch.int32, pin_memory=True
)
return ExpertTokensMetadata(
expert_num_tokens=expert_num_tokens_cpu.to(device, non_blocking=True),
expert_num_tokens_cpu=expert_num_tokens_cpu,
m_sum=sum(expert_num_tokens_list),
)

这样修改之后,空泡有所减小,


但还是不够,需要继续修改。
5 消除空泡方法2
那么继续看,还有什么,


那么接下来去看vllm在memcpy之后,cpu在干什么



那么vllm中间的cpu调用是哪些东西





这里把allocate_buffer里面的这个替换了一下

6 消除空泡方法3

刚才从prof看到,这里的import也占用了时间,于是这里加个判断,只有ep的时候才走下面的代码
7 消除空泡方法4
刚才有个误区,老是看memcpy之后的cpu调用,其实应该再往前看,看memcpy之前的有哪些调用可以优化,发现了一个

这个torchwhere在deepep_ht.py文件中,这里给他删掉


现在新路径,不用全局的了,不用expertmap了,探后topkids里面就是局部的,然后scatter也是直接用局部的,
就是本来吧,这个topk_ids在distapch之后收到的里面的是本地局部的专家,并且里面是带有负一的,然后这个torch.where给他加上了偏置,把局部的都给转成了全局的,然后scatter里面到时候还要根据expertmap给把这个topk_ids给再转成局部的才做scatter,
以前的路径多此一举
去掉torch.where之后,这四个算子都没了

8 其他消除空泡方法
其实就是和上面一样,还是看dispatch和scatter之间有哪些cpu调用消耗了时间,然后看看这些cpu调用能不能替换成更省时间的,或者直接删掉,就这样一步步来。
9 总结
下面是最终消除空泡前后的对比图,


更多推荐





所有评论(0)