CANN SHMEM 经典 MoE Dispatch 示例详解:基于对称窗口的 Token 路由与 Combine 辅助信息生成
【免费下载链接】shmemCANN SHMEM 是面向昇腾平台的多机多卡内存通信库,基于OpenSHMEM 标准协议,实现跨设备的高效内存访问与数据同步。项目地址: https://gitcode.com/cann/shmem
本篇文章深入解析 CANN SHMEM 开源仓库中 examples/dispatch/dispatch_classic 这一经典 MoE(Mixture of Experts)Dispatch 示例。它展示了如何基于 OpenSHMEM 风格的对称数据窗口(SHMEM window)与 MTE 数据搬运能力,将本 PE(Processing Element)的 token/topK 按expert_ids路由到目标 expert 所在 PE,并同步生成 combine 阶段所需的辅助信息。读完本文,你将掌握该算子的输入输出语义、Kernel 三阶段实现原理、构建与运行方法、全部命令行参数含义,以及如何用内置的性能测试框架对 dispatch 进行多 shape 扫描与基准测量。
示例定位与核心价值
该示例实现了非量化 MoE 的经典 dispatch 算子,对应设计文档DOC/moe_dispatch_combine_non_quant_architecture.md。在 MoE 前向计算中,token 需要根据路由结果(expert id)被分发到保存对应 expert 权重的设备上完成计算,计算完成后再把结果聚合回原设备,前者即为 Dispatch,后者为 Combine。本示例聚焦 Dispatch 阶段,其职责是:根据expert_ids将本 PE 的 token/topK 路由到目标 expert 所在 PE,并生成 combine 阶段需要的辅助信息。
从实现策略上,该示例与同目录下的 双平面 MoE Dispatch 示例 形成对照:经典版本优先保证逻辑清晰、可验证和输出顺序稳定,数据面统一使用 MTE 传输;双平面版本则在此基础上同时启用 MTE 与 SDMA 自适应选择传输路径,追求通信阶段性能。建议先以经典版本建立正确性基线,再用双平面版本做性能对比。
功能说明:输入与输出语义
输入
| 输入 | 说明 |
|---|---|
x | 本 PE 的 token hidden states,形状为(bs, h) |
expert_ids | 每个 token/topK 对应的全局 expert id,形状为(bs, topk) |
输出
| 输出 | 说明 |
|---|---|
expand_x | 按目标 expert 聚合后的 token 数据 |
assist_info_for_combine | combine 回传使用的辅助信息,格式为[src_rank_id, src_token_id, src_topk_id] |
ep_recv_count | 按(local_expert_id, src_rank_id)排列的累计接收计数 |
expert_token_nums | 每个本地 expert 实际收到的 token 数 |
输出顺序固定为(local_expert_id, src_rank_id),这一确定性的顺序是后续 combine 阶段与 golden 校验的基础。其中ep_recv_count保存的是累计计数(cumulative count),即从第一个 segment 累加到当前 segment 的 token 总数,等价于每个 segment 在最终expand_x中的前缀偏移边界;expert_token_nums则是按本地 expert 维度聚合的段内 token 数。
实现逻辑:Kernel 三阶段流水
Kernel 启动pe_size个 AIV core,每个 active core 负责一个目标 rank(core_id == dst_rank的 core 负责写入对应目标 PE 的对称窗口)。整个流程分为发送、等待与 compact 三个阶段。
Stage 1:发送阶段(路由与写入)
- 每个 active core 扫描本 PE 的
expert_ids,统计发往自己负责目标 PE 的每个 local expert 的 token 数(对应代码中的segment_counts与slot_offsets数组)。 - 对每个命中的 token/topK,通过 MTE
put_nbi写 payload 到目标 PE 的 SHMEM window。 - 写入
assist_info_for_combine、token ready flag 和 segment count 信号。 - 所有发送 core 完成后,
aclshmem_quiet()保证写操作对远端可见。
Stage 2:等待阶段
接收端(aiv_index == 0的 core)等待各来源 rank 的 count 信号,通过aclshmem_signal_wait_until轮询count_base中的 ready 标志,随后累加得到累计接收计数。
Stage 3:compact 阶段
根据 count 构造ep_recv_count和expert_token_nums,再按确定顺序((local_expert_id, src_rank_id))将 payload 与辅助信息 compact 到最终输出expand_x、assist_info_for_combine。compact 统一放在 core 0 上执行,保证前缀边界与输出写入按唯一确定顺序完成。
对称窗口的内存布局
从 dispatch_kernel.cpp 可以确认,SHMEM window 由四段连续区域组成:
payload(token 数据) -> assist(辅助信息) -> ready(token 级 ready flag) -> count(segment 计数)对应 host 侧 main.cpp 中的DispatchWindowBytes/DispatchCountOffset计算,每个 region 都按 32 字节对齐。Kernel 在栈上为每个目标 local expert 分配固定工作区(segment_counts与slot_offsets,容量上限为DISPATCH_MAX_LOCAL_EXPERT_NUM = 1024),这是expertPerPe上限的由来。
关键 API 调用链
aclshmemx_mte_put_nbi:非阻塞 MTE 远端写,用于 payload 与 ready 信号(数据面)。aclshmemx_signal_op(..., ACLSHMEM_SIGNAL_SET, dst_rank):向远端写辅助信息字段(控制面)。aclshmem_signal_wait_until(..., ACLSHMEM_CMP_EQ, ...):接收端等待 ready/count。aclshmem_quiet():同步本地写队列。aclshmemi_sync_core_soft():Kernel 内 core 间软同步。
构建
在仓库根目录执行bash scripts/build.sh,按目标平台追加参数:
- A2/A3 平台:
bash scripts/build.sh -examples- Ascend950 平台:
bash scripts/build.sh -soc_type Ascend950 -examples构建产物包括build/bin/dispatch(host 可执行程序)与链接进可执行程序的 kernel 库。示例通过 examples/CMakeLists.txt 中的aclshmem_add_collective_example(dispatch)宏注册:dispatch_kernel.cpp被编译为dispatch_kernel共享库,main.cpp编译为dispatch可执行文件,两者均链接shmem库并包含examples/utils等头文件目录。
运行
基础 2 卡测试
cd examples/dispatch/dispatch_classic bash scripts/run.sh -pes 2 -bs 8 -h 16 -topk 2 -expertPerPe 2 -type int32_t8 卡、64 expert 测试
cd examples/dispatch/dispatch_classic bash scripts/run.sh -pes 8 -bs 8 -h 16 -topk 2 -expertPerPe 8 -type int32_trun.sh的执行流程(见 scripts/run.sh):
- 调用 data_gen.py 生成输入
x、路由矩阵expert_ids与 golden 数据(写入golden/shape_<bs>_<h>_<topk>_<moe_expert_num>_<pes>/rank_<id>/,并附带meta.json记录 shape 元信息)。int32_t使用均匀随机整数,float16_t/bfloat16_t使用均匀浮点并依赖ml_dtypes包。 - 为每个 PE(rank)启动一个
dispatch进程,通过SHMEM_UID_SESSION_ID与自适应端口(8766 + case_index % 1000)区分多个 case 的 bootstrap 会话。 - 各进程从 golden 目录读入输入,完成 kernel 执行后把结果写入
output/(expand_x_<pe>.bin、assist_info_<pe>.bin、ep_recv_count_<pe>.bin、expert_token_nums_<pe>.bin)。 - 调用 check_dispatch.py 对四个输出逐一与 golden 比对:整数类型使用严格相等,浮点类型使用
rtol=1e-3, atol=1e-3的np.allclose容差校验。
host 侧 main.cpp 负责解析参数、aclInit/aclrtSetDevice、通过aclshmemx_init_attr初始化 SHMEM 运行时(单实例,堆大小 1GB)、按-type分发到int32_t/fp16_t/bf16_t模板实例,最后依次执行aclshmem_barrier_all、aclshmem_finalize与aclFinalize。
参数说明
run.sh支持的完整参数如下:
-pes <n> PE 数量,单机示例要求与 -gnpus 相同(默认 2)。 -gnpus <n> 本机启动的 NPU 数量,必须与 -pes 相同(默认 2)。 -bs <n> 每个 PE 的 token 数(默认 8)。 -h <n> token hidden size(默认 16)。 -topk <n> 每个 token 路由的 expert 数(默认 2)。 -expertPerPe <n> 每个 PE 上的 local expert 数,范围为 [1, 1024](默认 2)。 -type <dtype> 数据类型,支持 int32_t、float16_t、bfloat16_t(默认 int32_t)。 -fpe <id> 首个 PE 编号(默认 0)。 -fnpu <id> 起始 NPU id(默认 0)。 -ipport <url> SHMEM bootstrap 地址(默认 tcp://127.0.0.1:8766)。 --perf 性能测试模式:保留正确性校验,并将性能 CSV 写入 --output-dir 目录。 --warmup <n> 性能测试的预热迭代次数,不计入统计(默认 5)。 --loops <n> 性能测试的正式测量迭代次数(默认 50)。 -pes-list <a,b,...> 性能测试模式下扫描的 PE 数量列表。 -bs-list <a,b,...> 性能测试模式下扫描的每个 PE token 数列表。 -h-list <a,b,...> 性能测试模式下扫描的 hidden size 列表。 --topk-list <a,b,...> 性能测试模式下扫描的 topk 列表。 --expert-per-pe-list <a,b,...> 性能测试模式下扫描的 local expert 数列表。 --prof-pe <id|all> 性能采集的 PE 编号;为 all 时轮流采集每个 PE 并汇总(默认 0)。 --output-dir <dir> 性能测试 CSV 输出目录(默认 output/perf)。 -a|--analyse <mode> 性能结果处理方式:plot(图形化展示)、md(生成 Markdown 报告)、none(不处理,默认)。expertPerPe上限为 1024:Kernel 在 AI core 栈上为每个目标 local expert 分配固定工作区(segment_counts/slot_offsets),超过该上限会被 host 侧拒绝(main.cpp 会打印max supported value is 1024并返回错误)。此外,run.sh在非 perf 模式下会强制校验-gnpus与-pes相等,否则直接退出。
几个值得注意的实现细节:
-fpe在 run.sh 中仍被保留为一个 CLI 槽位以兼容共享示例脚本,但 host 侧实际将其忽略(见 main.cpp 的注释)。- 设备 id 由
args.pe_id % g_npus + f_npu推导,-fnpu用于支持非 0 起始的 NPU 编号。 run.sh中-pes与-gnpus被联动赋值,保证单机示例下两者始终一致。
性能测试
run.sh --perf会在每个 shape 上保留正确性校验,并在output/perf/下写入 CSV。
单 shape profiling
cd examples/dispatch/dispatch_classic bash scripts/run.sh --perf -pes 2 -bs 8 -h 256 -topk 2 -expertPerPe 2 -type int32_t \ --warmup 5 --loops 50多 shape、多卡数 sweep
cd examples/dispatch/dispatch_classic bash scripts/run.sh --perf --pes-list 2,4,8 --bs-list 8,16,32 --h-list 64,256,1024 \ --topk-list 2 --expert-per-pe-list 2,8 -type int32_t --prof-pe all \ --warmup 5 --loops 50sweep 模式下,run.sh对pes-list × bs-list × h-list × topk-list × expert-per-pe-list做全组合遍历;--prof-pe all时对每个 PE 轮流采集并汇总生成dispatch_perf_summary.csv。
CSV 指标
| 指标 | 含义 |
|---|---|
full_op | 完整 dispatch,包括通信、元数据构造、compact 和同步 |
comm_only | Stage 1 payload 通信及必要的元数据/status 协议 |
单 rank 文件名为dispatch_perf_rank<rank>.csv。
性能数据的产生机制
Kernel 内通过SHMEMI_PROF_START/END(见 dispatch_kernel.cpp)为full_frame_id=0与comm_frame_id=1两个计时框架打点,其中comm_only框架在 Stage 1 结束后即关闭。Host 侧 moe_perf_host.h 中的MoeAppendPerfCsvRows通过aclshmemx_get_prof读取各 core 的周期计数:
- 时间换算系数按 SoC 自适应:Ascend950 平台 1000 cycles 对应 1us,其余平台 50 cycles 对应 1us(
MoeGetCycleToUs)。 - 每行 CSV 记录
DataSize/B、Npus、Blocks(=pe_size)、UBsize/KB(= 190)、Bandwidth/GB/s、CoreMaxTime/us、Metric以及BS/H/TopK/ExpertPerPe/Dtype/Warmup/Loops/ProfPe/CaseId等字段,并追加每个 active core 的SingleCoreTime/us。 - 运行时会通过环境变量
SHMEM_CYCLE_PROF_PE指定被采集的 PE(MoeGetProfPe),非 perf 模式自动以 PE 0 兜底。 - 结果文件可通过
-a plot交给 perf_data_process.py 绘制图表,或以-a md生成 Markdown 报告。
DISPATCH_UB_SIZE_KB = 190对应 Kernel 侧UB_DMA_MAX_SIZE = 190 * 1024(UB 单次 DMA 最大搬运字节数),这是 MTE payload 传输按h个元素分块拷贝时的块大小上限。
从经典走向双平面
完成本示例后,可以继续探索同组的 双平面 MoE Dispatch 示例。双平面版本保持与经典 dispatch 完全一致的输入输出语义与输出顺序,因此同一组 golden/check 脚本可以验证两条路径。双平面的差异在于:当某个(dst_rank, dst_local_expert)segment 的 payload 字节数大于 2MB 且大于当前 PE 的远端平均 segment 大小(判定逻辑使用交叉相乘避免整数截断),该大段改用aclshmemx_sdma_put_nbi传输,小段与全部控制面信号仍走 MTE;SDMA 每提交 256 次 issue 即调用aclshmemx_sdma_quiet防止 outstanding 请求积压。需要注意的是,SDMA 功能要求 CANN 9.0.0 及以上,且当前暂不支持 Ascend950 平台,基础安装与独立 SDMA demo 可参考 examples/sdma/README.md。
对经典与双平面做性能对比时,可对相同 shape 分别执行--perf并比较 CSV 中的comm_only:若comm_only降低,说明大段 payload 走 SDMA 对通信阶段有效;若full_op收益不明显,则需要结合 shape、路由倾斜度和后续 compact/同步成本综合判断。
进一步阅读
- 双平面版本与经典版本的差异与选型建议:dispatch_doubleplane README
- 数据生成与 golden 构造:data_gen.py
- 正确性校验脚本:check_dispatch.py
- 性能 CSV 生成与单位换算:moe_perf_host.h
- 构建系统集成方式:examples/CMakeLists.txt
【免费下载链接】shmemCANN SHMEM 是面向昇腾平台的多机多卡内存通信库,基于OpenSHMEM 标准协议,实现跨设备的高效内存访问与数据同步。项目地址: https://gitcode.com/cann/shmem
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考