CANN ops-transformer MoeTokenUnpermuteWithEp 算子解析:EP 场景下 MoE Token 反重排与加权聚合原理与实战
2026/9/20 3:33:38 网站建设 项目流程

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_tokenssorted_indices,可选输入probs,输出unpermuted_tokens,属性num_topkrangepadded_moderestore_shape)。对外暴露的语义如下:

参数名输入/输出/属性描述数据类型数据格式
permutedTokens输入表示经过扩展并排序过的 tokens,对应公式中的permutedTokensFLOAT16、BFLOAT16、FLOAT32ND
sortedIndices输入表示需要计算的数据在 permutedTokens 中的位置,对应公式中的sortedIndicesINT32ND
probsOptional可选输入表示输入 tokens 对应的专家概率,对应公式中的probsFLOAT16、BFLOAT16、FLOAT32ND
numTopk属性被选中的专家个数INT64-
rangeOptional属性ep 切分的有效范围,size 为 2aclIntArray*-
paddedMode属性目前仅支持 falseBOOL-
restoreShapeOptional属性目前仅支持 nullptraclIntArray*-
out输出表示 permutedTokens 反重排的输出结果,对应公式中的outFLOAT16、BFLOAT16、FLOAT32ND

注意: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、FLOAT32ND(num_permuted_tokens, hidden_size)
sortedIndices输入需要计算的数据在 permutedTokens 中的位置;1D,长度为 num_tokens*numTopk,不支持空 Tensor;元素值不在 rangeOptional 范围内时对应输出贡献为 0INT32ND(num_tokens*numTopk)
probsOptional可选输入输入 tokens 对应的专家概率;2D,第 1 维长度必须等于 numTopk;传非空合法 Tensor 时做乘法,传空指针时不乘BFLOAT16、FLOAT16、FLOAT32ND(num_tokens, numTopk)
numTopk输入被选中的专家个数;必须 ≥1,probsOptional 非空时必须 ≤512INT64---
rangeOptional可选输入ep 切分的有效范围,size 为 2;允许传空指针,空指针时默认 {0,0},输出全 0----
paddedMode输入true 开启 paddedMode,false 关闭;目前仅支持 falsebool---
restoreShapeOptional可选输入预留参数,当前仅支持传入空指针----
out输出反重排的输出结果;2D,不支持空 Tensor与 permutedTokens 一致ND(num_tokens, hidden_size)×
workspaceSize输出需要在 Device 侧申请的 workspace 大小----
executor输出op 执行器,包含算子计算流程----

返回值与错误码

第一段接口完成入参校验,返回aclnnStatus(完整返回码说明参见 aclnn 返回码):

返回值错误码描述
ACLNN_ERR_PARAM_NULLPTR161001permutedTokens、sortedIndices、out 或 executor 为空指针
ACLNN_ERR_PARAM_INVALID161002输入和输出的数据类型和数据格式不在支持的范围之内

第二段接口参数细节

参数名输入/输出描述
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; }

调用步骤要点

  1. 初始化aclInitaclrtSetDeviceaclrtCreateStream,固定写法;
  2. 构造 Tensor:通过aclrtMalloc/aclrtMemcpy准备 device 侧数据,再用aclCreateTensor封装为 aclTensor;注意rangeOptional这类aclIntArray需用aclCreateIntArray构造,示例中 range 为 {1, 5};
  3. 第一段接口aclnnMoeTokenUnpermuteWithEpGetWorkspaceSize完成入参校验并返回 workspaceSize 与 executor;
  4. 申请 workspace:按返回值通过aclrtMalloc申请(workspaceSize 为 0 时可不申请);
  5. 第二段接口aclnnMoeTokenUnpermuteWithEp在指定 stream 上执行;
  6. 同步与取数aclrtSynchronizeStream后从 device 拷贝结果到 host;
  7. 资源释放:依次释放 aclTensor、device 内存、stream,并aclrtResetDeviceaclFinalize

底层实现原理:从 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 上做任务划分,核心步骤为:

  1. 参数校验与初始化:校验 numTopk ≥ 1、permutedTokens 为 2D 且非空、sortedIndices 非空;probs 非空时校验其第 1 维等于 numTopk、第 0 维 × numTopk 等于 sortedIndices 长度、numTopk ≤ 512;读取 range 属性(空则 start=end=0)。
  2. 核数分配usedCoreNum = min(tokensNum, maxCoreNum),即 token 数少于核数时只启用实际需要的核。
  3. hidden 维度切分:根据 UB(Unified Buffer)可用内存与数据类型大小,先计算单核能容纳的最大 hidden size(预留 5120 字节给 indices/probs,并对 512 对齐),若 hiddenSize 超限则按length/num/remain三段式切分。
  4. token 维度切分:先按核均分 token(处理余数尾块),再根据剩余内存空间判断每个核一次能处理的 token×topK 组数,必要时二次切分。
  5. tilingKey 计算:0 表示 probs 为 None;1/2/3 分别表示 probs 类型为 FLOAT/FLOAT16/BF16。kernel 入口通过TILING_KEY_IS(n)选择对应的模板实例化分支。
  6. 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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询