CANN ATB all_to_allvv2 算子深度解析:基于 HCCL 的可变长度 AllToAll 集合通信增强版
【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库,基于华为Ascend AI处理器,提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost
本文面向需要在华为昇腾(Ascend)多卡集群上做全交换(All-to-All)通信的开发者和 Agent,系统讲解 ascend-transformer-boost 中all_to_allvv2算子的设计定位、参数结构、张量契约、校验逻辑与底层 HCCL 调用链。读完本文,你将掌握该算子的完整数据流(从 Operation 参数校验、Shape 推导到HcclAlltoAllV落盘执行),能够正确构造 6 个输入张量并配置通信域,也能理解它与all_to_all、all_to_allv之间的演进关系。
1. 算子定位:通信域内的可变长度数据全交换
all_to_allvv2是 ATB(Ascend Transformer Boost)推理侧(ops_infer)提供的一个集合通信算子,知识条目将其归类为communication类型、tier: S、type: single(单一算子,非融合算子),底层仅依赖 HCCL 通信库(详见 .agent/knowledge/ops/communication/all_to_allvv2/index.md)。
它解决的是典型的"全交换(All-to-All)"问题:通信域内的每一张卡,都向其他所有卡发送各自定制数量的数据,同时从其他所有卡接收定制数量的数据。相比等长版all_to_all,其每次通信的数据量(count)与偏移(displs)均可按目标 rank 定制,属于可变长度(variable-length)AllToAll 语义;相比all_to_allv,它在参数组织、校验与增强能力上做了演进,是当前仓库中的"v2 增强版"。
- 源码目录:src/ops/ops_infer/all_to_allvv2/(共 4 个文件:operation 定义 + hccl runner)
- 参数头文件:include/atb/infer_op_params.h
- 相关算子:
all_to_all(等长版,双后端 HCCL/LCCL,见 src/ops/ops_infer/all_to_all/)、all_to_allv(可变长度版,仅 HCCL,见 src/ops/ops_infer/all_to_allv/)
从源码结构看,与all_to_all支持 HCCL/LCCL 双后端不同,all_to_allvv2只实现了 HCCL 单后端,知识条目中的 "仅 HCCL" 即源于此。
2. 参数结构 AllToAllVV2Param 详解
算子的超参数定义在 include/atb/infer_op_params.h 的infer::AllToAllVV2Param结构体中,全部字段如下:
| 字段 | 类型 | 默认值 | 含义与取值说明 |
|---|---|---|---|
rank | int | -1 | 当前卡在通信域内的编号。-1表示未指定:单卡场景(rankSize == 1)时会被自动修正为 0;多卡场景必须显式指定,否则参数校验直接报错 |
rankSize | int | 0 | 参与通信的卡的数量,必须与 sendCounts/recvCounts/sdispls/rdispls 等数组的长度一致 |
rankRoot | int | 0 | 主通信编号(root rank),用于构建通信域时的根节点定位 |
backend | std::string | "hccl" | 通信计算类型,当前仅支持"hccl",传入其他值会被 ParamCheck 拒绝 |
hcclComm | HcclComm | nullptr | HCCL 通信域指针。默认为空,此时由加速库根据 rank/rankSize/rankTableFile 自行创建;若用户希望自行管理通信域,则传入该指针,加速库将直接复用 |
commMode | CommMode | COMM_MULTI_PROCESS | 通信模式枚举值。注意:HCCL 多线程场景只支持外部传入通信域的方式 |
rankTableFile | std::string | 空 | 集群信息配置文件路径,适用单机及多机通信场景,当前仅支持 hccl 后端。若单机配置了 rankTable,则以 rankTable 初始化通信域 |
commDomain | std::string | 空 | 通信 device 组用通信域名标识,多通信域场景使用,当前仅支持 hccl |
rsv[64] | uint8_t[64] | 全 0 | 预留字段 |
该结构体还提供了operator==用于参数相等性比较(比较 rank、rankSize、rankRoot、hcclComm、commMode、backend、rankTableFile、commDomain),供 ATB 的算子参数缓存/复用机制使用。
3. 输入输出张量契约:6 入 1 出
算子的张量接口由 ops_configs/atb_ops_info.ini 中的[AllToAllVV2Operation]配置定义,源码中通过IN_TENSOR_NUM = 6、OUT_TENSOR_NUM = 1(all_to_allvv2_operation.cpp)与之对应:
| 序号 | 名称 | dtype | format | 说明 |
|---|---|---|---|---|
| input0 | x | float16 / float / int8 / int32 / int16 / int64 / bf16 | nd | 本卡待发送的数据张量,device 侧 |
| input1 | sendCount | int64 | nd | 长度为rankSize的 host 数组,指定发给每个目标 rank 的元素个数 |
| input2 | sdispls | int64 | nd | 长度为rankSize的 host 数组,指定每个发送分片在输入张量中的起始偏移 |
| input3 | recvCount | int64 | nd | 长度为rankSize的 host 数组,指定从每个源 rank 接收的元素个数 |
| input4 | rdispls | int64 | nd | 长度为rankSize的 host 数组,指定每个接收分片在输出张量中的起始偏移 |
| input5 | tensorForInferShape | int8 | nd | 仅用于 shape 推导的辅助张量,其dims[0]必须等于sum(recvCounts),dtype 恒为 int8 |
| output0 | output | 与 input0 相同 | nd | 接收并汇聚后的数据张量,device 侧 |
几点关键设计:
- count/displs 全部走 host 内存:
sendCount、sdispls、recvCount、rdispls四个张量在 runner 中通过hostData访问(见 all_to_allvv2_hccl_runner.cpp),它们是纯 host 侧的控制元数据,不占用 device 内存。 tensorForInferShape专为 InferShape 服务:由于 count/displs 是 host 数据,在纯 shape 推导阶段无法可靠读取,因此额外引入一个 device 张量,使其dims[0]携带sum(recvCounts)的信息,从而推导出输出的第二维(详见下节)。- 输入数据 dtype 覆盖面广:支持 float16、float、int8、int32、int16、int64、bf16 共 7 种,输入输出 dtype 一一对应。
4. 算子实现剖析:校验、Shape 推导与 Runner 决策
算子的核心逻辑集中在 all_to_allvv2_operation.cpp,包含四条主线。
4.1 参数校验 ParamCheck
构造函数之外,文件内的匿名命名空间定义了ParamCheck(L24-L40):
backend必须等于"hccl",否则报错 "backend must be hccl";- 当
rank == -1且rankSize > 1时,报错 "multi-card, rank must be specified"——即多卡场景不允许省略 rank 号; - 随后调用
OperationUtil::DistributedInitCheck<AllToAllVV2Param>做分布式初始化相关检查。
4.2 Shape 推导 InferShapeImpl
InferShapeImpl 逻辑非常简洁:输出张量直接继承输入 0 的描述,但强制改写为二维[1, dims0(tensorForInferShape)],即输出 shape 为[1, sum(recvCounts)]。这从侧面说明该算子要求最终输出在逻辑上按"一行、sum(recvCounts) 列"理解,各 rank 的接收数据按rdispls指定的偏移落位到这一行中。
4.3 张量校验 SetupCheckImpl
SetupCheckImpl 承担运行时校验,可从源码逐条提取:
- 数组长度校验:
sendCounts、sdispls、recvCounts、rdispls的dims[0]都必须等于param_.rankSize; - 非负校验:四个数组中每个元素都必须
>= 0; - 求和溢出校验:累加
sum(recvCounts)过程中检查 int64 溢出,且要求sum(recvCounts) > 0; - 接收边界校验:对每个 rank,要求
recvCounts[i] + rdispls[i] <= sum(recvCounts),防止接收分片越出输出范围; - 发送边界校验:对每个 rank,要求
sendCounts[i] + sdispls[i] <= inputCount(输入张量总元素数,由辅助函数AllToAllVV2CalculateTensorSize计算); - 输出维度校验:
outTensors[0].shape.dims[1]必须等于sum(recvCounts)。
这些检查在 Setup 阶段(SetupCheckImpl)执行,把绝大多数越界、负数和长度不匹配的错误拦截在执行之前,避免把非法参数直接传给 HCCL。
4.4 Runner 决策 CreateRunner
CreateRunner 依据参数选择执行后端:
- 仅当
backend == "hccl"时返回AllToAllVV2HcclRunner; - 当
hcclComm == nullptr时,走"由加速库创建通信域"路径,构造参数为(param_, !param_.rankTableFile.empty()),即根据是否配置了rankTableFile决定使用 rankTable 初始化还是 rank/rankSize/rankRoot/commDomain 初始化; - 当外部传入
hcclComm时,走"复用用户通信域"路径,构造参数为(param_, hcclComm); - 其他 backend 一律返回空指针(配合 ParamCheck 的双重保障)。
此外,构造函数会根据硬件平台选择不同的算子 IR 配置(L52-L56):310P 平台使用AllToAllVV2Operation310p,其余平台使用AllToAllVV2Operation,两个配置在 ops_configs/atb_ops_info.ini 中均有定义(310p 变体仅支持 float16/int8)。
5. Runner 与底层 HCCL 调用链
AllToAllVV2HcclRunner继承自HcclRunner(all_to_allvv2_hccl_runner.h),真正的执行发生在 ExecuteImpl,其调用链为:
AllToAllVV2Operation::Execute └─> AllToAllVV2HcclRunner::ExecuteImpl ├─ 校验 hcclComm_ 非空 ├─ 校验 input0.deviceData / output0.deviceData 非空 └─ HcclAlltoAllV( input0.deviceData, // sendbuf input1.hostData, // sendCounts input2.hostData, // sdispls GetHcclDtype(inDtype), // 输入 dtype output0.deviceData, // recvbuf input3.hostData, // recvCounts input4.hostData, // rdispls GetHcclDtype(outDtype), // 输出 dtype hcclComm_.get(), // 通信域 GetExecuteStream(context) // 执行流 )几个实现要点:
- 语义映射:ATB 层的
sendCount/sdispls/recvCount/rdispls与 HCCLHcclAlltoAllV的参数一一对应,ATB 在此处的职责是参数组织、校验与通信域管理,不涉及 kernel 计算,因此该算子没有独立 kernel 目录(知识条目中 "operation + hccl_runner" 的结构即源于此)。 - 错误处理:
HcclAlltoAllV返回非HCCL_SUCCESS时,通过ConvertHcclResultToStatus转换为 ATB 的Status并记录日志;hcclComm_为空或 device 张量为空时分别返回ERROR_INTERNAL_ERROR/ERROR_INVALID_PARAM。 - 注册机制:文件末尾通过
REG_RUNNER_TYPE(AllToAllVV2HcclRunner)将 Runner 注册进 ATB 的 Runner 工厂(all_to_allvv2_hccl_runner.cpp),供CreateRunner统一创建。
6. 通信域初始化三种方式与适用场景
由AllToAllVV2HcclRunner的三个构造函数可见,通信域有三种初始化途径:
| 方式 | 触发条件 | 初始化依据 | 适用场景 |
|---|---|---|---|
| rankTable 文件 | hcclComm == nullptr且rankTableFile非空 | HcclRunner(name, rank, rankTableFile, commDomain) | 单机/多机集群,按集群信息文件建域 |
| rank 参数组 | hcclComm == nullptr且rankTableFile为空 | HcclRunner(name, rank, rankSize, rankRoot, commDomain) | 单机多卡,按 rank/rankSize/rankRoot 建域 |
| 外部通信域 | hcclComm != nullptr | 直接复用用户传入的HcclComm | 用户自行管理通信域、多线程等高级场景 |
其中第三种方式对应commMode字段的说明:"hccl 多线程只支持外部传入通信域方式"。日常单卡调试时,rank == -1且rankSize == 1会被SetParam自动修正为rank = 0(all_to_allvv2_operation.cpp),无需显式指定。
7. 使用注意事项与参数速查
综合源码校验与配置,实际使用all_to_allvv2时需满足以下约束:
- 平台约束:参数头文件注明该算子当前仅支持 Atlas 800I A2 推理产品(310P 走独立的
AllToAllVV2Operation310p配置,且仅支持 float16/int8); - 后端约束:
backend仅支持"hccl",不支持 LCCL; - 多卡必须指定 rank:
rankSize > 1时rank不能为-1; - 数组契约:
sendCount/sdispls/recvCount/rdispls的长度必须等于rankSize,元素非负,且满足sendCounts[i] + sdispls[i] <= 输入总元素数、recvCounts[i] + rdispls[i] <= sum(recvCounts); - 输出形状:输出固定为二维
[1, sum(recvCounts)],tensorForInferShape的dims[0]必须与之相等; - dtype 对齐:输入输出 dtype 一致,支持 float16/float/int8/int32/int16/int64/bf16(310P 仅 float16/int8)。
8. 与 all_to_all / all_to_allv 的横向对比
仓库中同属 communication 分类的另外两个算子可作为对照(详见 .agent/knowledge/ops/communication/all_to_all/index.md 与 .agent/knowledge/ops/communication/all_to_allv/index.md):
| 维度 | all_to_all | all_to_allv | all_to_allvv2 |
|---|---|---|---|
| 长度语义 | 等长交换 | 可变长度 | 可变长度(增强版) |
| 后端 | HCCL / LCCL 双后端 | 仅 HCCL | 仅 HCCL |
| 源码文件数 | 6 | 4 | 4 |
| Runner | AllToAllHcclRunner / AllToAllLcclRunner | AllToAllvHcclRunner | AllToAllVV2HcclRunner |
| 输入数量 | 较少(等长,无需 count/displs 数组) | 含 count/displs 控制数组 | 6 输入(含 tensorForInferShape 辅助张量) |
可见all_to_allvv2面向"可变长度 + 精确位移控制 + 强校验 + 仅 HCCL"的增强场景,是三者中最精细的版本。若需要阅读更宏观的 Agent 知识索引,可参考 .agent/knowledge/README.md 与路由文件 .agent/knowledge/routing/all_to_allvv2.md。
【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库,基于华为Ascend AI处理器,提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考