CANN ops-transformer 通算融合算子 aclnnWeightQuantMatmulAllReduce 完全指南:权重伪量化 Matmul + AllReduce 融合计算
2026/9/20 1:16:05 网站建设 项目流程

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做矩阵乘,叠加可选biasx3,最后对结果做 AllReduce 集合通信。

它支持pertensor、perchannel、pergroup三种伪量化方式,从源码看,反量化类型通过 tiling 阶段的antiQuantType_字段区分(见 weight_quant_matmul_all_reduce_tiling_950.cpp)。

计算公式

$$ output = AllReduce(x1 @ ((x2 + antiquantOffset) * antiquantScale) + bias + x3) $$

各符号含义:

符号含义
x1MatMul 左矩阵(激活侧,不量化),BFLOAT16 / FLOAT16
x2MatMul 右矩阵(权重侧,量化存储),INT8 / INT4,或 950 系列上的 FLOAT8_E4M3FN / HIFLOAT8
antiquantOffset伪量化 offset,可空;x2为 FLOAT8 类数据类型时须为空指针
antiquantScale伪量化 scale,必填,与x2逐元素配合完成(x2 + offset) * scale反量化
bias偏置,可空,一维,长度与 output 最后一维相等
x3MatMul 后的残差加项,可空,shape 与 output 一致
outputMatMul + 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 计算与执行器创建;对于可选的biasantiquantOffsetx3等入参,还会通过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_NZ2×
bias输入对应计算公式中 bias 偏移,即计算公式中的 bias支持传入空指针,非空时当前版本仅支持一维输入参见约束说明ND1
antiquantScale输入即计算公式中的 antiquantScalepertensor 场景 shape 为 (1);perchannel 场景 shape 为 (n)/(1,n),n 为 x2 最后一维的大小;pergroup 场景 shape 为 (ceil(k,antiquantGroupSize),n)BFLOAT16、FLOAT16ND1-2
antiquantOffset输入对 x2 进行伪量化计算的 offset 参数,即计算公式中的 antiquantOffset支持传入空指针,非空时 shape 与 antiquantScale 一致;当 x2 的数据格式为 FLOAT8_E4M3FN 或者 HIFLOAT8 时,不支持该参数,填空指针BFLOAT16、FLOAT16ND1-2
x3输入MatMul 计算后的 add 计算,即计算公式中的 x3支持传入空指针,非空时 shape 与 mm 计算后的 shape 相同参见约束说明ND2-3
group输入通信域名称通过 Hccl 提供的接口extern HcclResult HcclGetCommName(HcclComm comm, char* commName);获取,其中 commName 即为 groupString---
reduceOp输入reduce 操作类型当前版本仅支持输入"sum"String---
commTurn输入通信数据切分数,即总数据量/单次通信量当前版本仅支持输入 0INT64---
streamMode输入流模式的枚举当前版本仅支持枚举值 1INT64---
antiquantGroupSize输入伪量化 pergroup 模式下,对 x2 进行反量化计算的 groupSize 输入pergroup 量化场景下需传入该参数,传入值的范围为 [32, min(k-1,INT_MAX)],且为 32 的倍数;k 取值范围与 mm 接口保持一致,为 [1,65535];非 pergroup 量化场景下仅支持传入 0INT64---
output输出MatMul 计算与 AllReduce 通信的结果,即计算公式中的 outputoutput 的维度与 x1 一致-ND2-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 架构上还会校验x2antiquantScaleantiquantOffset在非转置场景下的连续性(第 335-343 行),并检查 x2 转置时 scale/offset 是否与之一致(CheckContiguous,第 276-314 行),与文档"pergroup 场景下 x2 转置时,antiquantScale 和 antiquantOffset 需要一起转置,保持连续性"的约束呼应。

返回值与错误码

两个接口均返回aclnnStatus状态码,具体取值参见 aclnn 返回码。

第一段接口完成入参校验,出现以下场景报错:

返回值错误码描述
ACLNN_ERR_PARAM_NULLPTR161001传入的 x1、x2、antiquantScale 或 output 是空指针
ACLNN_ERR_PARAM_INVALID161002x1、x2、bias、antiquantScale、antiquantOffset、x3 或 output 的数据类型不符合要求
ACLNN_ERR_PARAM_INVALID161002reduceOp、streamMode、antiquantGroupSize 不在合法范围内
ACLNN_ERR_PARAM_INVALID161002x1、x2、bias、antiquantScale、antiquantOffset、x3、output、antiquantGroupSize 的 shape 不符合约束要求

在源码中,这三类错误分别由CheckNotNull(返回ACLNN_ERR_PARAM_NULLPTR)、CheckDtypeValidCheckAttrCheckShape(返回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 时,配合aclnnCalculateMatmulWeightSizeV2aclnnTransMatmulWeight完成输入 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)

x1x2biasantiquantScaleantiquantOffsetx3output限制
BFLOAT16INT8、INT4null、BFLOAT16BFLOAT16null、BFLOAT16null、BFLOAT16BFLOAT16-
FLOAT16INT8、INT4null、FLOAT16FLOAT16null、FLOAT16null、FLOAT16FLOAT16-

Ascend 950PR / Ascend 950DT

x1x2biasantiquantScaleantiquantOffsetx3output限制
BFLOAT16INT8、INT4null、BFLOAT16BFLOAT16null、BFLOAT16null、BFLOAT16BFLOAT16支持 pertensor、perchannel、pergroup 量化场景
BFLOAT16FLOAT8_E4M3FN、HIFLOAT8null、BFLOAT16BFLOAT16null、BFLOAT16null、BFLOAT16BFLOAT16仅支持 perchannel 量化场景
FLOAT16INT8、INT4null、FLOAT16FLOAT16null、FLOAT16null、FLOAT16FLOAT16支持 pertensor、perchannel、pergroup 量化场景
FLOAT16FLOAT8_E4M3FN、HIFLOAT8null、FLOAT16FLOAT16null、FLOAT16null、FLOAT16FLOAT16仅支持 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 集合通信库接口:HcclGetCommNameHcclCommInitAllHcclCommDestroy。代码对应(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; }

示例要点拆解

  1. 环境初始化aclInit→ 逐卡aclrtSetDeviceHcclCommInitAll(ndev, devices, comms)创建通信域 → 每卡创建 context 与 stream;
  2. 获取通信域名称:在线程内通过HcclGetCommName(args.hcclComm, hcom_name)拿到group字符串;
  3. 构造 aclTensorCreateAclTensor内部完成aclrtMallocaclrtMemcpy(HOST→DEVICE)与aclCreateTensor(ND 格式、连续 strides);
  4. 两段式调用:先GetWorkspaceSize获取 workspace 与 executor,workspaceSize > 0aclrtMalloc申请,再执行第二段接口并在固定位置aclrtSynchronizeStreamWithTimeout同步等待;
  5. 资源回收:依次销毁 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_quantgroup_size=32、scale shape(2,128)(即ceil(64,32)=2)→ SUCCESS;
    • pertensor_quant:scale shape(1)→ SUCCESS;
    • invalid_pergroup_sizegroup_size=16(小于 32)→ PARAM_INVALID;
    • invalid_pergroup_size_not_multiplegroup_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.cppaclnn 两段式接口实现与入参校验
op_api/aclnn_weight_quant_matmul_all_reduce.h对外头文件(接口声明与文档注释)
op_kernel/weight_quant_matmul_all_reduce_tiling_data.htiling 数据结构(含 tile/tail 两段 MatMul 切分)
op_host/op_tiling分架构 tiling 实现(arch22/arch31/arch35)
docs/aclnnWeightQuantMatmulAllReduceV2.mdV2 版本(新增 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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询