CANN ops-transformer 通算融合算子 aclnnWeightQuantMatmulAllReduce 完全指南:权重伪量化 Matmul + AllReduce 融合计算
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
本文是 CANN ops-transformer 开源仓库中aclnnWeightQuantMatmulAllReduce算子的实战技术指南。该算子位于 mc2/matmul_all_reduce 模块,属于 MC2(Matmul + Communication 融合)通算融合算子族:它在一次算子执行中完成"权重伪量化(anti-quant)→ MatMul → 加 bias/x3 → AllReduce 通信"的整条链路,可用于大模型全量推理/训练场景下的分布式线性层前向计算。读完本文,你将掌握该算子的计算语义、两段式 aclnn 接口的完整参数约束、不同产品(Ascend 950 系列与 Atlas A2 系列)上的数据类型与卡数支持差异、per-tensor/per-channel/per-group 三种伪量化模式的正确用法,以及基于 调用示例 编写可运行多卡样例的完整方法。
算子功能与计算语义
核心功能
aclnnWeightQuantMatmulAllReduce是"权重量化感知的 Matmul + AllReduce 融合算子":对入参x2(通常是量化后的权重矩阵)先做伪量化(anti-quant)反量化还原,再与x1做矩阵乘,叠加可选bias与x3,最后对结果做 AllReduce 集合通信。
它支持pertensor、perchannel、pergroup三种伪量化方式,从源码看,反量化类型通过 tiling 阶段的antiQuantType_字段区分(见 weight_quant_matmul_all_reduce_tiling_950.cpp)。
计算公式
$$ output = AllReduce(x1 @ ((x2 + antiquantOffset) * antiquantScale) + bias + x3) $$
各符号含义:
| 符号 | 含义 |
|---|---|
x1 | MatMul 左矩阵(激活侧,不量化),BFLOAT16 / FLOAT16 |
x2 | MatMul 右矩阵(权重侧,量化存储),INT8 / INT4,或 950 系列上的 FLOAT8_E4M3FN / HIFLOAT8 |
antiquantOffset | 伪量化 offset,可空;x2为 FLOAT8 类数据类型时须为空指针 |
antiquantScale | 伪量化 scale,必填,与x2逐元素配合完成(x2 + offset) * scale反量化 |
bias | 偏置,可空,一维,长度与 output 最后一维相等 |
x3 | MatMul 后的残差加项,可空,shape 与 output 一致 |
output | MatMul + AllReduce 融合后的结果 |
该公式与仓库 README.md 中描述的非量化融合场景(output = Allreduce(x1 @ ((x2 + antiquantOffset) * antiquantScale) + bias + x3))完全一致,是 MC2 算子族中"权重伪量化 + 通算融合"的代表实现。
产品支持情况与版本要求
支持矩阵
| 产品 | 是否支持 |
|---|---|
| Ascend 950PR / Ascend 950DT | 支持 |
| Atlas A3 训练系列产品 / Atlas A3 推理系列产品 | 不支持 |
| Atlas A2 训练系列产品 / Atlas A2 推理系列产品 | 支持 |
| Atlas 200I/500 A2 推理产品 | 不支持 |
| Atlas 推理系列产品 | 不支持 |
| Atlas 训练系列产品 | 不支持 |
从源码CheckDtypeValid的实现可以看出,不同架构的支撑差异会直接影响数据类型校验规则(见 aclnn_weight_quant_matmul_all_reduce.cpp):
- Atlas A2(910B 架构):
x1/scale/offset/x3/output支持 FLOAT16 与 BFLOAT16,x2仅支持 INT8/INT4; - Ascend 950(DAV_3510 架构):
x2额外支持 FLOAT8_E4M3FN 与 HIFLOAT8(对应源码中的dtypeSupportListQuantA5); - Atlas 310P(DAV_2002 架构):仅支持 FLOAT16 一种非量化侧数据类型(对应
DTYPE_SUPPORT_LIST_310P)。
版本要求
使用该接口时,请确保驱动固件包和 CANN 包都为配套的 8.0.RC2 版本或配套的更高版本,否则将引发报错,例如 BUS ERROR 等硬件级异常。
两段式接口与函数原型
该算子遵循 CANN aclnn 的两段式接口规范:必须先调用aclnnWeightQuantMatmulAllReduceGetWorkspaceSize获取计算所需的 workspace 大小以及包含算子计算流程的执行器(executor),再调用aclnnWeightQuantMatmulAllReduce执行计算。第二段接口不可重复调用。
第一段接口:获取 workspace 大小与执行器
aclnnStatus aclnnWeightQuantMatmulAllReduceGetWorkspaceSize( const aclTensor *x1, const aclTensor *x2, const aclTensor *bias, const aclTensor *antiquantScale, const aclTensor *antiquantOffset, const aclTensor *x3, const char *group, const char *reduceOp, int64_t commTurn, int64_t streamMode, int64_t antiquantGroupSize, const aclTensor *output, uint64_t *workspaceSize, aclOpExecutor **executor)第二段接口:执行计算
aclnnStatus aclnnWeightQuantMatmulAllReduce( void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, const aclrtStream stream)从源码看,第一段接口在完成参数校验后,会调用内部通用入口aclnnInnerMatmulAllReduceGetWorkspaceSize(见 aclnn_weight_quant_matmul_all_reduce.cpp)完成 workspace 计算与执行器创建;对于可选的bias、antiquantOffset、x3等入参,还会通过NnopbaseDisableOptionalInput在 IR 层标记为可选输入。第二段接口在 950 架构上会额外设置 HCCL 服务器类型为 AICPU 模式(见同文件第 435-439 行)。
aclnnWeightQuantMatmulAllReduceGetWorkspaceSize 参数详解
参数说明
| 参数名 | 输入/输出 | 描述 | 使用说明 | 数据类型 | 数据格式 | 维度(shape) | 非连续 Tensor |
|---|---|---|---|---|---|---|---|
| x1 | 输入 | MatMul 计算的左矩阵,即计算公式中的 x1 | 当前版本仅支持二维或者三维输入;支持不转置场景 | BFLOAT16、FLOAT16 | 参见约束说明 | 2-3 | × |
| x2 | 输入 | MatMul 计算的右矩阵,即计算公式中的 x2 | 当前版本仅支持二维输入;支持转置/不转置场景;ND 格式下支持最后两轴转置情况下的非连续 tensor,其他非连续 tensor 不支持 | 参见约束说明 | ND、FRACTAL_NZ | 2 | × |
| bias | 输入 | 对应计算公式中 bias 偏移,即计算公式中的 bias | 支持传入空指针,非空时当前版本仅支持一维输入 | 参见约束说明 | ND | 1 | √ |
| antiquantScale | 输入 | 即计算公式中的 antiquantScale | pertensor 场景 shape 为 (1);perchannel 场景 shape 为 (n)/(1,n),n 为 x2 最后一维的大小;pergroup 场景 shape 为 (ceil(k,antiquantGroupSize),n) | BFLOAT16、FLOAT16 | ND | 1-2 | √ |
| antiquantOffset | 输入 | 对 x2 进行伪量化计算的 offset 参数,即计算公式中的 antiquantOffset | 支持传入空指针,非空时 shape 与 antiquantScale 一致;当 x2 的数据格式为 FLOAT8_E4M3FN 或者 HIFLOAT8 时,不支持该参数,填空指针 | BFLOAT16、FLOAT16 | ND | 1-2 | √ |
| x3 | 输入 | MatMul 计算后的 add 计算,即计算公式中的 x3 | 支持传入空指针,非空时 shape 与 mm 计算后的 shape 相同 | 参见约束说明 | ND | 2-3 | √ |
| group | 输入 | 通信域名称 | 通过 Hccl 提供的接口extern HcclResult HcclGetCommName(HcclComm comm, char* commName);获取,其中 commName 即为 group | String | - | - | - |
| reduceOp | 输入 | reduce 操作类型 | 当前版本仅支持输入"sum" | String | - | - | - |
| commTurn | 输入 | 通信数据切分数,即总数据量/单次通信量 | 当前版本仅支持输入 0 | INT64 | - | - | - |
| streamMode | 输入 | 流模式的枚举 | 当前版本仅支持枚举值 1 | INT64 | - | - | - |
| antiquantGroupSize | 输入 | 伪量化 pergroup 模式下,对 x2 进行反量化计算的 groupSize 输入 | pergroup 量化场景下需传入该参数,传入值的范围为 [32, min(k-1,INT_MAX)],且为 32 的倍数;k 取值范围与 mm 接口保持一致,为 [1,65535];非 pergroup 量化场景下仅支持传入 0 | INT64 | - | - | - |
| output | 输出 | MatMul 计算与 AllReduce 通信的结果,即计算公式中的 output | output 的维度与 x1 一致 | - | ND | 2-3 | √ |
| workspaceSize | 输出 | 返回需要在 Device 侧申请的 workspace 大小 | - | - | - | - | - |
| executor | 输出 | 返回 op 执行器,包含了算子计算流程 | - | - | - | - | - |
关键属性的源码级说明
- group / reduceOp / streamMode / antiquantGroupSize 的校验:源码
CheckAttr(见 aclnn_weight_quant_matmul_all_reduce.cpp)中,reduceOp必须为"sum"(源码常量REDUCE_OP_SUM),streamMode必须为 1,antiquantGroupSize为 0 时视为非 pergroup 场景;非 0 时要求% 32 == 0且落在[32, min(k-1, INT32_MAX)]。当 k 为 0(空 tensor 场景)时跳过 groupSize 校验。 - antiquantScale 的 shape 校验:源码
IsAntiquantScaleShapeValid(见同文件第 148-173 行)进一步实现文档描述的三种量化模式 shape 规则:pertensor 为(1);perchannel 为(n)或(1,n);pergroup 为(ceil(k, groupSize), n)。 - 连续性与转置约束:源码
CheckParams在 910B 架构上还会校验x2、antiquantScale、antiquantOffset在非转置场景下的连续性(第 335-343 行),并检查 x2 转置时 scale/offset 是否与之一致(CheckContiguous,第 276-314 行),与文档"pergroup 场景下 x2 转置时,antiquantScale 和 antiquantOffset 需要一起转置,保持连续性"的约束呼应。
返回值与错误码
两个接口均返回aclnnStatus状态码,具体取值参见 aclnn 返回码。
第一段接口完成入参校验,出现以下场景报错:
| 返回值 | 错误码 | 描述 |
|---|---|---|
| ACLNN_ERR_PARAM_NULLPTR | 161001 | 传入的 x1、x2、antiquantScale 或 output 是空指针 |
| ACLNN_ERR_PARAM_INVALID | 161002 | x1、x2、bias、antiquantScale、antiquantOffset、x3 或 output 的数据类型不符合要求 |
| ACLNN_ERR_PARAM_INVALID | 161002 | reduceOp、streamMode、antiquantGroupSize 不在合法范围内 |
| ACLNN_ERR_PARAM_INVALID | 161002 | x1、x2、bias、antiquantScale、antiquantOffset、x3、output、antiquantGroupSize 的 shape 不符合约束要求 |
在源码中,这三类错误分别由CheckNotNull(返回ACLNN_ERR_PARAM_NULLPTR)、CheckDtypeValid、CheckAttr、CheckShape(返回ACLNN_ERR_PARAM_INVALID)按顺序完成检查(见 aclnn_weight_quant_matmul_all_reduce.cpp)。
约束说明
确定性计算
- Atlas A2 训练/推理系列产品(910B):
aclnnWeightQuantMatmulAllReduce默认非确定性实现,支持通过配置HCCL_DETERMINISTIC环境变量为 true 开启确定性计算。 - Ascend 950PR / Ascend 950DT:默认确定性实现。
形状与范围约束
- 增量场景不开启 MC2,全量场景开启 MC2。
- 输入 x1 可为二维或者三维,其 shape 为
(b, s, k)或者(m, k)。 - x2 必须是二维,其 shape 为
(k, n),k 轴满足 mm 算子入参要求(k 轴相等),m 的范围为[1, 2147483647],k、n 的范围为[1, 65535]。 - 传入的 x1、x2、antiquantScale 或者 output 不为空指针。
- 当输入 x1 的 shape 为
(b, s, k)时,x3(非空场景)与输出 output 的 shape 为(b, s, n);当输入 x1 的 shape 为(m, k)时,x3(非空场景)与输出 output 的 shape 为(m, n)。 - bias 若非空,shape 大小与 output 最后一维大小相等。antiquantScale 在 pertensor 场景下 shape 为
(1),在 perchannel 场景下 shape 为(1,n)/(n),在 pergroup 场景 shape 为(ceil(k,antiquantGroupSize), n)。antiquantOffset 若非空,其 shape 与 antiquantScale 一致。 - x1 和 x2,x3(非空场景)、antiquantScale、antiquantOffset(非空场景)、output、bias(非空场景)的数据类型和数据格式需要在支持的范围之内。
- x1、antiquantScale、antiquantOffset(非空场景)、x3(非空场景)、bias(非空场景)、output 的数据类型相同。antiquantGroupSize 取值满足取值范围且为 32 的倍数。
- pergroup 场景下,x2 转置时,antiquantScale 和 antiquantOffset 需要一起转置,保持连续性。
- 在长序列场景,随着 b/s 或者 m 的增大,可能出现 OOM 或者计算超时。
组网与卡数约束
仅支持 hccs 链路 all mesh 组网:
- Atlas A2 训练/推理系列产品(910B):支持 1、2、4、8 卡。
- Ascend 950PR / Ascend 950DT:支持 1、2、4、8、16、32、64 卡。
产品相关的格式与对齐约束
Atlas A2 训练/推理系列产品(910B):
- 一个模型中的通算融合 MC2 算子,仅支持相同通信域。
- 输入 x2 的数据格式支持 ND(当前版本仅支持二维输入)和 FRACTAL_NZ 格式(当前版本仅支持四维输入)。当 x2 的数据格式为 FRACTAL_NZ 时,配合
aclnnCalculateMatmulWeightSizeV2和aclnnTransMatmulWeight完成输入 ND 到 NZ 的转换,非连续的 tensor 仅支持 transpose 场景。
Ascend 950PR / Ascend 950DT:
- 输入 x2 的数据格式支持 ND(仅支持 2D 输入)。当前版本,当数据类型为 INT8 时,要求 N、K 为 32 对齐;当数据类型为 INT4 时,要求 N、K 为 64 对齐。
空 tensor 支持度
仅支持 k 为 0 的场景,此时输出为bias + x3;不支持 bs/m/n 为 0 的空 tensor 输入。该逻辑与源码CheckAttr中"kLen == 0 时跳过 antiquantGroupSize 校验"的分支相互印证,也与 UT 用例empty_K(k 为 0 时返回 SUCCESS)和empty_M(m 为 0 时返回 PARAM_INVALID)一一对应。
输入输出数据类型组合
Atlas A2 训练/推理系列产品(910B)
| x1 | x2 | bias | antiquantScale | antiquantOffset | x3 | output | 限制 |
|---|---|---|---|---|---|---|---|
| BFLOAT16 | INT8、INT4 | null、BFLOAT16 | BFLOAT16 | null、BFLOAT16 | null、BFLOAT16 | BFLOAT16 | - |
| FLOAT16 | INT8、INT4 | null、FLOAT16 | FLOAT16 | null、FLOAT16 | null、FLOAT16 | FLOAT16 | - |
Ascend 950PR / Ascend 950DT
| x1 | x2 | bias | antiquantScale | antiquantOffset | x3 | output | 限制 |
|---|---|---|---|---|---|---|---|
| BFLOAT16 | INT8、INT4 | null、BFLOAT16 | BFLOAT16 | null、BFLOAT16 | null、BFLOAT16 | BFLOAT16 | 支持 pertensor、perchannel、pergroup 量化场景 |
| BFLOAT16 | FLOAT8_E4M3FN、HIFLOAT8 | null、BFLOAT16 | BFLOAT16 | null、BFLOAT16 | null、BFLOAT16 | BFLOAT16 | 仅支持 perchannel 量化场景 |
| FLOAT16 | INT8、INT4 | null、FLOAT16 | FLOAT16 | null、FLOAT16 | null、FLOAT16 | FLOAT16 | 支持 pertensor、perchannel、pergroup 量化场景 |
| FLOAT16 | FLOAT8_E4M3FN、HIFLOAT8 | null、FLOAT16 | FLOAT16 | null、FLOAT16 | null、FLOAT16 | FLOAT16 | 仅支持 perchannel 量化场景 |
注意:当 x2 为 FLOAT8_E4M3FN 或 HIFLOAT8 时,antiquantOffset 不支持(须传空指针),且仅支持 perchannel 量化;当 x2 为 INT8/INT4 时三种量化方式均支持。从 tiling 源码看,x2 为 FRACTAL_NZ 格式时仅支持 perchannel 反量化(见 weight_quant_matmul_all_reduce_tiling_950.cpp)。
多卡调用示例(C++)
示例代码如下,仅供参考,具体编译和执行过程请参考仓库内编译与运行样例。本示例调用了部分 HCCL 集合通信库接口:HcclGetCommName、HcclCommInitAll、HcclCommDestroy。代码对应(m, k) @ (k, n) + bias + x3 → AllReduce的 FLOAT16 + INT8 场景,其中antiquantGroupSize = 0表示非 pergroup(示例为 perchannel,scale 与 offset shape 均为 (n))。
#include <iostream> #include <vector> #include <thread> #include <string.h> #include "hccl/hccl.h" #include "aclnn/opdev/fp16_t.h" #include "aclnnop/aclnn_weight_quant_matmul_all_reduce.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::vector<int64_t> &shape) { int64_t shapeSize = 1; for (auto i: shape) { shapeSize *= i; } return shapeSize; } 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; } struct Args { uint32_t rankId; HcclComm hcclComm; aclrtStream stream; aclrtContext context; }; int launchOneThreadweightQuantmatmulAllReduce(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::vector<int64_t> x1Shape = {32, 64}; std::vector<int64_t> x2Shape = {64, 128}; std::vector<int64_t> biasShape = {128}; std::vector<int64_t> antiquantScaleShape = {128}; std::vector<int64_t> antiquantOffsetShape = {128}; std::vector<int64_t> x3Shape = {32, 128}; std::vector<int64_t> outShape = {32, 128}; void *x1DeviceAddr = nullptr; void *x2DeviceAddr = nullptr; void *biasDeviceAddr = nullptr; void *antiquantScaleDeviceAddr = nullptr; void *antiquantOffsetDeviceAddr = nullptr; void *x3DeviceAddr = nullptr; void *outDeviceAddr = nullptr; aclTensor *x1 = nullptr; aclTensor *x2 = nullptr; aclTensor *bias = nullptr; aclTensor *antiquantScale = nullptr; aclTensor *antiquantOffset = nullptr; aclTensor *x3 = nullptr; aclTensor *out = nullptr; int64_t commTurn = 0; int64_t streamMode = 1; int64_t antiquantGroupSize = 0; 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 antiquantScaleShapeSize = GetShapeSize(antiquantScaleShape); long long antiquantOffsetShapeSize = GetShapeSize(antiquantOffsetShape); long long x3ShapeSize = GetShapeSize(x3Shape); long long outShapeSize = GetShapeSize(outShape); std::vector<op::fp16_t> x1HostData(x1ShapeSize, 1); std::vector<int8_t> x2HostData(x2ShapeSize, 1); std::vector<op::fp16_t> biasHostData(biasShapeSize, 1); std::vector<op::fp16_t> antiquantScaleHostData(antiquantScaleShapeSize, 1); std::vector<op::fp16_t> antiquantOffsetHostData(antiquantOffsetShapeSize, 1); std::vector<op::fp16_t> x3HostData(x3ShapeSize, 1); std::vector<op::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_INT8, &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(antiquantScaleHostData, antiquantScaleShape, &antiquantScaleDeviceAddr, aclDataType::ACL_FLOAT16, &antiquantScale); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(antiquantOffsetHostData, antiquantOffsetShape, &antiquantOffsetDeviceAddr, aclDataType::ACL_FLOAT16, &antiquantOffset); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(x3HostData, x3Shape, &x3DeviceAddr, aclDataType::ACL_FLOAT16, &x3); CHECK_RET(ret == ACL_SUCCESS, return ret); ret = CreateAclTensor(outHostData, outShape, &outDeviceAddr, aclDataType::ACL_FLOAT16, &out); CHECK_RET(ret == ACL_SUCCESS, return ret); // 调用第一段接口 ret = aclnnWeightQuantMatmulAllReduceGetWorkspaceSize(x1, x2, bias, antiquantScale, antiquantOffset, x3, hcom_name, "sum", commTurn, streamMode, antiquantGroupSize, out, &workspaceSize, &executor); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnWeightQuantMatmulAllReduceGetWorkspaceSize 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 = aclnnWeightQuantMatmulAllReduce(workspaceAddr, workspaceSize, executor, args.stream); CHECK_RET(ret == ACL_SUCCESS, LOG_PRINT("aclnnWeightQuantMatmulAllReduce 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 aclnnWeightQuantMatmulAllReduce execute success \n", args.rankId); // 释放device资源,需要根据具体API的接口定义修改 if (x1 != nullptr) { aclDestroyTensor(x1); } if (x2 != nullptr) { aclDestroyTensor(x2); } if (bias != nullptr) { aclDestroyTensor(bias); } if (antiquantScale != nullptr) { aclDestroyTensor(antiquantScale); } if (antiquantOffset != nullptr) { aclDestroyTensor(antiquantOffset); } if (x3 != nullptr) { aclDestroyTensor(x3); } if (out != nullptr) { aclDestroyTensor(out); } if (x1DeviceAddr != nullptr) { aclrtFree(x1DeviceAddr); } if (x2DeviceAddr != nullptr) { aclrtFree(x2DeviceAddr); } if (biasDeviceAddr != nullptr) { aclrtFree(biasDeviceAddr); } if (antiquantScaleDeviceAddr != nullptr) { aclrtFree(antiquantScaleDeviceAddr); } if (antiquantOffsetDeviceAddr != nullptr) { aclrtFree(antiquantOffsetDeviceAddr); } if (x3DeviceAddr != nullptr) { aclrtFree(x3DeviceAddr); } 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[]) { 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::vector<std::unique_ptr<std::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(&launchOneThreadweightQuantmatmulAllReduce, std::ref(args[rankId]))); } for (uint32_t rankId = 0; rankId < ndev; rankId++) { threads[rankId]->join(); } aclFinalize(); return 0; }示例要点拆解
- 环境初始化:
aclInit→ 逐卡aclrtSetDevice→HcclCommInitAll(ndev, devices, comms)创建通信域 → 每卡创建 context 与 stream; - 获取通信域名称:在线程内通过
HcclGetCommName(args.hcclComm, hcom_name)拿到group字符串; - 构造 aclTensor:
CreateAclTensor内部完成aclrtMalloc、aclrtMemcpy(HOST→DEVICE)与aclCreateTensor(ND 格式、连续 strides); - 两段式调用:先
GetWorkspaceSize获取 workspace 与 executor,workspaceSize > 0时aclrtMalloc申请,再执行第二段接口并在固定位置aclrtSynchronizeStreamWithTimeout同步等待; - 资源回收:依次销毁 tensor、device 内存、workspace、stream、HcclComm、context,并
aclrtResetDevice,最后aclFinalize。
仓库中的验证与配套资源
单元测试
仓库在 tests/ut/op_api 下提供了针对本接口的完整 UT:
- test_aclnn_weight_quant_matmul_all_reduce.cpp:通过参数化测试 + 若干专项用例,覆盖 NZ 格式权重、310P 预转置权重、950 INT4 权重、scale 与转置 x2 连续性不匹配、非连续 x2、非法 x3 数据类型、非法 pertensor scale shape 等边界场景;
- test_aclnn_weight_quant_matmul_all_reduce.csv:以 CSV 形式给出了 26 组入参-期望结果组合,可以直接对照理解各参数的合法取值范围,例如:
common_1:x1(32,64) FLOAT16 / x2(64,128) INT8 / scale(128) / output(32,128) → SUCCESS;pergroup_quant:group_size=32、scale shape(2,128)(即ceil(64,32)=2)→ SUCCESS;pertensor_quant:scale shape(1)→ SUCCESS;invalid_pergroup_size:group_size=16(小于 32)→ PARAM_INVALID;invalid_pergroup_size_not_multiple:group_size=24(非 32 倍数)→ PARAM_INVALID;empty_K:x1(32,0)、x2(0,128) → SUCCESS(k 为 0 的空 tensor 场景);empty_M:x1(0,64) → PARAM_INVALID(bs/m 为 0 不支持)。
Golden 验证脚本
tests/assets/impl/golden.py 提供了多卡 golden 参考实现,其中针对 WeightQuant 变体按(x2 + offset) * scale完成权重反量化,并将 pergroup 的 scale/offset 通过repeat_interleave(group_size, dim=0)展开为逐元素系数后参与matmul + all_reduce(SUM)的浮点参考计算,可用于校验算子数值结果。
工程结构速览
| 目录/文件 | 作用 |
|---|---|
| op_api/aclnn_weight_quant_matmul_all_reduce.cpp | aclnn 两段式接口实现与入参校验 |
| op_api/aclnn_weight_quant_matmul_all_reduce.h | 对外头文件(接口声明与文档注释) |
| op_kernel/weight_quant_matmul_all_reduce_tiling_data.h | tiling 数据结构(含 tile/tail 两段 MatMul 切分) |
| op_host/op_tiling | 分架构 tiling 实现(arch22/arch31/arch35) |
| docs/aclnnWeightQuantMatmulAllReduceV2.md | V2 版本(新增 commMode 通信引擎参数) |
| README.md | 模块总览与 MC2 算子族计算公式全集 |
常见问题与排查建议
- BUS ERROR 等硬件异常:优先确认驱动固件包与 CANN 包均为 8.0.RC2 或更高配套版本;
- 返回 161002(PARAM_INVALID):按上文错误码表逐项核对——
reduceOp必须为"sum"、streamMode必须为 1、commTurn必须为 0、antiquantGroupSize为 0 或[32, min(k-1,INT_MAX)]内的 32 倍数,且 x1 与 x2 的 k 轴必须相等; - pergroup 场景 scale 维度不匹配:antiquantScale/antiquantOffset 的 shape 必须为
(ceil(k, antiquantGroupSize), n),且 x2 转置时二者须随 x2 一起转置保持连续; - 多卡通信异常:仅支持 hccs 链路 all mesh 组网,910B 上限 8 卡、950 系列上限 64 卡;910B 上一个模型内的 MC2 算子只能使用相同通信域;
- 长序列 OOM / 超时:随 b/s 或 m 增大可能出现资源问题,需关注序列长度与显存/算力预算。
综上,aclnnWeightQuantMatmulAllReduce是 CANN ops-transformer MC2 算子族中面向"量化权重 + 分布式全量推理/训练"场景的融合算子:理解其伪量化语义、两段式接口的严格参数约束以及不同架构的差异化能力,是在 Atlas A2 与 Ascend 950 系列硬件上正确落地通算融合计算的关键。
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考