[Docs] PyPTO 性能优化案例:DeepSeek-V4 MoE 预处理算子融合与同步开销复盘 - #825
Conversation
📝 WalkthroughWalkthroughThe MoE pipeline now computes normalization intermediates during preprocessing, passes them into a dedicated gating entry point, and produces quantized routing inputs for dispatch. Expert sizing derived from ChangesDeferred MoE normalization and gating
Estimated code review effort: 4 (Complex) | ~45 minutes Sequence Diagram(s)sequenceDiagram
participant moe
participant hc_pre_moe_norm
participant gate_from_norm
participant dispatch
moe->>hc_pre_moe_norm: preprocess tokens and produce norm buffers
moe->>gate_from_norm: route using norm buffers
gate_from_norm->>dispatch: provide x_norm_i8, x_norm_scale, indices, weights
Possibly related PRs
Poem
🚥 Pre-merge checks | ✅ 5✅ Passed checks (5 passed)
Thanks for using CodeRabbit! It's free for OSS, and your support helps us grow. If you like it, consider giving us a shout-out. Comment |
74ba698 to
4ab10f4
Compare
Summary
split_pre_post,mix_x, andffn_norminto one external AIV task.GM soft barriers, while preserving
sync_start=True.comb_sinkhornbehinddispatch_gather.synchronization/ABI contract tests.
global / 16 * EPfor the updated partition.阅读导航(ADHD-friendly)
这部分故意写成可以复用的“算子融合操作手册”。不用一次全部读完:
0. 一屏结论
融合前
三个小 AIV 算子在关键路径上串行。任务边界原本提供正确的完成顺序,但也带来
三次 launch 和两个 Scheduler gap。
融合后
必须同时成立的 7 个不变量
core_num、C++ grid stride、soft barrier participant 数和workspace slot 数描述的是同一组 8 个逻辑 AIV。
get_block_idx(args)的 runtime 逻辑编号,不能使用裸
get_block_idx()的物理核编号。sync_start=True,保证 8 个会忙等的 participant 全部驻留后再启动。InOut[64, INT32]workspace。pre_val_store和x_mixed还需要单独做DCCI publish/invalidate + DSB。
1. 如何分析融合前的算子
1.1 先画数据流,不要先看“哪些函数写在一起”
对每个候选 scope 记录:
本例 decode
T=8的关键表格:hc_pre_rmsx_flatinv_rmshc_pre_linearx_flat,hc_fnmixes_rawsplit_pre_postinv_rms,mixes_raw, base/scalepre_val_store,postmix_xpre_val_store,x_flatx_mixedffn_normx_mixed,norm_wcomb_sinkhorninv_rms,mixes_rawcomb选择
split_pre_post + mix_x + ffn_norm的原因:1.2 找到“重新切分”的数据边
删除 task boundary 后,Scheduler 不再替 phase 提供完成顺序。只有在 consumer
可能读取其他 lane 产生的数据时,才需要 grid barrier。
本例有两条这样的边:
第一条边:
8 x 8 FP32;第二条边:
[1, 4096]。所以单个 lane 内的程序顺序不够,必须有两个跨 lane barrier。
1.3 为什么不把更多算子一起融合
没有融合
hc_pre_rms和hc_pre_linear:没有融合
comb_sinkhorn:hc_post之前完成。通用原则:
2. 依赖图如何修改
只写 fused kernel 不够。必须在 Scheduler graph 中重新表达全部 happens-before。
2.1 把外部 producer TaskId 显式返回
hc_pre_moe_producers()只保留 unfused RMS/linear,并捕获:fused submit 显式依赖它们:
不要在文档或测试中写死数值 TaskId;运行期 TaskId 会随生成结果变化,只检查符号边。
2.2 从 gate 中拆出
ffn_norm旧
gate()中的ffn_norm产生:融合步骤:
ffn_norm作为 fused 的第三个 generated body。gate_precomputed(),直接消费这些 buffer。fused_pre_norm_tid作为norm_tid传给 gate。x_norm_quant、gate_pre_route、gate使用deps=[norm_tid]。这样 standalone
ffn_norm仍可单测,MoE 路径则复用 fused 结果。2.3 把
comb_sinkhorn延后到dispatch_gatherdispatch()现在返回_gather_tid。comb_sinkhorn改成:为什么使用
with pl.spmd:with ... as tid表达整个 grid 的 TaskId 和显式deps;pl.tile.get_block_idx()取得 block index。延后为什么安全:
dispatch_gather是 producer 的传递后继。gather 完成时,inv_rms/mixes_raw必然已经完成,所以 comb 此时读取它们是安全的。
2.4 为什么需要
manual_dep=True显式 deps 会与自动依赖合并。如果
inv_rms/mixes_raw仍由 OverlapMap 自动建边,comb 的实际 fanin 会重新变成:
这会把想迁移的旧依赖加回来。生产路径因此对这两个跨显式边界的 buffer 使用
manual_dep=True:dispatch_gather_tid;manual_dep=True不是通用性能开关。只有能够给出完整的传递顺序证明时才可使用。2.5
sync_start和allow_early_resolve是两件事sync_start=True:allow_early_resolve=True:sync_start。依赖修改后的检查项:
comb_sinkhorn是否只有dispatch_gather一条显式 dep?3. 如何构造 fused external kernel
3.1 复用 external CCE bridge 的结构
整体结构参考 external-kernel PR(例如
#765):
这里复用的是 bridge/ABI 分层,不是 #765 的 attention 算法,也不照搬它的同步模板。
3.2 generated body 是数学实现的 source of truth
三个 standalone PyPTO scope 分别编译成:
做法:
logical_block和原 phase 的 logical work 数。fused_body.hpp只负责 ABI、循环、同步和调用顺序。generated instruction sequence 尽量保持原样。不要手工“简化”:
TASSIGN的 UB 地址;standalone scope 改变后,应重新生成对应 header,并重新跑 UB high-water/contract。
generated body 末尾的
PIPE_ALL只 drain 本核流水,不是跨核 barrier,也不能保证其他核看到业务 GM 的新值。
3.3 Python/C++ ABI 必须一一对应
PyPTO 会把全部 Tensor 参数放在 scalar 参数之前。production 使用 13 个 Tensor:
x_mixedx_flatinv_rmsmixes_rawhc_basenorm_wpre_val_storepostxg_bufffn_inv_rms_bufxn_scale_bufx_norm_scalesync_workspace随后是:
debug ABI 最后再追加
stop_after。C++ 侧锁定:
workspace 是内部同步状态,不是业务输出,所以返回 tuple 保持七项。随意改成八项会
移动 alias/解包顺序,给所有下游制造 ABI 风险。
3.4 保留每个 phase 原来的 logical mapping
wrapper 计算:
三个 phase 均使用:
decode
T=8的 work 是(1, 4, 8)。三个 phase 串行,所以参与核数是:不是
1 + 4 + 8 = 13。prefill
T=128可能是(16, 64, 128)。8 lane 通过 grid-stride 多轮执行,因此正确性仍成立;但性能最优 participant 数需要单独调优。
3.5 Python submit
if False中的 direct extern call 仅用于 materializeself.fused_pre_norm_ccespecialization,最终会被删除。4. 如何选择同步
4.1 通用决策顺序
每删除一条 task edge,都按顺序问:
sync_start=True。4.2 第一版:48-AIV 硬同步
第一版 correctness baseline 使用:
并启动 48 个 AIV。它证明了:
但它不是最终性能方案:
hc_pre_rms和后继 gate 也需要 AIV;本 kernel 是 AIV-only,不能直接照抄 mixed-core kernel 的
SyncAll<false>()模板;同步类型必须从实际 core topology 推导。4.3 最终版:indexed 8-AIV soft barrier
为什么不用当前 public 三参数 soft API:
0..7动态调度到任意物理 AIV;get_block_idx()返回物理编号,可能稀疏或不从 0 开始;[0,8)workspace 会越界或永远等不到某些 slot;get_block_idx(args)返回 dense runtime logical ID。所以统一使用:
同一个
lane同时用于:长期应给 PTO-ISA 增加公共四参数 indexed overload;本 PR 使用 simpler 当前 pin
中已有的 internal helper。
4.4 为什么绝不能删除
sync_start=TrueGM soft barrier 会忙等并持续占用 AIV。如果 grid 只启动一部分:
sync_start=True负责先集结完整的 8-lane cohort,再允许 kernel 运行。soft sync 的作用是把同步集合从 48 缩到 8;它不替代 Scheduler 的同时驻留保证。
4.5 GM workspace 和 UB scratch
PTO helper 的 slot ABI 是每个 participant 8 个
INT32:每个 fused grid 创建:
必须满足:
InOut;pl.no_dep/manual_dep隐藏它。generation:
共享 workspace 会让不同 grid 互相“代打卡”。
UB scratch 固定在 176 KiB:
static/contract 检查:
TASSIGN,确保所有 tile end 低于 176 KiB。注意两个不同概念:
不能为了“对齐 64B”只在 wrapper 中擅自把 slot stride 改为 16,否则会和 helper
寻址/Python ABI 不一致。这个 32B slot/64B cache-line 差异需要靠 A2/A3 并发压力
测试持续覆盖。
4.6 barrier 到达不等于业务 GM 可见
soft barrier 自带的 DCCI 只维护同步 counter,不会刷新
pre_val_store/x_mixed。正确顺序:
范围:
pre_val_store:一个 split block 写8*8*4=256B,即 4 条 64B line;x_mixed:一个 mix block 在 8 个不连续 row 中分别写 2048B slice;4096*2=8192B。x_mixed的 8 个 slice 不是连续区域,必须逐 row 处理;不能把它误当成一段连续16KiB,否则会触碰其他 lane 拥有的 slice,增加 false sharing。
当前 DCCI 是 correctness-first,可能比较贵。必须先通过真机 golden/压力测试,再
基于泳道证据缩小或合并范围。
5.
InOut、specialization 与 dump5.1 workspace 为什么是
InOutadd_inout(sync_workspace)。workspace 必须放在全部 Tensor 之后、scalar 之前;TaskArgs 不允许 scalar 后再放
Tensor。
5.2
if False为什么使用另一份 workspacePyPTO 会在 constant-false elimination 前检查 direct extern call。一个变量传入
InOut后,原 pre-call SSA value 就被消费。如果 specialization direct call 和真实
spmd_submit共用同一 workspace,会触发:因此分支中单独创建:
它随
if False一起消除。codegen audit 会拒绝最终 orchestration 中残留的specialization-only buffer。
5.3 正确的 dump 位置
pl.dump_tag是 forward-sticky 标记。只在真实 submit 前执行一次:runtime 会保存:
BEFORE_DISPATCH;AFTER_COMPLETION。submit 后不能再次引用旧的 workspace SSA value,否则违反 InOut use discipline。
七个业务输出在 submit 后单独
dump_tag。6. dump/debug 定位方法
6.1 phase-stop debug entry
--debug-stopsplitbarrier1mixbarrier2full示例:
6.2 workspace dump 期望值
8 个 slot 的 head 位于:
期望:
6.3 症状 -> 第一检查点
core_num==8、kAivLanes==8、sync_start=True、logical IDmix_x错pre_val_storepublish/invalidatemixstop 正确但 full 错x_mixed可见性或 ffn logical mappingInOutUseDisciplinemanual_dep缺失或生成 deps 又含 rms/linear7. 分层验证
Level 1:contract/static
contract 锁定:
sync_start=true;AscendC::SyncAll、隐式 public hard sync、FFTS sync、裸物理get_block_idx();Level 2:compile-only
分别编译:
检查生成 C++:
Level 3:单 kernel 真机 golden
至少覆盖:
0/1很重要:业务 loop 可以为空/很小,但 8 个 lane 仍必须完整通过 barrier。Level 4:debug stop + counter
按顺序跑:
同时检查业务输出和 workspace generation。
Level 5:完整分布式 MoE
至少两 rank、balanced routing,检查:
x_nextgolden;Level 6:并发/泳道/性能
需要覆盖:
0 -> 1 -> 2;Perfetto/泳道重点看:
rms -> fused、linear -> fused;fused -> x_norm_quant/gate;dispatch_gather -> comb,且 comb fanin 不应再出现 rms/linear;不要跨不同 profiling level/输入比较绝对时间。CPU sim 也不能证明动态物理核映射
和真实 DCCI 行为。
8. 可复用逐步清单
Phase A:证明融合边界
Phase B:准备 Scheduler 依赖
manual_dep=True。Phase C:构建 external kernel
Phase D:选择同步
sync_start=True。InOut。Phase E:让问题可定位
Phase F:逐级验收
9. 常见错误
(1,4,8)需要 8,不是 13。get_block_idx():物理 ID 不保证是 dense0..N-1。sync_start=True:partial grid 可占住全部资源等待未调度 lane。num_tokens=max:空/tail work 才容易暴露 barrier bug。利用率。
workspace 和
sync_start,不能依赖超时后行为。主要文件
models/deepseek/v4-flash/hc_pre.pycomb_sinkhorn。models/deepseek/v4-flash/gate.pyffn_norm;gate_precomputedconsumer。models/deepseek/v4-flash/fused_pre_norm_cce.pymodels/deepseek/v4-flash/kernels/fused_pre_norm_cce/models/deepseek/v4-flash/moe.pytests/contract/test_deepseek_v4_moe_fusion.pyVerification performed
tests/contract/test_deepseek_v4_moe_fusion.py: 14 passed.num_tokens=0/1/8.num_tokens=0/1/128.x_nextgolden and ABI audit.0 -> 1 -> 2.Follow-ups
可能需要更大 cap。