ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

ascend-transformer-boost 的 AllToAllV 算子:可变长度全交换通信的实现解析与使用指南

2026/9/18 17:30:20 拓冰建站 浏览量
ascend-transformer-boost 的 AllToAllV 算子:可变长度全交换通信的实现解析与使用指南 ascend-transformer-boost 的 AllToAllV 算子可变长度全交换通信的实现解析与使用指南【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库基于华为Ascend AI处理器提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost本文基于 CANN ascend-transformer-boost 仓库中AllToAllVall_to_allv算子的路由文档及其源码展开核心讲解该算子向通信域内所有卡发送、从所有卡接收、且各 rank 之间数据量可定制的集合通信能力。读者读完本文后将掌握 AllToAllV 的参数语义与合法性约束、InferShape 的输出形状推导规则、HcclRunner 的三种通信域初始化方式以及如何通过仓库自带的 PyTorch 高层测试用例验证该算子的正确性。AllToAllV是 ascend-transformer-boost 提供的可变长度 AllToAll集合通信算子对应等长版本的all_to_all用于 Transformer 分布式推理/训练中需要按 rank 定制收发数据量的场景如 MoE 路由分发、负载不均衡的张量搬运。与普通 AllToAll 相比其核心差异在于sendCounts[i]/recvCounts[i]允许每个 rank 对各不相同配合sdispls[i]/rdispls[i]偏移量可实现非均匀数据的精确搬运。1. 算子定位分类、复杂度与 Runner 类型路由文档.agent/knowledge/routing/all_to_allv.md给出的元信息如下分类infer推理侧集合通信算子复杂度S单算子结构简单文件数42 个 Operation 文件 2 个 Runner 文件Runner 类型实际实现为AllToAllVHcclRunner继承自仓库统一封装的HcclRunner见 src/atb/runner/hccl_runner.hACLNN否不提供 aclnn 接口直接以 Operation 形式使用硬件约束当前仅支持 Atlas 800I A2 推理产品源码在创建算子时通过Config::Is910B()校验见下文从源码结构看该算子没有独立的 AscendC kernel 实现其执行完全委托给 HCCL 通信库的HcclAlltoAllV原语因此路由文档中标注的 Kernel 目录属于路由模板的通用字段实际调用链不经过src/kernels/mixkernels/下的任何 kernel。2. 文件清单与推荐阅读顺序按路由文档整理AllToAllV算子共 4 个文件全部位于 src/ops/ops_infer/all_to_allv/ 目录#文件角色阅读重点1all_to_allv_operation.cppOperation 定义CreateOperation()参数校验、InferShapeImpl()形状推导、CreateRunner()决策逻辑2all_to_allv_operation.hOperation 定义输入/输出数量各 1 个、InferShapeImpl签名、持有的AllToAllVParam成员3all_to_allv_hccl_runner.cppRunner 实现ExecuteImpl()中HcclAlltoAllV的实参组装与错误处理4all_to_allv_hccl_runner.hRunner 头文件三个构造重载分别对应三种通信域初始化方式推荐的阅读顺序先读.h建立输入输出与参数的整体认知再读.cpp理解创建与推理逻辑最后深入 Runner 看实际通信执行。详细知识条目可参考.agent/knowledge/ops/communication/all_to_allv/index.md其中明确该算子 Pipeline 为集合通信 — 可变长度 AllToAll — HCCL 通信库Related 指向等长版all_to_all。3. 参数详解AllToAllVParam参数结构体定义在头文件 include/atb/infer_op_params.h 的AllToAllVParam约 2429 行起。各字段语义如下字段类型默认值说明rankint0当前卡所属通信编号rankSizeint0通信域内卡的数量rankRootint0主通信编号rank 0 为默认主卡sendCountsvectorint64_t空发送数据量数组。sendCounts[i] n表示本 rank 发给 rank i 共 n 个元素以元素个数计非字节数sdisplsvectorint64_t空发送偏移量数组。sdispls[i] n表示本 rank 从输入张量起始偏移 n 的位置开始向 rank i 发送数据recvCountsvectorint64_t空接收数据量数组。recvCounts[i] n表示本 rank 从 rank i 接收 n 个元素rdisplsvectorint64_t空接收偏移量数组。rdispls[i] n表示本 rank 将 rank i 的数据写入输出张量起始偏移 n 的位置backendstringhccl通信计算类型仅支持 hcclhcclCommHcclCommnullptrHCCL 通信域指针。默认为空由加速库创建若用户自行管理通信域则传入该指针加速库直接复用commModeCommModeCOMM_MULTI_PROCESS通信模式。hccl 多线程场景只支持外部传入通信域方式rankTableFilestring空集群信息配置文件路径适用单机与多机通信场景仅支持 hccl 后端单机配置了 rankTable 时以 ranktable 初始化通信域commDomainstring空通信域名标识多通信域时使用rsvuint8_t[64]{0}预留参数注意sendCounts/recvCounts的计数单位是元素个数而非字节数实际字节数由张量 dtype 决定例如 float32 下sendCounts[i]n表示 n 个 float32 数据。4. 创建与参数校验CreateOperation 的硬性约束算子通过模板特化CreateOperation(const infer::AllToAllVParam opParam, Operation **operation)创建all_to_allv_operation.cpp 第 27 行起。源码中的校验链决定了该算子的使用边界backend 必须为 hccl否则打印backend must be hccl并返回ERROR_INVALID_PARAM。硬件必须是 Atlas 800I A2通过GetSingletonConfig().Is910B()判断否则返回ERROR_INVALID_PARAM。这是仓库中明确写死的硬件前提其他昇腾产品上该算子不可用。分布式初始化检查调用OperationUtil::DistributedInitCheckinfer::AllToAllVParam(opParam)确保通信环境已正确初始化。四个数组长度必须等于 rankSizesendCounts、recvCounts、sdispls、rdispls的 size 必须同时等于opParam.rankSize否则报错... size should be equal to ranksize。数组元素必须非负四组数组的每个元素都不能小于 0否则报错... should more than zero。以上任一校验失败都会导致算子创建失败返回ERROR_INVALID_PARAM而不仅是运行时告警——这要求在组装参数时严格保证四个数组长度一致且取值非负。5. 形状推导InferShapeImpl 的输出规则AllToAllVOperation的输入、输出张量各为 1 个IN_TENSOR_NUM 1、OUT_TENSOR_NUM 1。InferShapeImpl()第 85 行起的逻辑如下先执行ParamCheck遍历sendCounts校验sendCounts[i] sdispls[i]不越界即不超过输入张量总元素数Utils::GetTensorNumel(inTensorDesc)防止发送区间超出输入范围。输出形状强制改写为二维dimNum 2dims[0] 1。累加所有recvCounts得到count过程中做 int64 溢出保护要求count 0。逐项校验recvCounts[i] rdispls[i]不超过count防止接收区间超出输出范围。最终输出形状为[1, sum(recvCounts)]即输出是一维展开后的拼接结果。这意味着输入张量是发送数据的连续缓冲区输出张量则是所有 rank 接收数据的按 rdispls 拼接结果形状恒为[1, ΣrecvCounts]与输入形状无关输入只作为数据源缓冲区。6. Runner 决策与 HCCL 执行链路CreateRunner()第 131 行起根据param_.backend是否为 hccl 做决策hcclComm nullptr且rankTableFile非空 →AllToAllVHcclRunner(param_, true)以 rankTable 方式初始化通信域hcclComm nullptr且rankTableFile为空 →AllToAllVHcclRunner(param_, false)以rank/rankSize/rankRoot/commDomain方式初始化hcclComm ! nullptr→AllToAllVHcclRunner(param_, hcclComm)直接复用用户传入的通信域。对应的三个构造重载定义在 all_to_allv_hccl_runner.h内部通过委托构造初始化HcclRunner基类src/atb/runner/hccl_runner.h基类提供CreateHcclCommInMulitProcess、CreateHcclCommInMulitProcessByRootInfo、CreateHcclCommInMulitProcessByRankFile等多种通信域创建路径并通过Mki::ShareMemory实现多进程间的 rootInfo 共享与 barrier 同步。实际执行在ExecuteImpl()all_to_allv_hccl_runner.cpp 第 36 行起前置检查包括hcclComm_非空否则返回ERROR_COMM_EMPTY输入/输出deviceData非空否则返回ERROR_INVALID_PARAM。随后组装 HCCL 原语调用HcclResult ret HcclAlltoAllV( runnerVariantPack.inTensors[0].deviceData, // 输入数据指针 static_castvoid *(param_.sendCounts.data()), // 发送数据量数组 static_castvoid *(param_.sdispls.data()), // 发送偏移量数组 GetHcclDtype(runnerVariantPack.inTensors[0].desc.dtype), // 输入 dtype 映射为 HcclDataType runnerVariantPack.outTensors[0].deviceData, // 输出数据指针 static_castvoid *(param_.recvCounts.data()), // 接收数据量数组 static_castvoid *(param_.rdispls.data()), // 接收偏移量数组 GetHcclDtype(runnerVariantPack.outTensors[0].desc.dtype), // 输出 dtype 映射 hcclComm_.get(), // 通信域 GetExecuteStream(runnerVariantPack.context)); // 当前执行流执行失败时通过ConvertHcclResultToStatus(ret)将 HCCL 返回码转换为 ATB 状态码并记录日志。Runner 通过REG_RUNNER_TYPE(AllToAllVHcclRunner)注册到仓库的 Runner 工厂由CreateRunner()返回后交给框架调度执行。7. 使用前提与限制基于源码事实仅 Atlas 800I A2Ascend910B 系列推理产品支持其余硬件创建算子直接失败backend 固定为 hccl不提供其他通信后端四个参数数组长度必须等于 rankSize 且元素非负否则创建失败输入输出均为 device 上的连续缓冲区输出形状被固定为[1, ΣrecvCounts]当前仅支持多进程通信模式COMM_MULTI_PROCESShccl 多线程需外部传入通信域。8. 测试验证从高层测试用例看正确用法仓库在 tests/high_level_test/AllToAllVOperation/Smoke/test_all_to_allv_operation.py 提供了完整的 PyTorch 侧验证用例可视为该算子的标准调用范例环境初始化每个进程torch_npu.npu.set_device(rank)绑定设备并通过torch.classes.load_library加载libatb_test_framework.so与libatb.so构造torch.classes.OperationTorch.OperationTorch(AllToAllVOperation)。参数构造随机生成sendCountsnp.random.randint(1, 8, size[world_size, world_size])并依据前序累加方式推导sdispls利用 AllToAllV 的转置对称性令recvout[i][j] sendcount[j][i]同样累加得到recvdis最后通过json.dumps序列化为参数字典并调用set_paramacl_param json.dumps( {rank: rank, rankSize: world_size, sendCounts: sendcount[rank], sdispls: senddisp[rank], recvCounts: recvout[rank], rdispls: recvdis[rank], rankRoot: 0, backend: hccl}) all_to_allv_operation.set_param(acl_param) acl_out_tensor all_to_allv_operation.execute([inTensors[rank].npu()])[0]Golden 校验golden 张量由各 rank 发送给自己的那一段按接收顺序拼接而成最终用torch.allclose(..., rtol0.001, atol0.001)对比测试随机抽取 1~8 的 world_size 分别跑三轮。硬件门控测试首行检查operation_test.get_soc_version() Ascend910B与源码Is910B()约束一致——这从测试侧再次印证了硬件限制。此外 tests/apitest/opstest/python/operations/all_to_allv/ 还提供单机与多机的 apitest 用例配合 tests/apitest/opstest/csv/all_to_allv.csv 中的 CSV 参数矩阵覆盖不同 dtype、shape、counts 组合可系统验证边界场景。9. 小结与延伸阅读AllToAllV以Operation HcclRunner两段式结构实现了可变长度全交换通信Operation 层负责参数合法性校验与[1, ΣrecvCounts]形状推导Runner 层直接封装 HCCL 的HcclAlltoAllV原语并通过 rankTable / rank-rankSize-rankRoot / 外部 hcclComm 三种方式灵活初始化通信域。对等长数据场景可对照阅读all_to_all路由与实现如需了解通信算子的整体分类可参考知识库主索引.agent/knowledge/README.md。【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库基于华为Ascend AI处理器提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考