news 2026/9/19 20:37:26

CANN ops-transformer 算子深度解析:MoeTokenUnpermuteWithRoutingMapGrad 反向传播原理与 aclnn 接口实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
CANN ops-transformer 算子深度解析:MoeTokenUnpermuteWithRoutingMapGrad 反向传播原理与 aclnn 接口实战

CANN ops-transformer 算子深度解析:MoeTokenUnpermuteWithRoutingMapGrad 反向传播原理与 aclnn 接口实战

【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer

本文是 CANN / ops-transformer 开源算子库中 MoE(Mixture of Experts)稀疏专家路由反向算子MoeTokenUnpermuteWithRoutingMapGrad的完整技术指南,围绕其 aclnn 双段式接口的数学原理、参数约束、错误码与可编译的 C++ 调用示例展开,并下沉到 op_host 注册与 tiling、op_kernel 内核实现的源码级原理。读完本文,你将能够在 Ascend 平台上正确构造输入输出并调用aclnnMoeTokenUnpermuteWithRoutingMapGrad完成 MoE Token 反重排的梯度回传,同时理解 dropAndPad 两种模式、probs 可选输入与混合精度组合背后的实现机制。

算子定位:MoE 稀疏路由中的 Token 反重排梯度算子

在 MoE 架构的 Transformer 大模型中,输入 token 会先通过路由(Routing)机制被**重排(Permute)并分发给对应专家处理,专家计算完成后需要将结果反重排(Unpermute)回原始 token 顺序。CANN ops-transformer 提供了配套的 MoeTokenUnpermuteWithRoutingMap 正向算子 完成这一数据搬运,而本文主角MoeTokenUnpermuteWithRoutingMapGrad则是该算子的反向传播(Backward)**接口:它把正向输出unpermutedTokens的梯度,依据outIndexpermuteTokenId索引反推回输入permutedTokens的梯度,并在存在probs(路由权重/概率)时进一步计算出probs的梯度。

根据 aclnnMoeTokenUnpermuteWithRoutingMapGrad.md 与 模块 README,该算子支持以下产品:

产品是否支持
Ascend 950PR/Ascend 950DT
Atlas A3 训练系列产品/Atlas A3 推理系列产品
Atlas A2 训练系列产品/Atlas A2 推理系列产品
Atlas 200I/500 A2 推理产品×
Atlas 推理系列产品×
Atlas 训练系列产品×

算子注册源码 moe_token_unpermute_with_routing_map_grad_def.cpp 中可看到对应的 AICore 平台配置:ascend910bascend910_93ascend950三种平台均注册了该算子,并统一开启了DynamicCompileStaticFlagDynamicRankSupportFlagDynamicShapeSupportFlag,说明该算子支持动态 shape 与动态 rank。

数学原理与计算公式

正向算子将permutedTokens按索引累加回unpermutedTokens时,若存在probs还会先做加权。反向算子需要精确还原这两条链路的梯度。文档给出了如下计算规则:

(1)probs 非 None 时

首先按索引完成unpermutedTokensGradpermutedTokensGrad的基础散射(Scatter),并计算permutedProbsGrad

$$ permutedTokensGrad[outIndex[i]] = unpermutedTokensGrad[permuteTokenId[i]] $$

$$ permutedProbsGrad = permutedTokensGrad * permutedTokensOptional $$

$$ probsGradExpertOrder = \sum_{j=0}^{hidden_size}(permutedProbsGrad_{i,j}) $$

其中hidden_sizeunpermutedTokensGrad的第 1 维大小。

dropAndPad 为 false(每个 token 可被不超过 topK_num 个专家处理)时:

$$ probsGradOut = masked_scatter(routingMapOptional^T, probsGradExpertOrder) $$

$$ permutedProbs = probsOptional^T.masked_select(routingMapOptional^T) $$

$$ permutedTokensGradOut = permutedProbs.unsqueeze(-1) * permutedTokensGrad $$

dropAndPad 为 true(每个专家固定处理 capacity 个 token)时:

$$ probsGradOut[permuteTokenId[i], outIndex[i]/capacity] = probsGradExpertOrder[outIndex[i]] $$

$$ permutedProbs[outIndex[i]] = probsOptional.view(1)[i] $$

$$ permutedTokensGradOut = permutedProbs * permutedTokensGrad $$

(2)probs 为 None 时

此时退化为纯索引散射:

$$ permutedTokensGradOut[outIndex[i]] = unpermutedTokensGrad[permuteTokenId[i]] $$

关键维度推导

  • hidden_sizeunpermutedTokensGrad的第 1 维大小(词向量维度)。
  • dropAndPad == true时,每个专家固定能够处理capacity个 token。输入routingMapOptional的第 1 维是experts_num(专家个数),输入outIndex的第 0 维是experts_num * capacity,据此可以算出capacity
  • dropAndPad == false时,每个 token 能被小于等于topK_num个专家处理。输入unpermutedTokensGrad的第 0 维是tokens_num(token 个数),输入outIndex的第 0 维是tokens_num * topK_num,据此可以算出topK_num

在正向算子文档中,topK_num = permutedTokens.size(0) // routingMapOptional.size(0),未使用的槽位在sortedIndices中以-1表示并在计算时跳过;反向算子同样遵循这一约定,源码中通过permuteTokenId < 0判断跳过无效槽位。

函数原型与两段式接口

该算子遵循 CANN aclnn 标准的 两段式接口:必须先调用第一段接口aclnnMoeTokenUnpermuteWithRoutingMapGradGetWorkspaceSize获取计算所需 workspace 大小以及包含算子计算流程的执行器,再调用第二段接口aclnnMoeTokenUnpermuteWithRoutingMapGrad执行计算。

aclnnStatus aclnnMoeTokenUnpermuteWithRoutingMapGradGetWorkspaceSize( const aclTensor* unpermutedTokensGrad, const aclTensor* outIndex, const aclTensor* permuteTokenId, const aclTensor* routingMapOptional, const aclTensor* permutedTokensOptional, const aclTensor* probsOptional, bool dropAndPad, const aclIntArray* restoreShapeOptional, const aclTensor* permutedTokensGradOut, const aclTensor* probsGradOutOptional, uint64_t* workspaceSize, aclOpExecutor** executor)
aclnnStatus aclnnMoeTokenUnpermuteWithRoutingMapGrad( void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, aclrtStream stream)

aclnnMoeTokenUnpermuteWithRoutingMapGradGetWorkspaceSize 参数说明

第一段接口完成入参校验与 workspace 计算,各参数含义如下:

参数名输入/输出描述使用说明数据类型数据格式维度(shape)非连续 Tensor
unpermutedTokensGrad输入计算公式中的 unpermutedTokensGrad,代表正向输出 unpermutedTokens 的梯度-BFLOAT16、FLOAT16、FLOATND(tokens_num,hidden_size)
outIndex输入计算公式中的 outIndex,代表输出位置索引dropAndPad 为 false 时取值范围 [0, tokens_num*topK_num-1];为 true 时取值范围 [0, experts_num*capacity-1]INT32NDdropAndPad 为 false 时 (tokens_num*topK_num);为 true 时 (experts_num*capacity)
permuteTokenId输入计算公式中的 permuteTokenId,代表输入 permutedTokens 每个位置对应的 Token 序号取值范围 [0, tokens_num-1]INT32ND与 outIndex 相同
routingMapOptional可选输入当输入 probsOptional 为空指针时不需要此输入,应传入空指针。代表对应位置的 Token 是否被对应专家处理INT8 类型取值支持 0、1;BOOL 类型取值支持 true、falseINT8、BOOLND(tokens_num, experts_num)
permutedTokensOptional可选输入当输入 probsOptional 为空指针时不需要此输入,应传入空指针数据类型与 unpermutedTokensGrad 相同BFLOAT16、FLOAT16、FLOATNDdropAndPad 为 false 时 (tokens_num*topK_num, hidden_size);为 true 时 (experts_num*capacity, hidden_size)
probsOptional可选输入当不需要时为空指针数据类型与 unpermutedTokensGrad 相同;或者当 unpermutedTokensGrad 是 BFLOAT16 时 probsOptional 支持 FLOATBFLOAT16、FLOAT16、FLOATND与 routingMapOptional 相同
dropAndPad属性true 表示开启 dropAndPad,false 表示关闭 dropAndPad-BOOL---
restoreShapeOptional属性INT64 类型的 aclIntArray。dropAndPad 为 true 时代表 unpermutedTokensGrad 的 shape-INT64---
permutedTokensGradOut输出计算公式中的 permutedTokensGradOut,代表输入 permutedTokens 的梯度数据类型与 unpermutedTokensGrad 相同BFLOAT16、FLOAT16、FLOATNDdropAndPad 为 false 时 (tokens_num*topK_num, hidden_size);为 true 时 (experts_num*capacity, hidden_size)×
probsGradOutOptional可选输出未输入 probsOptional 时为空指针。输入 probs 的梯度数据类型与 probsOptional 相同BFLOAT16、FLOAT16、FLOATND与 routingMapOptional 相同×
workspaceSize输出返回需要在 Device 侧申请的 workspace 大小-----
executor输出返回 op 执行器,包含算子计算流程-----

注意:输入侧(unpermutedTokensGrad、outIndex、permuteTokenId、routingMapOptional、permutedTokensOptional、probsOptional)均支持非连续 Tensor(标记为 √),而两个输出permutedTokensGradOutprobsGradOutOptional要求连续(标记为 ×)。

返回值与错误码

两段接口均返回aclnnStatus状态码,具体参见 aclnn 返回码。第一段接口完成入参校验,出现以下场景时报错:

返回值错误码描述
ACLNN_ERR_PARAM_NULLPTR161001传入的必选输入、必选输出或必选属性是空指针
ACLNN_ERR_PARAM_INVALID161002输入和输出的数据类型和数据格式不在支持范围之内
ACLNN_ERR_INNER_TILING_ERROR561002输入 probsOptional 非空,且 dropAndPad 为 false 时,topK_num > 512
ACLNN_ERR_INNER_TILING_ERROR561002输入 probsOptional 非空,且 dropAndPad 为 false 时,topK_num 大于 experts_num
ACLNN_ERR_INNER_TILING_ERROR561002输入 probsOptional 非空,且 dropAndPad 为 false 时,(ubSize - (probTypeLen + 1) * numExpertAlign - (tokenTypeLen + 8) * 256) / (6 * tokenTypeLen + 12) < 1
ACLNN_ERR_INNER_TILING_ERROR561002输入 probsOptional 非空,且 dropAndPad 为 true 时,capacity 大于 tokens_num
ACLNN_ERR_INNER_TILING_ERROR561002输入 probsOptional 非空,且 dropAndPad 为 true 时,hidden_size > 256 * (ubSize - 2080) / (8 + tokenTypeLen)
ACLNN_ERR_INNER_TILING_ERROR561002输入 probsOptional 非空时,输入 routingMapOptional 或 permutedTokensOptional 为空
ACLNN_ERR_INNER_TILING_ERROR561002输入 probsOptional 非空时,probsOptional 数据类型与 unpermutedTokensGrad 不同且 unpermutedTokensGrad 不是 BFLOAT16
ACLNN_ERR_INNER_TILING_ERROR561002输入或输出的 shape 不符合要求

这些约束在 moe_token_unpermute_with_routing_map_grad_tiling.cpp 中均有对应的OP_CHECK_IF校验逻辑(例如topK > MAX_TOP_Kcapacity > tokensNum等)。

aclnnMoeTokenUnpermuteWithRoutingMapGrad 参数说明

第二段接口参数如下:

参数名输入/输出描述
workspace输入在 Device 侧申请的 workspace 内存地址
workspaceSize输入在 Device 侧申请的 workspace 大小,由第一段接口 aclnnMoeTokenUnpermuteWithRoutingMapGradGetWorkspaceSize 获取
executor输入op 执行器,包含算子计算流程
stream输入指定执行任务的 Stream

约束说明

  • 确定性计算aclnnMoeTokenUnpermuteWithRoutingMapGrad默认确定性实现。
  • tokens_num表示输入的 token 数量,hidden_size表示词向量维度,experts_num表示专家个数。
  • 通过dropAndPad区分两种模式:dropAndPad == true时,每个专家固定能处理capacity个 token;dropAndPad == false时,每个 token 能被小于等于topK_num个专家处理。
  • 当输入probsOptional非空,且dropAndPad为 false 时:
    • 要求topK_num <= 512topK_num <= experts_num
    • 要求experts_num满足(ubSize - (probTypeLen + 1) * numExpertAlign - (tokenTypeLen + 8) * 256) / (6 * tokenTypeLen + 12) >= 1,其中ubSize是芯片 ub 空间大小,probTypeLen是输入probsOptional的数据类型对应字节数,tokenTypeLen是输入unpermutedTokensGrad的数据类型对应字节数,numExpertAlignexperts_num对 32 做向上对齐的结果。
  • 当输入probsOptional非空,且dropAndPad为 true 时:
    • 要求capacity <= tokens_num
    • 要求hidden_size <= 256 * (ubSize - 2080) / (8 + tokenTypeLen),其中ubSize是芯片 ub 空间大小,tokenTypeLen是输入unpermutedTokensGrad的数据类型对应字节数。

从源码看,MAX_TOP_K = 512INDICES_RESERVE_MAX_NUM = 256等常量定义在 tiling.cpp 中,与上述约束一一对应;同时 moe_token_unpermute_with_routing_map_grad_base.h 中定义了BLOCK_SIZE_512 = 512FP32_ONE_REPEAT = 64INDICES_PROBS_MAX_RESERVE_NUM = 512等内核侧常量。

完整调用示例

示例代码如下(来源:aclnnMoeTokenUnpermuteWithRoutingMapGrad.md 及 examples/test_aclnn_moe_token_unpermute_with_routing_map_grad.cpp),具体编译与执行过程请参考 编译与运行样例:

#include <iostream> #include <vector> #include "acl/acl.h" #include "aclnnop/aclnn_moe_token_unpermute_with_routing_map_grad.h" #include <iostream> #define CHECK_RET(cond, return_expr) \ do { \ if (!(cond)) { \ return_expr; \ } \ } while (0) #define LOG_PRINT(message, ...) \ do { \ printf(message, ##__VA_ARGS__); \ } while (0) int64_t GetShapeSize(const std::vector<int64_t> &shape) { int64_t shapeSize = 1; for (auto i : shape) { shapeSize *= i; } return shapeSize; } void PrintOutResult(std::vector<int64_t> &shape, void **deviceAddr) { auto size = GetShapeSize(shape); std::vector<float> resultData(size, 0); auto ret = aclrtMemcpy( resultData.data(), resultData.size() * sizeof(resultData[0]), *deviceAddr, size * sizeof(resultData[0]), ACL_MEMCPY_DEVICE_TO_HOST); CHECK_RET( ret == ACL_SUCCESS, LOG_PRINT("copy result from device to host failed. ERROR: %d\n", ret); return ); for (int64_t i = 0; i < size; i++) { LOG_PRINT("result[%ld] is: %f\n", i, resultData[i]); } } 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 CreateAclTensor(const std::vector<T> &hostData, const std::vector<int64_t> &shape, void **deviceAddr, aclDataType dataType, aclTensor **tensor) { auto size = GetShapeSize(shape) * sizeof(T); // 调用aclrtMalloc申请device侧内存 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); // 调用aclrtMemcpy将host侧数据拷贝到device侧内存上 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]; } // 调用aclCreateTensor接口创建aclTensor *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初始化,参考acl API手册 // 根据自己的实际device填写deviceId 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. 构造输入与输出,需要根据API的接口自定义构造 bool dropAndPad = false; int32_t tokenNum = 1; int32_t hiddenSize = 2; int32_t expertNum = 2; int32_t topK = 2; int32_t outTokenNum = tokenNum * topK; std::vector<int64_t> permutedTokensShape = {outTokenNum, hiddenSize}; std::vector<int64_t> unpermutedTokensGradShape = {tokenNum, hiddenSize}; std::vector<int64_t> probsShape = {tokenNum, expertNum}; std::vector<int64_t> outIndexShape = {outTokenNum}; std::vector<int64_t> permuteTokenIdShape = {outTokenNum}; std::vector<int64_t> routingMapShape = {tokenNum, expertNum}; std::vector<int64_t> permutedTokensGradShape = {outTokenNum, hiddenSize}; std::vector<int64_t> probsGradShape = {tokenNum, expertNum}; void* permutedTokensDeviceAddr = nullptr; void* unpermutedTokensGradDeviceAddr = nullptr; void* probsDeviceAddr = nullptr; void* outIndexDeviceAddr = nullptr; void* permuteTokenIdDeviceAddr = nullptr; void* routingMapDeviceAddr = nullptr; void* permutedTokensGradDeviceAddr = nullptr; void* probsGradDeviceAddr = nullptr; aclTensor* permutedTokens = nullptr; aclTensor* unpermutedTokensGrad = nullptr; aclTensor* probs = nullptr; aclTensor* outIndex = nullptr; aclTensor* permuteTokenId = nullptr; aclTensor* routingMap = nullptr; aclTensor *permutedTokensGrad = nullptr; aclTensor *probsGrad = nullptr; std::vector<float> permutedTokensHostData = {1, 1, 1, 1}; std::vector<float> unpermutedTokensGradHostData = {1, 1}; std::vector<float> probsHostData = {1, 1}; std::vector<int> outIndexHostData = {0, 1}; std::vector<int> permuteTokenIdHostData = {0, 0}; std::vector<int8_t> routingMapHostData = {1, 1}; std::vector<float> permutedTokensGradHostData = {0, 0, 0, 0}; std::vector<float> probsGradHostData = {0, 0}; ret = CreateAclTensor(unpermutedTokensGradHostData, unpermutedTokensGradShape, &unpermutedTokensGradDeviceAddr, aclDataType::ACL_FLOAT, &unpermutedTokensGrad); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(outIndexHostData, outIndexShape, &outIndexDeviceAddr, aclDataType::ACL_INT32, &outIndex); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(permuteTokenIdHostData, permuteTokenIdShape, &permuteTokenIdDeviceAddr, aclDataType::ACL_INT32, &permuteTokenId); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(routingMapHostData, routingMapShape, &routingMapDeviceAddr, aclDataType::ACL_BOOL, &routingMap); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(permutedTokensHostData, permutedTokensShape, &permutedTokensDeviceAddr, aclDataType::ACL_FLOAT, &permutedTokens); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(probsHostData, probsShape, &probsDeviceAddr, aclDataType::ACL_FLOAT, &probs); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(permutedTokensGradHostData, permutedTokensGradShape, &permutedTokensGradDeviceAddr, aclDataType::ACL_FLOAT, &permutedTokensGrad); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(probsGradHostData, probsGradShape, &probsGradDeviceAddr, aclDataType::ACL_FLOAT, &probsGrad); CHECK_RET(ret == ACL_SUCCESS, return ret); // 3. 调用CANN算子库API,需要修改为具体的Api名称 uint64_t workspaceSize = 0; aclOpExecutor *executor; // 调用aclnnMoeTokenUnpermuteWithRoutingMapGrad第一段接口 ret = aclnnMoeTokenUnpermuteWithRoutingMapGradGetWorkspaceSize(unpermutedTokensGrad, outIndex, permuteTokenId, routingMap, permutedTokens, probs, dropAndPad, nullptr, permutedTokensGrad, probsGrad, &workspaceSize, &executor); CHECK_RET( ret == ACL_SUCCESS, LOG_PRINT("aclnnMoeTokenUnpermuteWithRoutingMapGradGetWorkspaceSize failed. ERROR: %d\n", ret); return ret); // 根据第一段接口计算出的workspaceSize申请device内存 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); } // 调用aclnnMoeTokenUnpermuteWithRoutingMapGrad第二段接口 ret = aclnnMoeTokenUnpermuteWithRoutingMapGrad(workspaceAddr, workspaceSize, executor, stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnMoeTokenUnpermuteWithRoutingMapGrad 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侧,需要根据具体API的接口定义修改 LOG_PRINT("permutedTokensGrad \n"); PrintOutResult(permutedTokensGradShape, &permutedTokensGradDeviceAddr); LOG_PRINT("probsGrad \n"); PrintOutResult(probsGradShape, &probsGradDeviceAddr); // 6. 释放aclTensor和aclScalar,需要根据具体API的接口定义修改 aclDestroyTensor(permutedTokens); aclDestroyTensor(unpermutedTokensGrad); aclDestroyTensor(outIndex); aclDestroyTensor(permuteTokenId); aclDestroyTensor(routingMap); aclDestroyTensor(probs); aclDestroyTensor(permutedTokensGrad); aclDestroyTensor(probsGrad); // 7. 释放device资源 aclrtFree(permutedTokensDeviceAddr); aclrtFree(unpermutedTokensGradDeviceAddr); aclrtFree(probsDeviceAddr); aclrtFree(outIndexDeviceAddr); aclrtFree(permuteTokenIdDeviceAddr); aclrtFree(routingMapDeviceAddr); aclrtFree(permutedTokensGradDeviceAddr); aclrtFree(probsGradDeviceAddr); if (workspaceSize > 0) { aclrtFree(workspaceAddr); } aclrtDestroyStream(stream); aclrtResetDevice(deviceId); aclFinalize(); return 0; }

示例构造的是dropAndPad = falsetokenNum = 1hiddenSize = 2expertNum = 2topK = 2的最简场景:unpermutedTokensGrad形状为 (1, 2),outIndex/permuteTokenId形状为 (2),routingMapprobs形状为 (1, 2),输出permutedTokensGrad形状为 (2, 2)、probsGrad形状为 (1, 2)。实际使用中请根据 batch 的真实 token 数、topK 与专家数调整各 shape。

源码级实现解析

算子注册与平台配置(op_host)

moe_token_unpermute_with_routing_map_grad_def.cpp 使用OpDef注册了 6 个输入(unpermuted_tokens_gradout_indexpermute_token_id为 REQUIRED,routing_mappermuted_tokensprobs为 OPTIONAL)、2 个输出(permuted_tokens_grad为 REQUIRED,probs_grad为 OPTIONAL)以及 2 个属性(drop_and_pad默认 false、restore_shape默认空列表)。数据类型组合以 8 元组形式声明,覆盖 BF16/FLOAT16/FLOAT 与 INT32/BOOL/INT8 的合法搭配;其中probs支持 BF16/FLOAT16/FLOAT,并与unpermuted_tokens_grad的 BF16 组合出混合精度场景。

Tiling 策略(op_host)

tiling.cpp 实现了完整的动态 tiling 逻辑,按「probs 是否为 None × dropAndPad 是否为 true」拆分为三条路径:

  • TilingForProbIsNone:纯索引散射路径,按numOutTokens做核间均分,核内按hiddenSizeAlign切分 hidden 维循环搬入搬出;
  • TilingForProbNotNonePadTrue:校验capacity <= tokensNumhiddenSize上限,计算capacity = numOutTokens / numExpert
  • TilingForProbNotNonePadFalse:校验topK <= 512topK <= numExpert,并根据 UB 剩余空间(减去 indices、routingMap 对齐缓冲等)反推numExperthiddenSizeAlign的上限。

Tiling4MoeTokenUnpermuteWithRoutingMapGrad中通过ascendcPlatform.GetCoreNumAiv()获取核数并SetBlockDim,通过GetCoreMemSize(UB, ...)获取totalUbSize;tiling 结果包含tokensNum/topK/capacity/numExpert/hiddenSize/numOutTokens、核间切分信息(formerCoreNum/tailCoreNum/rowIdMapEachCore/rowIdMapTailCore)与核内切分信息(hiddenSizeAlign/hiddenSizeLoopTimes/hiddenSizeTail等),字段定义见 moe_token_unpermute_with_routing_map_grad_tiling.h。

此外 tiling 阶段会计算一个tilingKeytilingKey = mixKey * 100 + paddedModeKey * 10 + probKey,其中probKey表示 probs 是否存在(0/1)、paddedModeKey表示 dropAndPad(0/1)、mixKey表示 probs 是否与 tokens 混合精度(0/1),用于内核侧分支选择。

Kernel 实现(op_kernel)

内核入口 moe_token_unpermute_with_routing_map_grad.cpp 根据TILING_KEY分派到 6 个类模板实例:

TilingKey场景模板实例
0probs=None + dropAndPad=falseMoeTokenUnpermuteWithRoutingMapGradProbNoneDropPadFalse
10probs=None + dropAndPad=trueMoeTokenUnpermuteWithRoutingMapGradProbNoneDropPadTrue
1probs≠None + dropAndPad=false(同精度)MoeTokenUnpermuteWithRoutingMapGradProbNotNoneDropPadFalse
11probs≠None + dropAndPad=true(同精度)MoeTokenUnpermuteWithRoutingMapGradProbNotNoneDropPadTrue
101probs≠None + dropAndPad=false(probs 为 FLOAT 混合精度)MoeTokenUnpermuteWithRoutingMapGradProbNotNoneDropPadFalse
111probs≠None + dropAndPad=true(probs 为 FLOAT 混合精度)MoeTokenUnpermuteWithRoutingMapGradProbNotNoneDropPadTrue

ProbNotNoneDropPadFalse分支(moe_token_unpermute_with_routing_map_grad_prob_not_none_drop_pad_false.h)为例,其核心流程为:

  1. outIndex通过GetValue读入indicesArray,并依据routingMap逐 token 选择被选中的专家,将对应 probs 写入probsArray(未选中槽位补 0 且地址记为 -1);
  2. 按 hidden 维循环:将unpermutedTokensGrad一行搬入 UB,Cast/Copy 到 FP32 缓冲,用Muls乘以对应 probs 得到permutedTokensGrad行,写回 GM;同时用Mul计算permutedTokens * unpermutedTokensGrad并通过ReduceSumFunc沿 hidden 维归约得到该槽位的 probs 梯度;
  3. 所有 hidden 分块累加完成后,通过DataCopyPad按专家位置 scatter 回probsGrad

ReduceSumFunc定义在 moe_token_unpermute_with_routing_map_grad_base.h,按 hidden 长度分档:大于 4096 先二分累加至 8192 再用多级BlockReduceSum + WholeReduceSum归约到 1 个标量;小于等于 64 则直接WholeReduceSumBinaryAddFunc实现了二分组加法的并行累加。基类还提供了 MTE2/MTE3/V 单元之间的事件同步封装(SToMTE2SyncVToMTE3Sync等)与手动 ping-pong 双缓冲空间管理,用于隐藏搬运延迟。

ProbNoneDropPadFalse分支(moe_token_unpermute_with_routing_map_grad_prob_none_drop_pad_false.h)则退化为一个非常轻量的通路:对每个rowIdMap槽位读取permuteTokenId,若为 -1 直接跳过,否则将对应 token 的 hidden 行从unpermutedTokensGrad逐块拷贝到permutedTokensGrad,全程仅使用一个VECIN/VECOUT队列完成搬运。

测试用例佐证

仓库在 tests/st/aclnnMoeTokenUnpermuteWithRoutingMapGrad/atk_aclnnMoeTokenUnpermuteWithRoutingMapGrad.json 中提供了大量 ATK 用例,覆盖:

  • unpermutedTokensGrad的 FP32 / FP16 / BF16 三种 dtype,与probsOptional的 BF16+FP32 混合精度组合(is_mix: true);
  • routingMapOptional的 INT8 与 BOOL 两种取值类型;
  • padded_mode(即 dropAndPad)为 true / false 两种模式;
  • 从几十到数万量级的 token 数、hidden_size 从 1000 到 7000、专家数从十几到两百以上的多种 shape 组合。

同时 tests/ut/op_host/test_moe_token_unpermute_with_routing_map_grad_tiling.cpp 与 tests/ut/op_kernel/test_moe_token_unpermute_with_routing_map_grad.cpp 提供了 host 侧 tiling 与 kernel 侧的单元测试,可用于验证不同 shape 与 tilingKey 组合下的正确性。

平台编译配置

op_host/config 下按ascend910bascend910_93分别提供了 binary 配置与 simplified key 配置;moe_token_unpermute_with_routing_map_grad_simplified_key.ini 中default=0,指示 opc 工具以simplified_key_mode=0编译二进制 kernel。

进一步阅读

  • 模块主页与算子概述:moe_token_unpermute_with_routing_map_grad/README.md
  • 正向算子(MoeTokenUnpermuteWithRoutingMap):moe/moe_token_unpermute_with_routing_map/docs/aclnnMoeTokenUnpermuteWithRoutingMap.md
  • 两段式接口规范:docs/zh/context/two_phase_api.md
  • aclnn 返回码说明:docs/zh/context/aclnn_return_code.md
  • 样例编译与运行:docs/zh/context/compile_and_run_sample.md
  • 同系列 MoE 路由算子:可参考仓库 moe 目录 下的moe_token_permute_with_routing_mapmoe_token_unpermute_with_routing_map等配套实现,形成完整的「路由—重排—专家计算—反重排—梯度」闭环认知。

【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/19 20:36:55

AI生成PPT终于能编辑了?PPT Master 的3步上手路径

AI生成PPT终于能编辑了&#xff1f;PPT Master 的3步上手路径 【免费下载链接】ppt-master AI turns documents or topics into real, native PowerPoint decks—with native shapes, transitions and animations, data-backed charts and tables on demand, audio narration f…

作者头像 李华
网站建设 2026/9/19 20:35:40

WMS选型与落地指南:从库位建模到波次策略的系统化拆解

简介&#xff1a;这是一份由郎丰利于2023年整理完成的《WMS仓储系统解决方案》演示PPT&#xff0c;面向制造与物流企业的仓储管理人员、信息化规划及项目推进者&#xff0c;重点解决出入库流程混乱、库存数据失真、拣货补货效率低、盘点困难等典型痛点。内容围绕“物动帐动、可…

作者头像 李华
网站建设 2026/9/19 20:32:38

Miniconda安装配置与虚拟环境管理:Windows新手避坑指南

1. 为什么我劝你从Miniconda开始&#xff0c;而不是直接装Python很多人第一次接触Python&#xff0c;第一反应是去官网下载一个安装包&#xff0c;双击、下一步、完成&#xff0c;然后就开始写代码。这个流程本身没错&#xff0c;但只要你开始接触数据分析、机器学习或者需要同…

作者头像 李华