十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

CANN ops-nn 算子实战:transpose_quant_batch_mat_mul 的 MX 量化批量矩阵乘使用指南

CANN ops-nn 算子实战:transpose_quant_batch_mat_mul 的 MX 量化批量矩阵乘使用指南 人工智能算子库深度学习CANNAscend【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库实现网络在NPU上加速计算。项目地址https://gitcode.com/cann/ops-nn点击查看免费下载本指南围绕 CANN ops-nn 仓库中的transpose_quant_batch_mat_mulTransposeQuantBatchMatMulTorch 接口展开详细说明其在 Ascend 950 系列产品上完成 MXFP8/MXFP4 量化矩阵乘的功能原理、全部参数语义、约束条件与可运行的调用示例。阅读完成后你将掌握如何为三维低比特张量含 FP4 拼包格式构造缩放因子、配置量化分组并正确调用该接口获得 bfloat16/float16 输出同时理解其底层对aclnnTransposeQuantBatchMatMul与aclnnTransposeQuantBatchMatMulWeightNz的分发逻辑。一、接口功能与计算公式cann_ops_nn.transpose_quant_batch_mat_mul完成张量 x1 与张量 x2 的 MX 量化 矩阵乘计算底层封装aclnnTransposeQuantBatchMatMul当 x2 为 FRACTAL_NZ 格式时自动切换封装aclnnTransposeQuantBatchMatMulWeightNz。该 Torch 接口仅支持 MX 量化模式MXFP8 与 MXFP4K-C、T-C 量化模式需使用 aclnn 接口详见 README 与 aclnnTransposeQuantBatchMatMul 文档。计算公式以 batch 维的每个切片为例$$ y[m, n] \sum_{j0}^{K/32-1} \left(\left(\sum_{k0}^{31} x1[m, j \times 32 k] \cdot x2[j \times 32 k, n]\right) \cdot x1Scale[m, j] \cdot x2Scale[j, n]\right) $$其中 K 为矩阵乘的 K 轴长度x1Scale、x2Scale 为torch.float8_e8m0fnu编码的 MX 量化缩放因子矩阵乘中间结果按 K 轴每 32 个元素一组进行缩放累加。示例假设 x1 的 shape 是(M, B, K)x2 的 shape 是(B, K, N)输出 y 的 shape 是(M, B, N)。从仓库实现看该公式的 K 轴按 32 分组逻辑在 aclnnTransposeQuantBatchMatMul 参数校验 中体现为numGroup CeilDivision(CeilDivision(k, 32), 2)的 scale 分组数推导而 scale 最后一维恒为 2对应每个 64 元素块上的两组 32 元素缩放。二、产品支持情况产品支持情况Ascend 950PR/Ascend 950DT支持Atlas A3 训练系列产品/Atlas A3 推理系列产品不支持Atlas A2 训练系列产品/Atlas A2 推理系列产品不支持Atlas 200I/500 A2 推理产品不支持Atlas 推理系列产品不支持Atlas 训练系列产品不支持该支持范围与仓库中算子 AICore 配置一致transpose_quant_batch_mat_mul_def.cpp 仅为ascend950与ascend350注册了AICore().AddConfig(...)而 aclnn 接口 的CheckParams中亦通过IsNpuArch3510Series()做了平台限制校验非目标平台会直接返回ACLNN_ERR_PARAM_INVALID。三、函数原型cann_ops_nn.transpose_quant_batch_mat_mul( x1, x2, *, dtype, biasNone, x1_scaleNone, x2_scaleNone, group_sizesNone, perm_x1None, perm_x2None, perm_yNone, batch_split_factorNone, x1_dtypeNone, x2_dtypeNone, x1_scale_dtypeNone, x2_scale_dtypeNone, ) - Tensor该原型与 torch_extension/transpose_quant_batch_mat_mul.py 中注册的 schema 一一对应参数通过torch.library.impl注册为PrivateUse1后端的自定义算子最终在 csrc 绑定 中转换为 aclnn 调用。四、参数说明参数名参数类型可选/必选描述数据类型维度(shape)x1Tensor必选矩阵乘运算中的左矩阵shape 为 (M, B, K)。MXFP4 场景 Tensor 最后一维为 FP4 拼包后的物理长度 K/2。torch.float8_e4m3fnMXFP4 场景为实际存储类型如 torch.uint8并通过 x1_dtype 指定为 torch_npu.float4_e2m1fn_x23 维(M, B, K)x2Tensor必选矩阵乘运算中的右矩阵数据类型与 x1 一致K 轴长度与 x1 一致。perm_x2 为 [0, 1, 2] 时 shape 为 (B, K, N)perm_x2 为 [0, 2, 1] 时 shape 为 (B, N, K)。MXFP4 场景 Tensor 最后一维为 FP4 拼包后的物理长度N/2 或 K/2。同 x13 维dtypeint必选输出 y 的数据类型枚举值1 表示 torch.float1627 表示 torch.bfloat16。int64-biasTensor可选矩阵乘运算后累加的偏置。预留参数当前暂不支持必须传入 None。--x1_scaleTensor必选x1 的 MX 量化缩放因子。torch.float8_e8m0fnu4 维(M, B, K/64, 2)x2_scaleTensor必选x2 的 MX 量化缩放因子。torch.float8_e8m0fnu4 维perm_x2 为 [0, 1, 2] 时为 (B, K/64, N, 2)perm_x2 为 [0, 2, 1] 时为 (B, N, K/64, 2)group_sizesList[int]必选量化分组大小 [groupSizeM, groupSizeN, groupSizeK]每个元素取值范围为 [0, 65535]。MX 量化场景 groupSizeM 和 groupSizeN 仅支持 0 或 1取值为 0 时由接口根据 scale 的 shape 推断groupSizeK 仅支持 32。int64-perm_x1List[int]可选x1 的转置序列仅支持 [1, 0, 2]默认值 [1, 0, 2]。int64-perm_x2List[int]可选x2 的转置序列支持 [0, 1, 2] 和 [0, 2, 1]默认值 [0, 1, 2]。int64-perm_yList[int]可选输出矩阵的转置序列仅支持 [1, 0, 2]默认值 [1, 0, 2]。int64-batch_split_factorint可选输出矩阵 B 维的切分大小当前仅支持取值 1默认值 1。int64-x1_dtypeint可选x1 的数据类型枚举值不传入时根据 x1 的数据类型自动推导。MXFP4 场景必须传入该参数指定为 torch_npu.float4_e2m1fn_x2。int64-x2_dtypeint可选x2 的数据类型枚举值不传入时根据 x2 的数据类型自动推导。MXFP4 场景必须传入该参数指定为 torch_npu.float4_e2m1fn_x2。int64-x1_scale_dtypeint可选x1_scale 的数据类型枚举值不传入时根据 x1_scale 的数据类型自动推导。int64-x2_scale_dtypeint可选x2_scale 的数据类型枚举值不传入时根据 x2_scale 的数据类型自动推导。int64-补充说明依据 csrc 实现dtype 的取值在 C 侧直接映射为 acl 数据类型ACL_FLOAT161与ACL_BF1627其余取值会在TORCH_CHECK中报错only support float16(1) and bfloat16(27)。x1_dtype/x2_dtype 显式传入时会优先通过GetAclDataType解析为 acl 数据类型而非依赖 Tensor 自身的scalar_type()——这正是 MXFP4 场景用torch.uint8存储、却必须显式声明float4_e2m1fn_x2的原因。PyTorch 层还会对非 MX 输入FP8-INT8、Hifp8在 入口处直接拒绝避免进入底层后产生难以理解的报错。五、返回值说明输出名输出类型描述数据类型维度(shape)yTensorMX 量化矩阵乘的输出。dtype 为 1 时为 torch.float16为 27 时为 torch.bfloat16(M, B, N)输出 shape 的推导可在 meta 函数 中看到完整逻辑M取x1.size(perm_x1[1])、B取x1.size(perm_x1[0])、N取x2.size(perm_x2[2])当输入为 FP4 且对应维度位于最后一维时该维度数值会乘以 2将物理长度还原为逻辑长度。若batch_split_factor 1输出 shape 会变为(batch_split_factor, M, B*N/batch_split_factor)但当前接口约束其只能取 1。六、约束说明以下约束同时被 infershape 与 aclnn 参数校验 双重检查该接口当前支持单算子模式调用。仅支持 MX 量化模式x1、x2、x1_scale 和 x2_scale 必须是 NPU Tensor且 x1_scale 和 x2_scale 为必选输入infershape 中x1Scale or x2Scale is null直接报错。x1 与 x2 的数据类型必须一致MXFP8 场景为 torch.float8_e4m3fnMXFP4 场景用 torch.uint8 表示。x1_scale 与 x2_scale 的数据类型必须为 torch.float8_e8m0fnu即 acl 侧DT_FLOAT8_E8M0IsMicroScaling判定两者均为 E8M0 才进入 MX 分支。仅支持 3 维 Tensorx1 的 shape 为 (M, B, K)x2 的 shape 为 (B, K, N) 或 (B, N, K)x1 与 x2 的 batch 轴和 K 轴必须一致不支持 batch 轴广播。MX 量化场景 K 仅支持 64 的倍数aclnn 侧x1KDim % 64 ! 0即报错。x1_scale 的 shape 必须为 (M, B, K/64, 2)x2_scale 的 shape 在 perm_x2 为 [0, 1, 2] 时必须为 (B, K/64, N, 2)在 perm_x2 为 [0, 2, 1] 时必须为 (B, N, K/64, 2)最后一维必须为 2。group_sizes 必须显式传入且 groupSizeM、groupSizeN 取值为 0 或 1groupSizeK 取值为 32。perm_x1 仅支持 [1, 0, 2]perm_x2 支持 [0, 1, 2] 和 [0, 2, 1]perm_y 仅支持 [1, 0, 2]infershape 的CheckPerm逐一校验。batch_split_factor 当前仅支持取值 1。bias 为预留参数当前暂不支持aclnn 侧bias ! nullptr即报错Torch 接口必须传 None。不支持空 Tensor。MXFP4 场景数据以 torch_npu.float4_e2m1fn_x2 格式存储两个 FP4 元素拼包Tensor 最后一维为物理长度即逻辑长度的一半x1 的最后一维为 K/2x2 在 perm_x2 为 [0, 1, 2] 时最后一维为 N/2在 perm_x2 为 [0, 2, 1] 时最后一维为 K/2。此时 Tensor 实际存储类型如 torch.uint8无法自动推导出 FP4 类型必须通过 x1_dtype、x2_dtype 指定为 torch_npu.float4_e2m1fn_x2。仅 x2 支持 FRACTAL_NZ 格式仅 MX 量化模式x1、输出 y、x1_scale、x2_scale 不允许为 FRACTAL_NZ且非 MX 模式下 x2 也不允许 NZ。最后一条在 csrc 绑定 中的实现逻辑为通过get_npu_format(x2)判断 x2 是否为ACL_FORMAT_FRACTAL_NZ若 x2 为 NZ 而 x1 为 ND则调用aclnnTransposeQuantBatchMatMulWeightNz系列接口否则走aclnnTransposeQuantBatchMatMul。若当前 CANN 版本不支持 WeightNz 路径会提示升级到 9.1 及以上版本或改用 ND 模式。七、group_sizes 的位打包原理Torch 接口的group_sizes是[groupSizeM, groupSizeN, groupSizeK]三段列表但 aclnn 底层groupSize是单个 int64。两者通过 csrc 中的check_and_get_group_size完成打包groupSize groupSizeK | groupSizeN 16 | groupSizeM 32即 K 占低 16 位、N 占 1631 位、M 占 3247 位每个分量取值范围 [0, 65535]。反解时aclnn 的InferGroupSize通过0xFFFF掩码与 16/32 位右移还原三段值并校验 MX 场景下groupSizeM/groupSizeN ∈ {0, 1}且groupSizeK 32随后把 groupSize 规约回纯 K 分组值下发。其中 groupSizeM/groupSizeN 取 0 时表示由接口根据 scale 的 shape 推断推断公式为groupSizeM M / scaleMM 与 x1 shape 的 M 一致scaleM 与 x1Scale shape 的 M 一致同理 groupSizeN 由 N 与 x2Scale 的 N 维度推断推导逻辑可参考 README。八、确定性计算默认支持确定性计算。九、调用示例单算子 eager 模式MXFP8 场景import torch import torch_npu import cann_ops_nn M, B, K, N 64, 16, 128, 256 x1 torch.randn(M, B, K).to(torch.float8_e4m3fn).npu() x2 torch.randn(B, K, N).to(torch.float8_e4m3fn).npu() x1_scale torch.ones(M, B, K // 64, 2, dtypetorch.float8_e8m0fnu).npu() x2_scale torch.ones(B, K // 64, N, 2, dtypetorch.float8_e8m0fnu).npu() y cann_ops_nn.transpose_quant_batch_mat_mul( x1, x2, dtype27, x1_scalex1_scale, x2_scalex2_scale, group_sizes[1, 1, 32], ) print(y.shape, y.dtype)MXFP4 场景import torch import torch_npu import cann_ops_nn M, B, K, N 64, 16, 128, 256 # FP4以torch_npu.float4_e2m1fn_x2格式存储两个FP4元素拼包Tensor最后一维为物理长度逻辑长度的一半 x1 torch.randint(0, 256, (M, B, K // 2), dtypetorch.uint8).npu() x2 torch.randint(0, 256, (B, K, N // 2), dtypetorch.uint8).npu() x1_scale torch.ones(M, B, K // 64, 2, dtypetorch.float8_e8m0fnu).npu() x2_scale torch.ones(B, K // 64, N, 2, dtypetorch.float8_e8m0fnu).npu() y cann_ops_nn.transpose_quant_batch_mat_mul( x1, x2, dtype27, x1_scalex1_scale, x2_scalex2_scale, group_sizes[1, 1, 32], x1_dtypetorch_npu.float4_e2m1fn_x2, x2_dtypetorch_npu.float4_e2m1fn_x2, ) print(y.shape, y.dtype)两个示例中dtype27均表示输出为 bfloat16如需 float16 输出将 dtype 改为 1 即可。示例代码中的 shape 均满足 K 为 64 的倍数、scale 形状 (M, B, K/64, 2) 与 (B, K/64, N, 2) 的约束。十、底层实现与验证线索算子定义transpose_quant_batch_mat_mul_def.cpp 声明了 x1/x2含FLOAT8_E4M3FN、FLOAT4_E2M1等 29 种数据类型组合、bias可选、x1_scale/x2_scaleMX 模式为FLOAT8_E8M0与输出 y并针对ascend950、ascend350配置 AICore同时为 x2 支持 FRACTAL_NZ 格式。shape 推导transpose_quant_batch_mat_mul_infershape.cpp 完成 perm 合法性校验、x1/x2 的 K 轴与 batch 轴一致性检查、MX 模式 dtype 组合检查并按perm转置推导输出 (M, B, N)。aclnn 图构建aclnn_transpose_quant_batch_mat_mul.cpp 在GetWorkspaceSize阶段完成入参检查后通过l0op::Contiguous/ReFormat预处理输入调用l0op::TransposeQuantBatchMatMul构建计算图再经l0op::Cast与ViewCopy输出到目标 Tensor。UT 验证仓库提供了完整的单测覆盖例如 MXFP4 kernel UTM64、B1、K64、N64验证 FP4 拼包输入与 E8M0 scale 的核函数路径、MXFP8/高精度 UT、infershape UT 以及 tiling UT可作为理解各路径行为与回归验证的参考。十一、进一步阅读torchapi_transpose_quant_batch_mat_mul 文档本文来源aclnnTransposeQuantBatchMatMul 接口说明aclnnTransposeQuantBatchMatMulWeightNz 接口说明TransposeQuantBatchMatMul 算子 README含 K-C/MX/T-C 三种量化模式对比MX 量化模式总览已按仓库根目录转换的相对路径aclnn 调用示例C 与 WeightNz 调用示例赞分享人工智能算子库深度学习CANNAscend【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库实现网络在NPU上加速计算。项目地址https://gitcode.com/cann/ops-nn点击查看免费下载相关推荐CANN ops-nn 算子指南aclnnTransposeQuantBatchMatMulWeightNz 接口详解与 MX 量化矩阵乘实战CANN ops nn 算子指南aclnnTransposeQuantBatchMatMulWeightNz 接口详解与 MX 量化矩阵乘实战 本篇技术指南以人工智能算子库深度学习CANNAscendCANN ops-nn 二级量化 MXFP4 矩阵乘算子 aclnnDualLevelQuantMatmulWeightNz 使用指南CANN ops nn 二级量化 MXFP4 矩阵乘算子 aclnnDualLevelQuantMatmulWeightNz 使用指南 aclnnDualLev人工智能算子库深度学习CANNAscend玩转 Ventoy 可启动U盘5个隐藏技巧让多系统安装与随身系统盘一劳永逸玩转 Ventoy 可启动U盘5个隐藏技巧让多系统安装与随身系统盘一劳永逸 做系统维护的老手都懂那种反复折腾的痛一个 U 盘只能烧录一个 ISO装完 W人工智能算子库深度学习CANNAscend上一篇【亲测免费】 Wux Weapp微信小程序开发的利器下一篇如何在5分钟内开始使用localllm零基础入门教程创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表