拓冰建站拓冰建站
首页 / 资讯中心 / 正文

CANN ops-transformer 通算融合算子 aclnnWeightQuantMatmulAllReduce 完全指南:权重伪量化 Matmul + AllReduce 融合计算

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 模块属于 MC2Matmul 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) $$各符号含义符号含义x1MatMul 左矩阵激活侧不量化BFLOAT16 / FLOAT16x2MatMul 右矩阵权重侧量化存储INT8 / INT4或 950 系列上的 FLOAT8_E4M3FN / HIFLOAT8antiquantOffset伪量化 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.cppAtlas A2910B 架构x1/scale/offset/x3/output支持 FLOAT16 与 BFLOAT16x2仅支持 INT8/INT4Ascend 950DAV_3510 架构x2额外支持 FLOAT8_E4M3FN 与 HIFLOAT8对应源码中的dtypeSupportListQuantA5Atlas 310PDAV_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)非连续 Tensorx1输入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 操作类型当前版本仅支持输入sumString---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_SUMstreamMode必须为 1antiquantGroupSize为 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_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、CheckDtypeValid、CheckAttr、CheckShape返回ACLNN_ERR_PARAM_INVALID按顺序完成检查见 aclnn_weight_quant_matmul_all_reduce.cpp。约束说明确定性计算Atlas A2 训练/推理系列产品910BaclnnWeightQuantMatmulAllReduce默认非确定性实现支持通过配置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 和 x2x3非空场景、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_Kk 为 0 时返回 SUCCESS和empty_Mm 为 0 时返回 PARAM_INVALID一一对应。输入输出数据类型组合Atlas A2 训练/推理系列产品910Bx1x2biasantiquantScaleantiquantOffsetx3output限制BFLOAT16INT8、INT4null、BFLOAT16BFLOAT16null、BFLOAT16null、BFLOAT16BFLOAT16-FLOAT16INT8、INT4null、FLOAT16FLOAT16null、FLOAT16null、FLOAT16FLOAT16-Ascend 950PR / Ascend 950DTx1x2biasantiquantScaleantiquantOffsetx3output限制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 集合通信库接口HcclGetCommName、HcclCommInitAll、HcclCommDestroy。代码对应(m, k) (k, n) bias x3 → AllReduce的 FLOAT16 INT8 场景其中antiquantGroupSize 0表示非 pergroup示例为 perchannelscale 与 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::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 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::vectorint64_t x1Shape {32, 64}; std::vectorint64_t x2Shape {64, 128}; std::vectorint64_t biasShape {128}; std::vectorint64_t antiquantScaleShape {128}; std::vectorint64_t antiquantOffsetShape {128}; std::vectorint64_t x3Shape {32, 128}; std::vectorint64_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::vectorop::fp16_t x1HostData(x1ShapeSize, 1); std::vectorint8_t x2HostData(x2ShapeSize, 1); std::vectorop::fp16_t biasHostData(biasShapeSize, 1); std::vectorop::fp16_t antiquantScaleHostData(antiquantScaleShapeSize, 1); std::vectorop::fp16_t antiquantOffsetHostData(antiquantOffsetShapeSize, 1); std::vectorop::fp16_t x3HostData(x3ShapeSize, 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_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::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(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字符串构造 aclTensorCreateAclTensor内部完成aclrtMalloc、aclrtMemcpyHOST→DEVICE与aclCreateTensorND 格式、连续 strides两段式调用先GetWorkspaceSize获取 workspace 与 executorworkspaceSize 0时aclrtMalloc申请再执行第二段接口并在固定位置aclrtSynchronizeStreamWithTimeout同步等待资源回收依次销毁 tensor、device 内存、workspace、stream、HcclComm、context并aclrtResetDevice最后aclFinalize。仓库中的验证与配套资源单元测试仓库在 tests/ut/op_api 下提供了针对本接口的完整 UTtest_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_1x1(32,64) FLOAT16 / x2(64,128) INT8 / scale(128) / output(32,128) → SUCCESSpergroup_quantgroup_size32、scale shape(2,128)即ceil(64,32)2→ SUCCESSpertensor_quantscale shape(1)→ SUCCESSinvalid_pergroup_sizegroup_size16小于 32→ PARAM_INVALIDinvalid_pergroup_size_not_multiplegroup_size24非 32 倍数→ PARAM_INVALIDempty_Kx1(32,0)、x2(0,128) → SUCCESSk 为 0 的空 tensor 场景empty_Mx1(0,64) → PARAM_INVALIDbs/m 为 0 不支持。Golden 验证脚本tests/assets/impl/golden.py 提供了多卡 golden 参考实现其中针对 WeightQuant 变体按(x2 offset) * scale完成权重反量化并将 pergroup 的 scale/offset 通过repeat_interleave(group_size, dim0)展开为逐元素系数后参与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/arch35docs/aclnnWeightQuantMatmulAllReduceV2.mdV2 版本新增 commMode 通信引擎参数README.md模块总览与 MC2 算子族计算公式全集常见问题与排查建议BUS ERROR 等硬件异常优先确认驱动固件包与 CANN 包均为 8.0.RC2 或更高配套版本返回 161002PARAM_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),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门