CANN ops-transformer MoeTokenUnpermuteWithEp 算子解析:EP 场景下 MoE Token 反重排与加权聚合原理与实战
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
导读
MoeTokenUnpermuteWithEp 是 CANN ops-transformer 中面向专家并行(Expert Parallel,EP)推理/训练场景的核心算子:它根据 sortedIndices 中记录的排序下标,从 permutedTokens 中取回各 token 在对应专家上的中间结果,乘以专家概率(probs)后按 topk 分组做累加合并,还原出每个原始 token 的最终输出。本文以 moe/moe_token_unpermute_with_ep/README.md 为主体,结合 算子 API 文档、调用示例 以及 host 端 tiling 与 kernel 源码,系统讲解其数学模型、参数语义、两段式 aclnn 接口调用流程与底层实现机制。读完本文,你将掌握如何在 CANN 环境下通过 aclnnMoeTokenUnpermuteWithEp 接口完成 MoE 专家输出的反重排与加权合并。
产品支持情况
MoeTokenUnpermuteWithEp 算子在不同硬件产品上的支持情况如下(与 README 及 算子定义源码 中 AICore 配置一致):
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR / Ascend 950DT | √ |
| Atlas A3 训练系列产品 / Atlas A3 推理系列产品 | √ |
| Atlas A2 训练系列产品 / Atlas A2 推理系列产品 | √ |
| Atlas 200I/500 A2 推理产品 | × |
| Atlas 推理系列产品 | × |
| Atlas 训练系列产品 | × |
| Kirin X90 处理器系列产品 | √ |
| Kirin 9030 处理器系列产品 | √ |
在算子定义文件中,op_def为 Ascend 910B(Atlas A2)、Ascend 910_93(Atlas A3)、Ascend 950 三类 AICore 添加了支持动态 Shape、动态 Rank、动态 Format 的config_dyn配置,并为 Kirin X90 / Kirin 9030 添加了config_kirin配置;从源码结构可以推断,产品支持表正是由这些 AICore 配置与对应的 binary 配置文件 驱动的。
功能说明与数学模型
算子功能
在 MoE 模型中,token 会先经过路由被"重排(permute)"并按专家分发计算,MoeTokenUnpermuteWithEp 负责与之对应的"反重排(unpermute)"环节:
- 根据 sortedIndices 存储的下标位置,去获取 permutedTokens 中的输入数据;
- 若提供了 probs,则与 probs 中对应位置的专家概率相乘;
- 按每个原始 token 对应的 numTopk 个位置进行合并累加,得到最终输出 out。
从 kernel 源码 可见,核心判断逻辑是:当acl_token_idx落在[start, end)的有效区间内时执行数据搬运与累加,否则(越界索引或 prob 为 0)该位置的贡献为 0,这一行为与下文"约束说明"中的 rangeOptional 语义完全对应。
计算公式
首先按 rangeOptional 对 sortedIndices 做范围裁剪:
$$ sortedIndices = sortedIndices[rangeOptional[0] \le i < rangeOptional[1]] $$
(1)probs 非 None 时,其中 $i \in {0, 1, 2, ..., num_tokens - 1}$,$j \in {0, 1, 2, ..., numTopk - 1}$,$k \in {0, 1, 2, ..., num_tokens \times numTopk}$:
$$ permutedTokens = permutedTokens.indexSelect(0, sortedIndices) $$
$$ permutedTokens_{k} = permutedTokens_{k} \times probs_{i,j} $$
$$ out_{i} = \sum_{k=i \times numTopk}^{(i+1) \times numTopk - 1} permutedTokens_{k} $$
(2)probs 为 None 时,其中 $i \in {0, 1, 2, ..., num_tokens - 1}$,$j \in {0, 1, 2, ..., numTopk - 1}$:
$$ permutedTokens = permutedTokens.indexSelect(0, sortedIndices) $$
$$ out_{i} = \sum_{k=i \times numTopk}^{(i+1) \times numTopk - 1} permutedTokens_{k} $$
直观理解:indexSelect等价于"按下标收集(gather)",即out[i] = Σ(permutedTokens[sortedIndices[i*numTopk + j]] * probs[i][j])。probs 为 None 时退化为纯求和,不进行乘法。
参数说明
算子参数定义在 算子定义源码 中(输入permuted_tokens、sorted_indices,可选输入probs,输出unpermuted_tokens,属性num_topk、range、padded_mode、restore_shape)。对外暴露的语义如下:
| 参数名 | 输入/输出/属性 | 描述 | 数据类型 | 数据格式 |
|---|---|---|---|---|
| permutedTokens | 输入 | 表示经过扩展并排序过的 tokens,对应公式中的permutedTokens | FLOAT16、BFLOAT16、FLOAT32 | ND |
| sortedIndices | 输入 | 表示需要计算的数据在 permutedTokens 中的位置,对应公式中的sortedIndices | INT32 | ND |
| probsOptional | 可选输入 | 表示输入 tokens 对应的专家概率,对应公式中的probs | FLOAT16、BFLOAT16、FLOAT32 | ND |
| numTopk | 属性 | 被选中的专家个数 | INT64 | - |
| rangeOptional | 属性 | ep 切分的有效范围,size 为 2 | aclIntArray* | - |
| paddedMode | 属性 | 目前仅支持 false | BOOL | - |
| restoreShapeOptional | 属性 | 目前仅支持 nullptr | aclIntArray* | - |
| out | 输出 | 表示 permutedTokens 反重排的输出结果,对应公式中的out | FLOAT16、BFLOAT16、FLOAT32 | ND |
注意:Kirin X90 / Kirin 9030 处理器系列产品不支持 BFLOAT16。从
GetKirinCoreConfig()可以看出,Kirin 平台仅注册了 FLOAT16、FLOAT 两类数据类型组合,这与 README 中的限制一致。
约束说明
- numTopk 必须大于等于 1;probsOptional 非空时,numTopk 必须小于等于 512。tiling 源码中通过
OP_CHECK_IF(inputTopK < 1, ...)与OP_CHECK_IF(topK > 512, ...)强制执行该约束。 - 不支持 Broadcast。
- 不支持 paddedMode 为
True(目前仅支持 false,restoreShapeOptional 仅支持空指针)。 - 当 rangeOptional 为空时,使用默认值
{0, 0},输出为全 0,不会回退调用其他算子。tiling 源码中rangePtr == nullptr时设置start = 0; end = 0,kernel 侧needCopyIn = acl_token_idx >= start && acl_token_idx < end恒为 false,因而输出全 0,两者完全吻合。 - API 层还有额外约束:aclnn 接口默认确定性实现;aclnnTensor 的 shape 不支持使用 -1 表示动态维度或 -2 表示动态 Rank(详见 aclnnMoeTokenUnpermuteWithEp 文档)。
调用说明:两段式 aclnn 接口
本算子通过 aclnn API 调用,入口为 test_aclnn_moe_token_unpermute_with_ep.cpp。与 CANN 其他算子一致,aclnnMoeTokenUnpermuteWithEp 采用两段式接口:先调用aclnnMoeTokenUnpermuteWithEpGetWorkspaceSize获取 workspace 大小与执行器,再调用aclnnMoeTokenUnpermuteWithEp真正执行计算。
函数原型
aclnnStatus aclnnMoeTokenUnpermuteWithEpGetWorkspaceSize( const aclTensor *permutedTokens, const aclTensor *sortedIndices, const aclTensor *probsOptional, int64_t numTopk, const aclIntArray *rangeOptional, bool paddedMode, const aclIntArray *restoreShapeOptional, const aclTensor *out, uint64_t *workspaceSize, aclOpExecutor **executor); aclnnStatus aclnnMoeTokenUnpermuteWithEp( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream);第一段接口参数细节(GetWorkspaceSize)
| 参数名 | 输入/输出 | 描述与使用说明 | 数据类型 | 数据格式 | 维度(shape) | 非连续 Tensor |
|---|---|---|---|---|---|---|
| permutedTokens | 输入 | 经过扩展并排序过的 tokens;支持 2D,不支持空 Tensor;rangeOptional 非空时第 0 维长度不小于 rangeOptional[1]-rangeOptional[0] | BFLOAT16、FLOAT16、FLOAT32 | ND | (num_permuted_tokens, hidden_size) | √ |
| sortedIndices | 输入 | 需要计算的数据在 permutedTokens 中的位置;1D,长度为 num_tokens*numTopk,不支持空 Tensor;元素值不在 rangeOptional 范围内时对应输出贡献为 0 | INT32 | ND | (num_tokens*numTopk) | √ |
| probsOptional | 可选输入 | 输入 tokens 对应的专家概率;2D,第 1 维长度必须等于 numTopk;传非空合法 Tensor 时做乘法,传空指针时不乘 | BFLOAT16、FLOAT16、FLOAT32 | ND | (num_tokens, numTopk) | √ |
| numTopk | 输入 | 被选中的专家个数;必须 ≥1,probsOptional 非空时必须 ≤512 | INT64 | - | - | - |
| rangeOptional | 可选输入 | ep 切分的有效范围,size 为 2;允许传空指针,空指针时默认 {0,0},输出全 0 | - | - | - | - |
| paddedMode | 输入 | true 开启 paddedMode,false 关闭;目前仅支持 false | bool | - | - | - |
| restoreShapeOptional | 可选输入 | 预留参数,当前仅支持传入空指针 | - | - | - | - |
| out | 输出 | 反重排的输出结果;2D,不支持空 Tensor | 与 permutedTokens 一致 | ND | (num_tokens, hidden_size) | × |
| workspaceSize | 输出 | 需要在 Device 侧申请的 workspace 大小 | - | - | - | - |
| executor | 输出 | op 执行器,包含算子计算流程 | - | - | - | - |
返回值与错误码
第一段接口完成入参校验,返回aclnnStatus(完整返回码说明参见 aclnn 返回码):
| 返回值 | 错误码 | 描述 |
|---|---|---|
| ACLNN_ERR_PARAM_NULLPTR | 161001 | permutedTokens、sortedIndices、out 或 executor 为空指针 |
| ACLNN_ERR_PARAM_INVALID | 161002 | 输入和输出的数据类型和数据格式不在支持的范围之内 |
第二段接口参数细节
| 参数名 | 输入/输出 | 描述 |
|---|---|---|
| workspace | 输入 | 在 Device 侧申请的 workspace 内存地址 |
| workspaceSize | 输入 | Device 侧申请的 workspace 大小,由第一段接口获取 |
| executor | 输入 | op 执行器,包含算子计算流程 |
| stream | 输入 | 指定执行任务的 Stream |
实战:完整调用示例
以下示例来自仓库 examples 目录,展示了从资源初始化、Tensor 构造到两段式调用的完整流程。编译与运行方法参考 编译与运行样例。
示例数据解读
示例构造了 4 个 permuted token(shape 为 {4, 2})、6 个排序下标、3 个 token 各 2 个专家的概率:
- permutedTokensData = {2, 2, 1, 1, 3, 3, 2, 2},shape = {4, 2};
- sortedIndicesData = {2, 0, 4, 1, 5, 3},shape = {6},即 num_tokens=3、numTopk=2;
- probsOptionalData = {1, 1, 1, 1, 1, 1},shape = {3, 2};
- numTopk = 2;rangeOptional = {1, 5}(有效索引区间为 [1,5));
- 输出 out shape = {3, 2}。
根据公式,out[i] = Σ_j permutedTokens[sortedIndices[i*2+j]] * probs[i][j],且仅当下标落在 [1, 5) 内才有贡献。逐项计算:out[0] = t[2]*1 + t[0]*1,但 t[0] 的下标 0 不在 [1,5) 内贡献为 0,故 out[0]=3;out[1] = t[4]*1 + t[1]*1 = 3+1=4;out[2] = t[5]*1 + t[3]*1 = 3+1=4(下标 5、3 均在有效范围内)。读者可自行运行示例验证输出。
核心代码框架
#include "acl/acl.h" #include "aclnnop/aclnn_moe_token_unpermute_with_ep.h" #include <iostream> #include <vector> // ... CHECK_RET / LOG_PRINT / GetShapeSize / PrintOutResult 等工具宏与函数 ... int Init(int32_t deviceId, aclrtStream *stream) { // 固定写法,资源初始化 auto ret = aclInit(nullptr); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclInit failed. ERROR: %d\n", ret); return ret); ret = aclrtSetDevice(deviceId); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSetDevice failed. ERROR: %d\n", ret); return ret); ret = aclrtCreateStream(stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtCreateStream failed. ERROR: %d\n", ret); return ret); return 0; } template <typename T> int CreateAclIntArray(const std::vector<T>& hostData, void** deviceAddr, aclIntArray** intArray) { auto size = GetShapeSize(hostData) * sizeof(T); auto ret = aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMalloc failed. ERROR: %d\n", ret); return ret); ret = aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMemcpy failed. ERROR: %d\n", ret); return ret); *intArray = aclCreateIntArray(hostData.data(), hostData.size()); return 0; } template <typename T> int CreateAclTensor(const std::vector<T> &hostData, const std::vector<int64_t> &shape, void **deviceAddr, aclDataType dataType, aclTensor **tensor) { auto size = GetShapeSize(shape) * sizeof(T); auto ret = aclrtMalloc(deviceAddr, size, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMalloc failed. ERROR: %d\n", ret); return ret); ret = aclrtMemcpy(*deviceAddr, size, hostData.data(), size, ACL_MEMCPY_HOST_TO_DEVICE); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtMemcpy failed. ERROR: %d\n", ret); return ret); // 计算连续 tensor 的 strides std::vector<int64_t> strides(shape.size(), 1); for (int64_t i = shape.size() - 2; i >= 0; i--) { strides[i] = shape[i + 1] * strides[i + 1]; } *tensor = aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND, shape.data(), shape.size(), *deviceAddr); return 0; } int main() { // 1. device/stream 初始化 int32_t deviceId = 0; aclrtStream stream; auto ret = Init(deviceId, &stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("Init acl failed. ERROR: %d\n", ret); return ret); // 2. 构造输入与输出 std::vector<float> permutedTokensData = {2, 2, 1, 1, 3, 3, 2, 2}; std::vector<int64_t> permutedTokensShape = {4, 2}; void *permutedTokensAddr = nullptr; aclTensor *permutedTokens = nullptr; ret = CreateAclTensor(permutedTokensData, permutedTokensShape, &permutedTokensAddr, aclDataType::ACL_FLOAT, &permutedTokens); CHECK_RET(ret == ACL_SUCCESS, return ret); std::vector<int> sortedIndicesData = {2, 0, 4, 1, 5, 3}; std::vector<int64_t> sortedIndicesShape = {6}; void *sortedIndicesAddr = nullptr; aclTensor *sortedIndices = nullptr; ret = CreateAclTensor(sortedIndicesData, sortedIndicesShape, &sortedIndicesAddr, aclDataType::ACL_INT32, &sortedIndices); CHECK_RET(ret == ACL_SUCCESS, return ret); std::vector<float> probsOptionalData = {1, 1, 1, 1, 1, 1}; std::vector<int64_t> probsOptionalShape = {3, 2}; void *probsOptionalAddr = nullptr; aclTensor *probsOptional = nullptr; ret = CreateAclTensor(probsOptionalData, probsOptionalShape, &probsOptionalAddr, aclDataType::ACL_FLOAT, &probsOptional); CHECK_RET(ret == ACL_SUCCESS, return ret); int64_t num_topk = 2; void* rangeDeviceAddr = nullptr; aclIntArray* range = nullptr; std::vector<int64_t> rangeHostData = {1, 5}; ret = CreateAclIntArray(rangeHostData, &rangeDeviceAddr, &range); CHECK_RET(ret == ACL_SUCCESS, return ret); std::vector<float> outData = {0, 0, 0, 0, 0, 0}; std::vector<int64_t> outShape = {3, 2}; void *outAddr = nullptr; aclTensor *out = nullptr; ret = CreateAclTensor(outData, outShape, &outAddr, aclDataType::ACL_FLOAT, &out); CHECK_RET(ret == ACL_SUCCESS, return ret); // 3. 两段式调用 uint64_t workspaceSize = 0; aclOpExecutor *executor; ret = aclnnMoeTokenUnpermuteWithEpGetWorkspaceSize(permutedTokens, sortedIndices, probsOptional, num_topk, range, false, nullptr, out, &workspaceSize, &executor); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("GetWorkspaceSize failed. ERROR: %d\n", ret); return ret); void *workspaceAddr = nullptr; if (workspaceSize > 0) { ret = aclrtMalloc(&workspaceAddr, workspaceSize, ACL_MEM_MALLOC_HUGE_FIRST); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("allocate workspace failed. ERROR: %d\n", ret); return ret); } ret = aclnnMoeTokenUnpermuteWithEp(workspaceAddr, workspaceSize, executor, stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnMoeTokenUnpermuteWithEp failed. ERROR: %d\n", ret); return ret); // 4. 同步等待任务执行结束 ret = aclrtSynchronizeStream(stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclrtSynchronizeStream failed. ERROR: %d\n", ret); return ret); // 5. 将结果从 device 拷贝至 host 并打印 PrintOutResult(outShape, &outAddr); // 6. 释放 aclTensor aclDestroyTensor(permutedTokens); aclDestroyTensor(sortedIndices); aclDestroyTensor(probsOptional); aclDestroyTensor(out); // 7. 释放 device 资源 aclrtFree(permutedTokensAddr); aclrtFree(sortedIndicesAddr); aclrtFree(probsOptionalAddr); aclrtFree(outAddr); aclrtFree(rangeDeviceAddr); if (workspaceSize > 0) { aclrtFree(workspaceAddr); } aclrtDestroyStream(stream); aclrtResetDevice(deviceId); aclFinalize(); return 0; }调用步骤要点
- 初始化:
aclInit→aclrtSetDevice→aclrtCreateStream,固定写法; - 构造 Tensor:通过
aclrtMalloc/aclrtMemcpy准备 device 侧数据,再用aclCreateTensor封装为 aclTensor;注意rangeOptional这类aclIntArray需用aclCreateIntArray构造,示例中 range 为 {1, 5}; - 第一段接口:
aclnnMoeTokenUnpermuteWithEpGetWorkspaceSize完成入参校验并返回 workspaceSize 与 executor; - 申请 workspace:按返回值通过
aclrtMalloc申请(workspaceSize 为 0 时可不申请); - 第二段接口:
aclnnMoeTokenUnpermuteWithEp在指定 stream 上执行; - 同步与取数:
aclrtSynchronizeStream后从 device 拷贝结果到 host; - 资源释放:依次释放 aclTensor、device 内存、stream,并
aclrtResetDevice、aclFinalize。
底层实现原理:从 Tiling 到 Kernel
Shape 推导
infershape 源码 明确了输出 shape 的推导规则:输出恒为 2D,out_shape = (tokens_num, hidden_size)。其中 tokens_num 在 probs 非空时取 probs 第 0 维;probs 为空时由sortedIndices 第 0 维 / numTopk计算得到。输出数据类型与 permutedTokens 保持一致。
Tiling 策略
tiling 源码 展示了该算子如何在多核 AIV 上做任务划分,核心步骤为:
- 参数校验与初始化:校验 numTopk ≥ 1、permutedTokens 为 2D 且非空、sortedIndices 非空;probs 非空时校验其第 1 维等于 numTopk、第 0 维 × numTopk 等于 sortedIndices 长度、numTopk ≤ 512;读取 range 属性(空则 start=end=0)。
- 核数分配:
usedCoreNum = min(tokensNum, maxCoreNum),即 token 数少于核数时只启用实际需要的核。 - hidden 维度切分:根据 UB(Unified Buffer)可用内存与数据类型大小,先计算单核能容纳的最大 hidden size(预留 5120 字节给 indices/probs,并对 512 对齐),若 hiddenSize 超限则按
length/num/remain三段式切分。 - token 维度切分:先按核均分 token(处理余数尾块),再根据剩余内存空间判断每个核一次能处理的 token×topK 组数,必要时二次切分。
- tilingKey 计算:0 表示 probs 为 None;1/2/3 分别表示 probs 类型为 FLOAT/FLOAT16/BF16。kernel 入口通过
TILING_KEY_IS(n)选择对应的模板实例化分支。 - workspace 申请:固定申请
16 * 1024 * 1024字节(16MB)的 workspace。
Kernel 计算流程
kernel 头文件 中的执行链为Process → CalMultiOutToken → CalSingleOutToken → CalPartOutToken → CopyTokenIn/CalFirstToken/CalToken/CopyOut:
- 每个 AIV 核按 tiling 参数处理自己负责的 token 段,先整体搬入对应的 sortedIndices 段(及 probs 段,非 float 类型会先 Cast 成 float 参与计算);
- 对每个输出 token,先校验
acl_token_idx ∈ [start, end)(有 probs 时还要求prob_value != 0),满足才搬入该 token 的 hidden 切片;不满足则用Duplicate直接填充 0; - 首 token 作为累加初值(CalFirstToken),后续 numTopk-1 个 token 通过 CalToken 做乘加累加;
- 一个 token 的所有专家贡献累加完成后 CopyOut 写回输出,hidden 维度按切分循环处理直至完整。
这种"先 gather、再逐 token 乘加、按 topk 分组归约"的实现与 README 中的公式完全一一对应,也解释了为什么 sortedIndices 越界或 prob 为 0 时"对应输出贡献为 0"——kernel 在数据搬入前就通过条件判断跳过了这些位置。
测试验证
仓库为该算子提供了完整的测试配套,可用于验证行为正确性:
- 单测(UT):host 端 tiling/infershape 单测(如 test_moe_token_unpermute_with_ep_tiling.cpp、test_moeTokenUnpermuteWithEp_infershape.cpp)与 kernel 单测;
- 系统测试(ST):ST 用例配置 及对应的 ATK 执行脚本。
总结
MoeTokenUnpermuteWithEp 是 CANN ops-transformer 中 MoE 专家并行流水线的关键收尾算子,通过 sortedIndices 完成 gather 式反重排,按 numTopk 分组加权求和还原原始 token 输出。其设计要点可归纳为:rangeOptional控制 EP 切分后的有效区间(空指针时输出全 0 的约定需特别注意);probsOptional为空时退化为纯累加;两段式 aclnn 接口配合 workspace 机制保证 host/device 内存管理清晰可控;tiling 在 hidden 与 token 两个维度上同时切分以适配多核 AIV 与 UB 容量。对从事 MoE 推理/训练框架开发、需要在 NPU 上实现或移植专家并行 unpermute 逻辑的工程师,本文给出的参数语义、示例代码与底层机制可直接作为接入参考。
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考