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的梯度,依据outIndex与permuteTokenId索引反推回输入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 平台配置:ascend910b、ascend910_93、ascend950三种平台均注册了该算子,并统一开启了DynamicCompileStaticFlag、DynamicRankSupportFlag、DynamicShapeSupportFlag,说明该算子支持动态 shape 与动态 rank。
数学原理与计算公式
正向算子将permutedTokens按索引累加回unpermutedTokens时,若存在probs还会先做加权。反向算子需要精确还原这两条链路的梯度。文档给出了如下计算规则:
(1)probs 非 None 时
首先按索引完成unpermutedTokensGrad到permutedTokensGrad的基础散射(Scatter),并计算permutedProbsGrad:
$$ permutedTokensGrad[outIndex[i]] = unpermutedTokensGrad[permuteTokenId[i]] $$
$$ permutedProbsGrad = permutedTokensGrad * permutedTokensOptional $$
$$ probsGradExpertOrder = \sum_{j=0}^{hidden_size}(permutedProbsGrad_{i,j}) $$
其中hidden_size指unpermutedTokensGrad的第 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_size指unpermutedTokensGrad的第 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、FLOAT | ND | (tokens_num,hidden_size) | √ |
| outIndex | 输入 | 计算公式中的 outIndex,代表输出位置索引 | dropAndPad 为 false 时取值范围 [0, tokens_num*topK_num-1];为 true 时取值范围 [0, experts_num*capacity-1] | INT32 | ND | dropAndPad 为 false 时 (tokens_num*topK_num);为 true 时 (experts_num*capacity) | √ |
| permuteTokenId | 输入 | 计算公式中的 permuteTokenId,代表输入 permutedTokens 每个位置对应的 Token 序号 | 取值范围 [0, tokens_num-1] | INT32 | ND | 与 outIndex 相同 | √ |
| routingMapOptional | 可选输入 | 当输入 probsOptional 为空指针时不需要此输入,应传入空指针。代表对应位置的 Token 是否被对应专家处理 | INT8 类型取值支持 0、1;BOOL 类型取值支持 true、false | INT8、BOOL | ND | (tokens_num, experts_num) | √ |
| permutedTokensOptional | 可选输入 | 当输入 probsOptional 为空指针时不需要此输入,应传入空指针 | 数据类型与 unpermutedTokensGrad 相同 | BFLOAT16、FLOAT16、FLOAT | ND | dropAndPad 为 false 时 (tokens_num*topK_num, hidden_size);为 true 时 (experts_num*capacity, hidden_size) | √ |
| probsOptional | 可选输入 | 当不需要时为空指针 | 数据类型与 unpermutedTokensGrad 相同;或者当 unpermutedTokensGrad 是 BFLOAT16 时 probsOptional 支持 FLOAT | BFLOAT16、FLOAT16、FLOAT | ND | 与 routingMapOptional 相同 | √ |
| dropAndPad | 属性 | true 表示开启 dropAndPad,false 表示关闭 dropAndPad | - | BOOL | - | - | - |
| restoreShapeOptional | 属性 | INT64 类型的 aclIntArray。dropAndPad 为 true 时代表 unpermutedTokensGrad 的 shape | - | INT64 | - | - | - |
| permutedTokensGradOut | 输出 | 计算公式中的 permutedTokensGradOut,代表输入 permutedTokens 的梯度 | 数据类型与 unpermutedTokensGrad 相同 | BFLOAT16、FLOAT16、FLOAT | ND | dropAndPad 为 false 时 (tokens_num*topK_num, hidden_size);为 true 时 (experts_num*capacity, hidden_size) | × |
| probsGradOutOptional | 可选输出 | 未输入 probsOptional 时为空指针。输入 probs 的梯度 | 数据类型与 probsOptional 相同 | BFLOAT16、FLOAT16、FLOAT | ND | 与 routingMapOptional 相同 | × |
| workspaceSize | 输出 | 返回需要在 Device 侧申请的 workspace 大小 | - | - | - | - | - |
| executor | 输出 | 返回 op 执行器,包含算子计算流程 | - | - | - | - | - |
注意:输入侧(unpermutedTokensGrad、outIndex、permuteTokenId、routingMapOptional、permutedTokensOptional、probsOptional)均支持非连续 Tensor(标记为 √),而两个输出permutedTokensGradOut与probsGradOutOptional要求连续(标记为 ×)。
返回值与错误码
两段接口均返回aclnnStatus状态码,具体参见 aclnn 返回码。第一段接口完成入参校验,出现以下场景时报错:
| 返回值 | 错误码 | 描述 |
|---|---|---|
| ACLNN_ERR_PARAM_NULLPTR | 161001 | 传入的必选输入、必选输出或必选属性是空指针 |
| ACLNN_ERR_PARAM_INVALID | 161002 | 输入和输出的数据类型和数据格式不在支持范围之内 |
| ACLNN_ERR_INNER_TILING_ERROR | 561002 | 输入 probsOptional 非空,且 dropAndPad 为 false 时,topK_num > 512 |
| ACLNN_ERR_INNER_TILING_ERROR | 561002 | 输入 probsOptional 非空,且 dropAndPad 为 false 时,topK_num 大于 experts_num |
| ACLNN_ERR_INNER_TILING_ERROR | 561002 | 输入 probsOptional 非空,且 dropAndPad 为 false 时,(ubSize - (probTypeLen + 1) * numExpertAlign - (tokenTypeLen + 8) * 256) / (6 * tokenTypeLen + 12) < 1 |
| ACLNN_ERR_INNER_TILING_ERROR | 561002 | 输入 probsOptional 非空,且 dropAndPad 为 true 时,capacity 大于 tokens_num |
| ACLNN_ERR_INNER_TILING_ERROR | 561002 | 输入 probsOptional 非空,且 dropAndPad 为 true 时,hidden_size > 256 * (ubSize - 2080) / (8 + tokenTypeLen) |
| ACLNN_ERR_INNER_TILING_ERROR | 561002 | 输入 probsOptional 非空时,输入 routingMapOptional 或 permutedTokensOptional 为空 |
| ACLNN_ERR_INNER_TILING_ERROR | 561002 | 输入 probsOptional 非空时,probsOptional 数据类型与 unpermutedTokensGrad 不同且 unpermutedTokensGrad 不是 BFLOAT16 |
| ACLNN_ERR_INNER_TILING_ERROR | 561002 | 输入或输出的 shape 不符合要求 |
这些约束在 moe_token_unpermute_with_routing_map_grad_tiling.cpp 中均有对应的OP_CHECK_IF校验逻辑(例如topK > MAX_TOP_K、capacity > 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 <= 512且topK_num <= experts_num; - 要求
experts_num满足(ubSize - (probTypeLen + 1) * numExpertAlign - (tokenTypeLen + 8) * 256) / (6 * tokenTypeLen + 12) >= 1,其中ubSize是芯片 ub 空间大小,probTypeLen是输入probsOptional的数据类型对应字节数,tokenTypeLen是输入unpermutedTokensGrad的数据类型对应字节数,numExpertAlign是experts_num对 32 做向上对齐的结果。
- 要求
- 当输入
probsOptional非空,且dropAndPad为 true 时:- 要求
capacity <= tokens_num; - 要求
hidden_size <= 256 * (ubSize - 2080) / (8 + tokenTypeLen),其中ubSize是芯片 ub 空间大小,tokenTypeLen是输入unpermutedTokensGrad的数据类型对应字节数。
- 要求
从源码看,MAX_TOP_K = 512、INDICES_RESERVE_MAX_NUM = 256等常量定义在 tiling.cpp 中,与上述约束一一对应;同时 moe_token_unpermute_with_routing_map_grad_base.h 中定义了BLOCK_SIZE_512 = 512、FP32_ONE_REPEAT = 64、INDICES_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 = false、tokenNum = 1、hiddenSize = 2、expertNum = 2、topK = 2的最简场景:unpermutedTokensGrad形状为 (1, 2),outIndex/permuteTokenId形状为 (2),routingMap与probs形状为 (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_grad、out_index、permute_token_id为 REQUIRED,routing_map、permuted_tokens、probs为 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 <= tokensNum与hiddenSize上限,计算capacity = numOutTokens / numExpert;TilingForProbNotNonePadFalse:校验topK <= 512、topK <= numExpert,并根据 UB 剩余空间(减去 indices、routingMap 对齐缓冲等)反推numExpert与hiddenSizeAlign的上限。
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 阶段会计算一个tilingKey:tilingKey = 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 | 场景 | 模板实例 |
|---|---|---|
| 0 | probs=None + dropAndPad=false | MoeTokenUnpermuteWithRoutingMapGradProbNoneDropPadFalse |
| 10 | probs=None + dropAndPad=true | MoeTokenUnpermuteWithRoutingMapGradProbNoneDropPadTrue |
| 1 | probs≠None + dropAndPad=false(同精度) | MoeTokenUnpermuteWithRoutingMapGradProbNotNoneDropPadFalse |
| 11 | probs≠None + dropAndPad=true(同精度) | MoeTokenUnpermuteWithRoutingMapGradProbNotNoneDropPadTrue |
| 101 | probs≠None + dropAndPad=false(probs 为 FLOAT 混合精度) | MoeTokenUnpermuteWithRoutingMapGradProbNotNoneDropPadFalse |
| 111 | probs≠None + dropAndPad=true(probs 为 FLOAT 混合精度) | MoeTokenUnpermuteWithRoutingMapGradProbNotNoneDropPadTrue |
以ProbNotNoneDropPadFalse分支(moe_token_unpermute_with_routing_map_grad_prob_not_none_drop_pad_false.h)为例,其核心流程为:
- 将
outIndex通过GetValue读入indicesArray,并依据routingMap逐 token 选择被选中的专家,将对应 probs 写入probsArray(未选中槽位补 0 且地址记为 -1); - 按 hidden 维循环:将
unpermutedTokensGrad一行搬入 UB,Cast/Copy 到 FP32 缓冲,用Muls乘以对应 probs 得到permutedTokensGrad行,写回 GM;同时用Mul计算permutedTokens * unpermutedTokensGrad并通过ReduceSumFunc沿 hidden 维归约得到该槽位的 probs 梯度; - 所有 hidden 分块累加完成后,通过
DataCopyPad按专家位置 scatter 回probsGrad。
ReduceSumFunc定义在 moe_token_unpermute_with_routing_map_grad_base.h,按 hidden 长度分档:大于 4096 先二分累加至 8192 再用多级BlockReduceSum + WholeReduceSum归约到 1 个标量;小于等于 64 则直接WholeReduceSum。BinaryAddFunc实现了二分组加法的并行累加。基类还提供了 MTE2/MTE3/V 单元之间的事件同步封装(SToMTE2Sync、VToMTE3Sync等)与手动 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 下按ascend910b、ascend910_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_map、moe_token_unpermute_with_routing_map等配套实现,形成完整的「路由—重排—专家计算—反重排—梯度」闭环认知。
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考