十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

CANN ops-transformer aclnnMatmulAlltoAll 算子接口深度解析:Matmul 与 AlltoAll 通算融合实践指南

CANN ops-transformer aclnnMatmulAlltoAll 算子接口深度解析:Matmul 与 AlltoAll 通算融合实践指南 CANN ops-transformer aclnnMatmulAlltoAll 算子接口深度解析Matmul 与 AlltoAll 通算融合实践指南【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformeraclnnMatmulAlltoAll 是 CANN ops-transformer 算子库mc2/matmul_allto_all中面向多卡 MoE 分布式场景的核心融合算子它将Matmul 矩阵乘计算、Permute 数据重排与 AlltoAll 集合通信融合为单次算子调用遵循先计算后通信的执行语义用于解决注意力MoE门控路由后专家并行下的 AlltoAll 数据交换性能瓶颈。读完本文你将掌握该接口的两段式调用流程、全部参数语义、设备差异约束与可编译运行的多卡示例并能结合仓库源码理解其内部实现链路。产品支持情况aclnnMatmulAlltoAll 并非在所有昇腾产品上都可用支持情况如下表产品是否支持Ascend 950PR / Ascend 950DT支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品支持Atlas 200I/500 A2 推理产品不支持Atlas 推理系列产品不支持Atlas 训练系列产品不支持功能说明与计算公式接口功能完成Matmul 计算、Permute保证通信后地址连续和 AlltoAll 通信的融合先计算后通信。假设x1的 shape 为(BS, H1)x2的 shape 为(H1, H2)rankSize为 NPU 卡数则计算过程为$$ computeOut x1 x2 bias \ permutedOut computeOut.view(BS, rankSize, H2/rankSize).permute(1, 0, 2) \ output AlltoAll(permutedOut).view(rankSize*BS, H2/rankSize) $$其中H2为 Matmul 输出维度被切分为rankSize份view permute使每个 rank 需要的数据在本地地址连续从而保证 AlltoAll 通信后各 rank 拿到(BS, H2/rankSize)的数据块并拼接为(rankSize*BS, H2/rankSize)的连续输出。该语义与图模式算子原型 matmul_allto_all_proto.h 中 Fusion op of alltoall and matmul 的定位一致。从仓库 READMEmc2/matmul_allto_all/README.md可见MatmulAlltoAll 家族还提供aclnnMatmulAlltoAllV2、aclnnQuantMatmulAlltoAllK-C 量化、aclnnQuantMatmulAlltoAllV2mx 量化等变体本接口为非量化场景的基础版本。两段式接口与函数原型每个算子分为两段式接口必须先调用aclnnMatmulAlltoAllGetWorkspaceSize接口获取计算所需 workspace 大小以及包含了算子计算流程的执行器再调用aclnnMatmulAlltoAll接口执行计算。第一段接口原型aclnnStatus aclnnMatmulAlltoAllGetWorkspaceSize( const aclTensor* x1, const aclTensor* x2, const aclTensor* biasOptional, const aclIntArray* alltoAllAxesOptional, const char* group, bool transposeX1, bool transposeX2, const aclTensor* output, uint64_t* workspaceSize, aclOpExecutor** executor)第二段接口原型aclnnStatus aclnnMatmulAlltoAll( void* workspace, uint64_t workspaceSize, aclOpExecutor* executor, aclrtStream stream)两段式接口的语义参见仓库 docs/zh/context/two_phase_api.mdworkspace 是算子在 NPU 上完成计算所需的临时内存除输入/输出外第二段接口不可重复调用。aclnnMatmulAlltoAllGetWorkspaceSize 参数说明参数明细参数名输入/输出描述使用说明数据类型数据格式维度(shape)非连续tensorx1输入融合算子的左矩阵输入对应公式中的 x1该输入作为 MatMul 计算的左矩阵输入FLOAT16、BFLOAT16ND2维shape 为 (BS, H1)xx2输入融合算子的右矩阵输入也是 MatMul 计算的右矩阵直接作为 MatMul 计算的右矩阵输入FLOAT16、BFLOAT16ND2维shape 为 (H1, H2)不同设备型号支持情况不同参见约束说明biasOptional输入阵乘运算后累加的偏置对应公式中的 bias支持传入空指针场景根据设备型号对数据类型有不同限制详细参见约束说明FLOAT16、BFLOAT16、FLOAT32ND1维shape 为 (H2)xalltoAllAxesOptional输入可选输入AlltoAll 和 Permute 数据交换的方向支持配置空或者 [-1, -2]传入空时默认按 [-1, -2] 处理表示将输入由 (BS, H2) 转为 (BS * rankSize, H2 / rankSize)aclIntArray*元素类型 INT64-1维shape 为 (2)-group输入标识列组的字符串即通信域名称通过 Hccl 接口 HcclGetCommName 获取 commName 作为该参数字符串长度要求 (0, 128)----transposeX1输入标识左矩阵是否转置过暂不支持配置为 True----transposeX2输入标识右矩阵是否转置过配置为 True 时右矩阵 Shape 为 (H2, H1)----output输出最终的计算结果数据类型与输入 x1 保持一致FLOAT16、BFLOAT16ND2维shape 为 (BS*rankSize, H2/rankSize)xworkspaceSize输出返回需要在 Device 侧申请的 workspace 大小-----executor输出返回 op 执行器包含了算子的计算流程-----返回值与错误码返回aclnnStatus状态码具体参见 aclnn返回码。第一段接口完成入参校验出现以下场景时报错返回值错误码描述ACLNN_ERR_PARAM_NULLPTR161001输入和输出的必选参数 Tensor 是空指针ACLNN_ERR_PARAM_INVALID161002输入和输出的数据类型不在支持的范围内ACLNN_ERR_PARAM_INVALID161002输入 Tensor 为空 TensorACLNN_ERR_PARAM_INVALID161002alltoAllAxesOptional 非法ACLNN_ERR_PARAM_INVALID161002transposeX1 为 trueACLNN_ERR_PARAM_INVALID161002通信域长度非法ACLNN_ERR_PARAM_INVALID161002输入输出 Tensor 维度不合法ACLNN_ERR_PARAM_INVALID161002输入输出 format 为私有格式上述校验在源码中的对应实现位于 matmul_allto_all_base.cpp 的CheckAndHandleParams依次检查必选参数非空CheckNotNull返回ACLNN_ERR_PARAM_NULLPTR、空 Tensor、shape 合法性、数据类型支持范围950 与 910B 的 bias 类型限制不同分别走CheckAllDtypesValid与CheckAllDtypesValid910B、format 是否为 ND、alltoAllAxes 是否为空或[-1, -2]、transposeX1 是否非法、group 长度是否超过 128。其中 format 若为私有格式直接报ACLNN_ERR_PARAM_INVALID若非 ND 但合法则通过l0op::ReFormat自动转换到 ND 后再计算ReFormatNotND。aclnnMatmulAlltoAll 执行接口参数说明参数名输入/输出描述workspace输入在 Device 侧申请的 workspace 内存地址workspaceSize输入在 Device 侧申请的 workspace 大小由第一段接口 aclnnMatmulAlltoAllGetWorkspaceSize 获取executor输入op 执行器包含了算子计算流程stream输入指定执行任务的 Stream返回值返回aclnnStatus状态码具体参见 aclnn返回码。从源码看第一段接口在 aclnn_matmul_allto_all.cpp 中会根据当前 NPU 架构自动选择通信模式DAV_3510A3 架构或ASCEND910_93走ai_cpu其余如 A2/910B走aiv即 MTE/AIV 通信随后转调统一的aclnnMatmulAlltoAllBaseGetWorkspaceSize第二段接口则直接转调aclnnMatmulAlltoAllBase。非量化接口内部还会补齐量化相关输入的默认值后调用 L0 层 Inner 接口见 matmul_allto_all_base.cpp 附近量化与图模式共用同一套 L0 接口。约束说明使用 aclnnMatmulAlltoAll 时必须遵守以下约束默认支持确定性计算。NPU 卡数rankSize限制按设备型号Atlas A2 训练系列产品 / Atlas A2 推理系列产品支持 2、4、8 卡Atlas A3 训练系列产品 / Atlas A3 推理系列产品支持 2、4、8、16 卡Ascend 950PR / Ascend 950DT支持 2、4、8、16 卡。参数说明中 shape 使用的变量H2 必须整除 NPU 卡数。H1 范围仅支持 [1, 65535]。H2 取值范围按设备型号Atlas A2 训练系列产品 / Atlas A2 推理系列产品不得超过 368640不得小于 2Atlas A3 训练系列产品 / Atlas A3 推理系列产品不得超过 2147483647INT32_MAX不得小于 2Ascend 950PR / Ascend 950DT不得超过 2147483647INT32_MAX不得小于 2。BS*rankSize 的值不得超过 2147483647INT32_MAX不得小于 0。空 tensor 支持度按设备型号Atlas A2 训练系列产品 / Atlas A2 推理系列产品不支持任何空 tensorAtlas A3 训练系列产品 / Atlas A3 推理系列产品不支持任何空 tensorAscend 950PR / Ascend 950DT仅支持输入 x1 的第一维度BS为 0 的空 tensor其它空 tensor 均不支持对应源码 matmul_allto_all_base.cpp 中CheckNotEmptyTensor对 950 上 x1.dimM 为 0 的放行逻辑该场景即空 token 提示词。非连续 tensor 支持度按设备型号Atlas A2 训练系列产品 / Atlas A2 推理系列产品不支持任何非连续 tensorAtlas A3 训练系列产品 / Atlas A3 推理系列产品不支持任何非连续 tensorAscend 950PR / Ascend 950DT仅支持 x2 为非连续 tensor其它非连续 tensor 均不支持。x1、x2 计算输入的数据类型要和 output 计算输出的数据类型一致传入的 x1、x2 与 output 均不为空指针。biasOptional 数据类型限制按设备型号Atlas A2 训练系列产品 / Atlas A2 推理系列产品output 计算输出的数据类型为 FLOAT16 时biasOptional 计算输入的数据类型支持 FLOAT16output 计算输出的数据类型为 BFLOAT16 时biasOptional 计算输入的数据类型支持 FLOAT32Atlas A3 训练系列产品 / Atlas A3 推理系列产品x1/x2 计算输入的数据类型为 FLOAT16 时biasOptional 计算输入的数据类型支持 FLOAT16x1/x2 计算输入的数据类型为 BFLOAT16 时biasOptional 计算输入的数据类型支持 FLOAT32Ascend 950PR / Ascend 950DTx1/x2 计算输入的数据类型为 FLOAT16 时biasOptional 计算输入的数据类型支持 FLOAT16 和 FLOAT32x1/x2 计算输入的数据类型为 BFLOAT16 时biasOptional 计算输入的数据类型支持 BFLOAT16 和 FLOAT32。通算融合算子不支持并发调用不同的通算融合算子也不支持并发调用。不支持跨超节点通信只支持超节点内。通信约束按设备型号Atlas A2 训练系列产品 / Atlas A2 推理系列产品支持 MTE 通信且通信缓冲区大于等于 200MBAtlas A3 训练系列产品 / Atlas A3 推理系列产品支持 AI_CPU 通信Ascend 950PR / Ascend 950DT支持 AI_CPU 通信。从图模式算子原型 matmul_allto_all_proto.h 可以看到all2all_axes属性默认值为{-1, -2}、x1_quant_mode/x2_quant_mode默认 0不量化、comm_mode默认ai_cpu与上述约束一一对应。infershape 实现 matmul_allto_all_infershape.cpp 中同样校验了 rank 数必须在{2, 4, 8, 16}内、alltoAllAxes 必须为[-1, -2]、K 轴H1上限 65535 等规则。调用示例示例代码如下仅供参考具体编译和执行过程请参考编译与运行样例仓库内可直接参考 examples/test_aclnn_matmul_allto_all.cpp 与测试用例 tests/ut/op_api/test_aclnn_matmul_allto_all.cpp。说明本示例代码调用了部分 HCCL 集合通信库接口HcclGetCommName、HcclCommInitAll、HcclCommDestroy请参考《HCCL API (C)》文档。Atlas A2 训练系列产品 / Atlas A2 推理系列产品示例#include thread #include iostream #include string #include cstring #include vector #include acl/acl.h #include hccl/hccl.h #include aclnn/opdev/fp16_t.h #include aclnnop/aclnn_matmul_allto_all.h int ndev 2; #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::vectorint64_t shape) { int64_t shapeSize 1; for (auto i: shape) { shapeSize * i; } return shapeSize; } templatetypename T int CreateAclTensor(const std::vectorT hostData, const std::vectorint64_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::vectorint64_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; } struct Args { uint32_t rankId; HcclComm hcclComm; aclrtStream stream; aclrtContext context; }; int launchOneThreadMatmulAlltoAll(Args args) { int ret; ret aclrtSetCurrentContext(args.context); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtSetCurrentContext failed. ERROR: %d\n, ret); return ret); char hcom_name[128]; ret HcclGetCommName(args.hcclComm, hcom_name); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT([ERROR] HcclGetCommName failed. ret %d \n, ret); return -1); LOG_PRINT([INFO] rank %d hcom: %s stream: %p, context : %p\n, args.rankId, hcom_name, args.stream, args.context); std::vectorint64_t x1Shape {32, 64}; std::vectorint64_t x2Shape {64, 128}; std::vectorint64_t biasShape {128}; std::vectorint64_t outShape {64, 64}; void *x1DeviceAddr nullptr; void *x2DeviceAddr nullptr; void *biasDeviceAddr nullptr; void *outDeviceAddr nullptr; aclTensor *x1 nullptr; aclTensor *x2 nullptr; aclTensor *bias nullptr; aclTensor *out nullptr; int64_t a2aAxes[2] {-1, -2}; aclIntArray* alltoAllAxesOptional aclCreateIntArray(a2aAxes, static_castuint64_t(2)); uint64_t workspaceSize 0; aclOpExecutor *executor; void *workspaceAddr nullptr; long long x1ShapeSize GetShapeSize(x1Shape); long long x2ShapeSize GetShapeSize(x2Shape); long long biasShapeSize GetShapeSize(biasShape); long long outShapeSize GetShapeSize(outShape); std::vectorop::fp16_t x1HostData(x1ShapeSize, 1); std::vectorop::fp16_t x2HostData(x2ShapeSize, 1); std::vectorop::fp16_t biasHostData(biasShapeSize, 1); std::vectorop::fp16_t outHostData(outShapeSize, 0); // 创建tensor ret CreateAclTensor(x1HostData, x1Shape, x1DeviceAddr, aclDataType::ACL_FLOAT16, x1); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(x2HostData, x2Shape, x2DeviceAddr, aclDataType::ACL_FLOAT16, x2); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(biasHostData, biasShape, biasDeviceAddr, aclDataType::ACL_FLOAT16, bias); CHECK_RET(ret ACL_SUCCESS, return ret); ret CreateAclTensor(outHostData, outShape, outDeviceAddr, aclDataType::ACL_FLOAT16, out); CHECK_RET(ret ACL_SUCCESS, return ret); // 调用第一段接口 ret aclnnMatmulAlltoAllGetWorkspaceSize(x1, x2, bias, alltoAllAxesOptional, hcom_name, false, false, out, workspaceSize, executor); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclnnMatmulAlltoAllGetWorkspaceSize failed. ERROR: %d\n, ret); return ret); // 根据第一段接口计算出的workspaceSize申请device内存 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 aclnnMatmulAlltoAll(workspaceAddr, workspaceSize, executor, args.stream); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclnnMatmulAlltoAll failed. ERROR: %d\n, ret); return ret); //固定写法同步等待任务执行结束 ret aclrtSynchronizeStreamWithTimeout(args.stream, 10000); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtSynchronizeStream failed. ERROR: %d\n, ret); return ret); LOG_PRINT(device%d aclnnMatmulAlltoAll execute success \n, args.rankId); // 释放device资源需要根据具体API的接口定义修改 if (x1 ! nullptr) { aclDestroyTensor(x1); } if (x2 ! nullptr) { aclDestroyTensor(x2); } if (bias ! nullptr) { aclDestroyTensor(bias); } if (out ! nullptr) { aclDestroyTensor(out); } if (x1DeviceAddr ! nullptr) { aclrtFree(x1DeviceAddr); } if (x2DeviceAddr ! nullptr) { aclrtFree(x2DeviceAddr); } if (biasDeviceAddr ! nullptr) { aclrtFree(biasDeviceAddr); } if (outDeviceAddr ! nullptr) { aclrtFree(outDeviceAddr); } if (workspaceSize 0) { aclrtFree(workspaceAddr); } aclrtDestroyStream(args.stream); HcclCommDestroy(args.hcclComm); aclrtDestroyContext(args.context); aclrtResetDevice(args.rankId); return 0; } int main(int argc, char *argv[]) { // 本样例基于Atlas A2实现必须在Atlas A2上运行 int ret; int32_t devices[ndev]; for (int i 0; i ndev; i) { devices[i] i; } HcclComm comms[128]; ret aclInit(nullptr); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclInit failed. ERROR: %d\n, ret); return ret); // 初始化集合通信域 for (int i 0; i ndev; i) { ret aclrtSetDevice(devices[i]); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtSetDevice failed. ERROR: %d\n, ret); return ret); } ret HcclCommInitAll(ndev, devices, comms); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(HcclCommInitAll failed. ERROR: %d\n, ret); return ret); Args args[ndev]; aclrtStream stream[ndev]; aclrtContext context[ndev]; for (uint32_t rankId 0; rankId ndev; rankId) { ret aclrtSetDevice(rankId); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtSetDevice failed. ERROR: %d\n, ret); return ret); ret aclrtCreateContext(context[rankId], rankId); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtCreateContext failed. ERROR: %d\n, ret); return ret); ret aclrtCreateStream(stream[rankId]); CHECK_RET(ret ACL_SUCCESS, LOG_PRINT(aclrtCreateStream failed. ERROR: %d\n, ret); return ret); } // 启动多线程 std::vectorstd::unique_ptrstd::thread threads(ndev); for (uint32_t rankId 0; rankId ndev; rankId) { args[rankId].rankId rankId; args[rankId].hcclComm comms[rankId]; args[rankId].stream stream[rankId]; args[rankId].context context[rankId]; threads[rankId].reset(new(std::nothrow) std::thread(launchOneThreadMatmulAlltoAll, std::ref(args[rankId]))); } for (uint32_t rankId 0; rankId ndev; rankId) { threads[rankId]-join(); } aclFinalize(); return 0; }Ascend 950PR / Ascend 950DT 示例Ascend 950 版本与 A2 版本整体代码结构完全一致仅存在两处差异头文件差异950 版本无需包含aclnn/opdev/fp16_t.h仅包含acl/acl.h、hccl/hccl.h与aclnnop/aclnn_matmul_allto_all.hHost 侧数据容器差异950 版本使用std::vectorint16_t作为 fp16 数据的 Host 容器不再使用op::fp16_tstd::vectorint16_t x1HostData(x1ShapeSize, 1); std::vectorint16_t x2HostData(x2ShapeSize, 1); std::vectorint16_t biasHostData(biasShapeSize, 1); std::vectorint16_t outHostData(outShapeSize, 0); // 创建tensor ret CreateAclTensor(x1HostData, x1Shape, x1DeviceAddr, aclDataType::ACL_FLOAT16, x1); ret CreateAclTensor(x2HostData, x2Shape, x2DeviceAddr, aclDataType::ACL_FLOAT16, x2); ret CreateAclTensor(biasHostData, biasShape, biasDeviceAddr, aclDataType::ACL_FLOAT16, bias); ret CreateAclTensor(outHostData, outShape, outDeviceAddr, aclDataType::ACL_FLOAT16, out);其余部分CHECK_RET/LOG_PRINT宏、GetShapeSize、CreateAclTensor、Args结构体、launchOneThreadMatmulAlltoAll函数体、main中aclInit→aclrtSetDevice→HcclCommInitAll→ 建流建上下文 → 多线程启动 →join→aclFinalize的流程与 A2 版本完全相同主函数注释标明本样例基于 Ascend 950PR/Ascend 950DT 实现必须在 Ascend 950PR/Ascend 950DT 上运行。示例要点解读通信域初始化HcclCommInitAll一次性初始化ndev个设备示例中ndev 2的集合通信域每个线程使用各自的HcclComm通信域名称获取HcclGetCommName获取的hcom_name直接作为group参数传入第一段接口字符串缓冲区char hcom_name[128]与 group 长度约束(0, 128)对应shape 语义示例中x1 (32, 64)、x2 (64, 128)、bias (128)、output (64, 64)即BS32, H164, H2128, rankSize2满足H2 % rankSize 0与output (rankSize*BS, H2/rankSize)的约束两段式固定流程第一段接口获取workspaceSize与executor→workspaceSize 0时aclrtMalloc申请 workspace → 第二段接口执行 →aclrtSynchronizeStreamWithTimeout同步等待 → 按序释放 tensor、device 内存、workspace、stream、通信域与 context。内部实现链路从 aclnn 接口到 NPU Kernel结合仓库源码可以还原该算子的完整实现链路便于理解其在 MC2Multi-Compute-Communication通算融合框架中的位置aclnn 入口Host L2 层aclnnMatmulAlltoAllGetWorkspaceSize按当前 NPU 架构选择通信模式ai_cpu / aiv后转调aclnnMatmulAlltoAllBaseGetWorkspaceSizeop_api/aclnn_matmul_allto_all.cpp、op_api/matmul_allto_all_base.cpp。基类实现完成上述全部参数校验并补齐量化输入的默认值后调用 L0 层 Inner 接口非量化与量化共用同一 L0 接口。图模式定义算子原型注册于 op_graph/matmul_allto_all_proto.h定义了groupString必选、world_sizeInt必选、all2all_axes默认{-1,-2}、x1_quant_mode/x2_quant_mode/comm_quant_mode、transpose_x1/transpose_x2默认 false、comm_mode默认ai_cpu等属性图模式下还通过 fusion_pass 做算子融合与转置优化。Infershape 与 TilingHost 层matmul_allto_all_infershape.cpp 负责图模式下的形状推导与合法性校验rank 数、alltoAllAxes、H1 上限 65535、量化模式枚举等op_host/op_tiling 下按架构分 arch22910B/A2 系列与 arch35950 系列提供切分策略其中 allto_all_formulaic_tiling.cpp 会估算 Matmul 计算时间与 AlltoAll 通信时间的比值ratioCalcComm_通过动态调整切分参数实现计算掩盖通信的平衡allto_all_comm_algo_table.cpp 则按通信引擎CCU 等与拓扑选择具体的 AlltoAll 算法名如CcuSchedAllToAllSoleMesh。Kernel 执行Device 层op_kernel 下matmul_allto_all_kernel_base.h、matmul_allto_all_pipeline.h与 arch22/arch35 目录中的实现负责最终在 NPU 上的流水执行——先完成 Matmul含 bias 累加再对输出做view permute重排使各 rank 数据本地连续最后触发 AlltoAll 通信。该算子的单算子 API 调用方式非图模式与图模式GE 构图共用同一份算子原型与 Kernel 实现aclnn 接口是面向开发者最直接的使用入口。总结与选型建议aclnnMatmulAlltoAll 适用于需要矩阵乘 列切分 跨卡数据交换一步到位的分布式训练/推理场景典型如 MoE 专家并行中的数据分发阶段。选型与使用时的关键决策点确认设备型号仅 A2/A3/950 系列支持且各型号在卡数、H2 上限、空 tensor、非连续 tensor、bias 类型、通信引擎MTE / AI_CPU上存在差异务必按实际型号对照约束说明配置shape 规划确保H2 % rankSize 0、H1 ∈ [1, 65535]、BS * rankSize INT32_MAX输出 shape 固定为(rankSize*BS, H2/rankSize)通信域准备通过HcclCommInitAll建域、HcclGetCommName取commName作为group参数如需量化K-C 量化INT8 输入与 mx 量化场景请改用同目录下的aclnnQuantMatmulAlltoAll/aclnnQuantMatmulAlltoAllV2接口详见 mc2/matmul_allto_all/README.md 及其文档 docs/aclnnQuantMatmulAlltoAll.md。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表