尧图网站建设 尧图网络
  • 首页
  • 关于我们
  • 服务项目
  • 案例展示
  • 建站流程
  • 资讯中心
  • 联系我们
首页/资讯中心/详情

大模型推理引擎vLLM(30): 参考sglang代码,重构vllm021中EP高吞吐代码,消除空泡问题:400us减小到25us

大模型推理引擎vLLM(30): 参考sglang代码,重构vllm021中EP高吞吐代码,消除空泡问题:400us减小到25us
📅 发布时间:2026/7/24 23:26:27

目录

1 问题描述

2 sglang的这个过程:一件事做完再干下一件

2.1 代码第 1 步 —— dispatch 只交作业,不拷名单

2.2 代码:第 2 步 —— 同一个函数里:拷名单 + scatter

2.2.1 拆开 dispatch 的结果

2.2.2 纯 CPU:算总长度,开输出 buffer

2.2.3 立刻 HtoD

2.2.4 下一行就是 scatter(里面先 scan 再 scatter)

3 vLLM:同样两件事,但拆成两个房间做

3.1 代码:第 1 步 —— _receiver 里就把名单拷了

3.1.1 等通信

3.1.2 改 topk(sglang 这条 DeepGEMM 路径基本不做)

3.1.3 立刻 HtoD,也就是memcpy

3.2 代码:中间走廊 —— modular_kernel

3.3 代码:第 2 步 —— 很晚才 scatter

3.3.1 用之前拷上来的 meta 算对齐长度(CPU)

3.3.2 两个 native fill(你 profile 里看到的)

3.3.3 再数一遍 counts(不用刚才 Memcpy 上去的那份做 scatter)

3.3.4 这才 scan + scatter

4 消除空泡方法1

5 消除空泡方法2

6 消除空泡方法3

7 消除空泡方法4

8 总结


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()

和 sglangdispatch_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 总结

下面是最终消除空泡前后的对比图,

相关新闻

  • 器,生成复指数信号与本振信号相乘,在ip核设置的过程中主要由三个模式 BYPASS 这个又叫直通模式,即不进行任何数字混频,基带信号 ...
  • 广州服饰品牌企业做GEO服务商怎么选?2026年五家服务商深度测评与靠谱选型指南 - 企业新闻快传
  • 模型网关不是多接几家 API 当接口人——给厨房配个会比价、会兜底、会记账的采购总管

最新新闻

  • XHS-Downloader:小红书内容采集的终极实用指南
  • 梧州专业丙烯酸球场地坪漆厂家选择指南 - 热点品牌推荐
  • 安庆酒店推拉门厂商选型指南及本地优质供应方详情 - 热点品牌推荐
  • 用 Agent Skills 武装自媒体起号:一个面向中文内容运营的 AI 框架包解析
  • 捷克的Payroll是什么?主要涉及哪些关键领域?
  • 算法入门(6)——线性数据结构

日新闻

  • 从国家条件到买方清单,深入理解 ABAP CDS 单值过滤器派生
  • 2026 年当下,齐齐哈尔专业的不锈钢闸门批发厂家哪个好,揭秘!这个工业“铁门”如何实现成本翻倍的效率提升? - 行业甄选官
  • 2026阳极氧化加工厂推荐:从设备规模看硬质氧化技术的成熟应用推荐百正机械 - 栗子测评

周新闻

  • SaaS软件行业GEO实践:AI搜索时代的品牌可见性与获客新路径
  • 什么是PCTFE?医药高端包装的“防潮王牌“材料
  • 【JVM调优实战】16-可视化利器-JConsole-VisualVM-JMC

月新闻

  • 2026年6月公司网站搭建最新热门渠道测评:四大低成本/零代码平台对比+避坑
  • 【Linux】Linux arm 编译QT程序,出现expected “}“报错
  • 【MATLAB例程】四基站二维AOA定位与距离辅助增强对比仿真。基于角度观测和测距修正的固定目标平面定位精度分析

关于尧图

  • 公司简介
  • 团队介绍
  • 企业文化
  • 荣誉资质

服务项目

  • 定制开发
  • 电商建站
  • UI 设计
  • 运维服务

快速链接

  • 案例展示
  • 建站流程
  • 常见问题
  • 资讯中心

联系方式

  • 📍北京市朝阳区互联网产业园 A 座 10 层
  • 📞400-888-8888
  • ✉️contact@rkmt.cn
  • 🕐周一至周日 9:00-21:00

© 2024 北京尧图网络科技有限公司 版权所有 | 京 ICP 备 XXXXXXXX 号