资讯动态

CANN PyPTO `pypto.scaled_mm` 完全指南:MX 量化矩阵乘的 API 用法、参数约束与工程实践

发布时间:2026/9/19 17:05:45 来源:尧图企业网站定制
CANN PyPTOpypto.scaled_mm完全指南MX 量化矩阵乘的 API 用法、参数约束与工程实践【免费下载链接】pyptoPyPTO发音: pai p-t-oParallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto导读pypto.scaled_mm是 CANN PyPTO 提供的 MXMicroscaling量化矩阵乘接口用于在 Ascend 950 系列产品上执行带量化参数的 FP8/FP4 矩阵乘核心计算公式为out (mat_a * scale_a) (mat_b * scale_b)。本文以该接口的官方文档为主体结合 matmul.py 源码与 ST 测试用例完整讲解函数原型、全部参数语义、数据类型支持矩阵、量化/Bias/ReLU 扩展能力、对齐约束及可复现的调用示例帮助你在 PyPTO 内核编程中正确、高效地使用 MX 量化矩阵乘。产品支持情况pypto.scaled_mm对当前产品的支持范围明确如下Ascend 950PR / Ascend 950DT支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品不支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品不支持。该能力与仓库源码中的架构判断逻辑一致在 matmul.py 中定义了A2A3_ARCHS (DAV_1001, DAV_2201)并通过__is_a2a3_arch()判断当前 NPU 架构用于在扩展参数校验时排除 A2/A3 平台上不支持的特性。因此本文介绍的所有能力均以 Ascend 950 系列950PR/950DT为运行前提。功能说明pypto.scaled_mm实现mat_a、mat_b矩阵的 MX 量化矩阵乘运算计算公式为out (mat_a * scale_a) (mat_b * scale_b)mat_a、mat_b、scale_a、scale_b为源操作数。其中mat_a为左矩阵mat_b为右矩阵scale_a为左矩阵的量化参数缩放因子scale_b为右矩阵的量化参数。out为目的操作数存放矩阵乘结果的矩阵。从源码实现看scaled_mm 函数体 会先经过四层校验__validate_inputs、__validate_scaled_inputs、__validate_scaled_shape、__validate_out_dtype随后根据输入维度分流2 维输入映射到pypto_impl.MatmulMX3/4 维输入映射到pypto_impl.BatchMatmulMX。这也解释了接口同时支持矩阵乘与批量矩阵乘的原因。函数原型scaled_mm(mat_a, mat_b, out_dtype, scale_a, scale_b, *, a_trans False, b_trans False, scale_a_trans False, scale_b_trans False, c_matrix_nz False, extend_paramsNone) - Tensormat_a、mat_b、out_dtype、scale_a、scale_b为位置参数其余为关键字参数。Python 接口定义可在 matmul.py 中查阅函数被op_wrapper装饰返回值类型为Tensor。参数说明下表表 1完整列出pypto.scaled_mm的全部参数语义与官方文档保持一致参数名输入/输出说明mat_a输入表示输入左矩阵。不支持输入空 Tensor。数据类型详见表 3。矩阵维度2 维、3 维、4 维。FormatTILEOP_ND、TILEOP_NZDT_FP8E5M2 输入不支持 TILEOP_NZ 格式。内轴外轴当输入矩阵 mat_a 非转置时对应数据排布为 [M, K]外轴为 M内轴为 K当 mat_a 转置时对应数据排布为 [K, M]外轴为 K内轴为 M。对齐要求Format 为 TILEOP_NDND 格式时外轴范围为 [1, 2^31 - 1]内轴范围为 [1, 65535]Format 为 TILEOP_NZNZ 格式时Shape 维度需满足内轴 32 字节对齐、外轴 16 元素对齐。使用 pypto.view 接口的场景传入 View 的 Shape 维度同样需满足内轴 32 字节对齐、外轴 16 元素对齐。mat_b输入表示输入右矩阵。不支持输入空 Tensor。数据类型详见表 3。矩阵维度2 维、3 维、4 维。FormatTILEOP_ND、TILEOP_NZDT_FP8E5M2 输入不支持 TILEOP_NZ 格式。内轴外轴当 mat_b 非转置时对应数据排布为 [K, N]外轴为 K内轴为 N当 mat_b 转置时对应数据排布为 [N, K]外轴为 N内轴为 K。对齐要求与 mat_a 相同ND 格式外轴范围 [1, 2^31 - 1]、内轴范围 [1, 65535]NZ 格式需内轴 32 字节对齐、外轴 16 元素对齐pypto.view 场景同理。out_dtype输出表示输出矩阵数据类型。基础场景输出数据类型支持情况详见表 3量化场景输出数据类型支持情况详见表 4。scale_a输入表示输入左矩阵量化参数。不支持输入空 Tensor。数据类型详见表 3。量化参数维度比 mat_a 多 1 维mat_a 为 2/3/4 维时scale_a 分别为 3/4/5 维。FormatTILEOP_ND。量化参数 shape记 Ks CeilAlign(K, 64) / 64。二维场景非转置为 [M, Ks, 2]转置为 [Ks, M, 2]三维、四维场景在上述 shape 前增加 mat_a 自身的 Batch 前缀。scale_a 的每个 Batch 轴必须与 mat_a 严格相等并与 mat_a 使用相同的广播索引。scale_b输入表示输入右矩阵量化参数。不支持输入空 Tensor。数据类型详见表 3。量化参数维度比 mat_b 多 1 维mat_b 为 2/3/4 维时scale_b 分别为 3/4/5 维。FormatTILEOP_ND。量化参数 shape记 Ks CeilAlign(K, 64) / 64。二维场景非转置为 [Ks, N, 2]转置为 [N, Ks, 2]三维、四维场景在上述 shape 前增加 mat_b 自身的 Batch 前缀。scale_b 的每个 Batch 轴必须与 mat_b 严格相等并与 mat_b 使用相同的广播索引。a_trans输入表示输入左矩阵是否转置默认为 False。b_trans输入表示输入右矩阵是否转置默认为 False。scale_a_trans输入表示输入左矩阵量化参数是否转置默认为 False。scale_b_trans输入表示输入右矩阵量化参数是否转置默认为 False。c_matrix_nz输入表示输出矩阵的 Format 是否采用 NZ 格式默认为 False当前仅支持设置 False即输出矩阵仅支持 ND 格式。extend_params输入支持 bias 及 fixpipe 的量化功能数据类型为字典格式默认为 None详见表 2。其中CeilAlign(value, align)元素对齐实现为(value align - 1) / align * align。量化参数 shape 的源码级验证上述 Ks 的计算与 shape 校验逻辑在 matmul.py 的校验函数 中得到了严格实现__validate_scale_k_alignment校验k_a_scale0_dim (ka_dim 64 - 1) // 64即scale_a的 Ks 维必须等于输入 K 维按 64 对齐向上取整后的块数这与文档中Ks CeilAlign(K, 64) / 64完全对应注意CeilAlign(K,64)/64等价于ceil(K/64)__validate_scale_k1_dimensions要求 scale_a、scale_b 最后一维必须都为 2__validate_scale_m_dimensions/__validate_scale_n_dimensions校验 scale 的 M/N 维度与矩阵的 M/N 维度严格相等__validate_scaled_shape中循环逐 Batch 轴校验scale_a的 Batch 维必须等于mat_a的 Batch 维scale_b的 Batch 维必须等于mat_b的 Batch 维这就是文档所述每个 Batch 轴必须严格相等的落地实现__validate_scaled_inputs要求 scale_a、scale_b 维度一致且都等于input_dim 1同时明确 scale 张量只允许 TILEOP_ND 格式matmul.py与文档FormatTILEOP_ND一致。此外NZ_UNSUPPORTED_INPUT_DTYPES 包含DT_FP32 / DT_FP8E5M2 / DT_HF8__validate_nz_input会在这些数据类型搭配 TILEOP_NZ 格式时抛出错误印证了文档中DT_FP8E5M2 输入不支持 TILEOP_NZ 格式的约束。extend_params 参数说明表 2 完整说明extend_params字典中支持的扩展功能参数名说明scale表示 pertensor 量化场景使用同一个缩放因子将高精度数映射到低精度数输出矩阵量化的参数。输入为 float 类型取 1 位符号位 8 位指数位 10 位尾数位参与运算。输入输出数据类型支持情况详见表 4。不支持叠加多核切 k 功能。scale_tensor表示 perchannel 量化场景对每一个输出通道独立计算一套量化参数输出矩阵量化的矩阵。scale_tensor 输入固定为 uint64_t 或 int64_t 的 Tensor计算时会转换 64bit 为 float 类型的低 32bit 后取 1 位符号位 8 位指数位 10 位尾数位参与运算。输入输出数据类型支持情况详见表 4。scale_tensor 的第一维度必须置 1且 N 维度需要与 mat_b 矩阵的 N 维度相等。scale_tensor 只支持 ND 格式。仅支持矩阵维度为 2 维场景。不支持叠加多核切 k 功能。量化输出类型为 DT_INT8 场景时需要在 function 外提前调用 torch_npu.npu_trans_quant_param 并传入 float32 类型的 torch.tensor 来获取 int64 数据类型的 scale_tensor。bias_tensor表示偏置矩阵。输入为 Tensor 类型。Bias 矩阵数据类型可选 DT_FP16、DT_BF16 和 DT_FP32。bias_tensor 只支持 ND 格式。仅支持矩阵维度为 2/3/4 维场景。当输入矩阵为 3 维时Bias 维度可以为 [B, 1, N] 或 [1, N]且 N 维度需要与 mat_b 矩阵的 N 维度相等。当输入矩阵为 4 维时Bias 维度只能为 [1, N]且 N 维度需要与 mat_b 矩阵的 N 维度相等。不支持叠加多核切 K 功能。relu_type表示输出矩阵是否进行 ReLU 操作。输入为 ReLuType 类型。支持 RELU 和 NO_RELU 两种模式。不支持叠加多核切 k 功能。关于 Bias 维度约束的源码佐证__validate_bias_dimension 规定 2 维/4 维输入只接受 2 维 Bias3 维输入接受 2 维或 3 维 Bias__validate_scale_dimension则要求scale_tensor必须为 2 维。ReLuType枚举定义在 ReLuType.md原型为class ReLuType(enum.Enum): NO_RELU ... # 不使能ReLu功能 RELU ... # 使能ReLu功能数据类型支持情况表 3基础场景支持的数据类型mat_amat_bout_dtypescale_ascale_bbias_tensorDT_FP8E5M2DT_FP8E5M2 / DT_FP8E4M3DT_FP16 / DT_BF16 / DT_FP32DT_FP8E8M0DT_FP8E8M0DT_FP16 / DT_BF16 / DT_FP32DT_FP8E4M3DT_FP8E5M2 / DT_FP8E4M3DT_FP16 / DT_BF16 / DT_FP32DT_FP8E8M0DT_FP8E8M0DT_FP16 / DT_BF16 / DT_FP32DT_FP4_E2M1DT_FP4_E2M1DT_FP16 / DT_BF16 / DT_FP32DT_FP8E8M0DT_FP8E8M0DT_FP16 / DT_BF16 / DT_FP32表 4量化场景支持的数据类型mat_amat_bout_dtypeDT_FP8E5M2DT_FP8E5M2、DT_FP8E4M3DT_INT8DT_FP8E4M3DT_FP8E5M2、DT_FP8E4M3DT_INT8DT_FP4_E2M1DT_FP4_E2M1DT_INT8这两张表的约束在源码中有对应的常量定义MX_INPUT_COMBOS 限定了 mat_a/mat_b 的合法组合FP8E5M2/FP8E4M3 的四种交叉组合 FP4_E2M1 自组合MX_BASIC_OUT_DTYPES 定义了 FP8/FP4 输入的基础输出类型均为DT_FP16/DT_BF16/DT_FP32QUANT_OUT_DTYPES 定义了量化场景extend_params 携带scale或scale_tensor时由 __has_quant_scale 判定下输出类型收敛为DT_INT8。若传入不支持的组合__validate_inputs与__validate_out_dtype会抛出PyptoError错误信息会提示查阅 API 文档的数据类型表。返回值说明返回值为out矩阵Tensor即矩阵乘计算结果。约束说明使用pypto.scaled_mm前需满足以下约束当输入为 DT_FP4_E2M1 量化场景时需保证内轴为偶数。调用pypto.scaled_mm接口前需要通过pypto.set_cube_tile_shapes设置 M、K、N 轴上的切分大小。调用pypto.scaled_mm接口的输入为调用pypto.reshape后的 NZ 格式时需要调用pypto.set_matrix_size接口设置pypto.reshape前的输入到 matmul 的原始 Shape 的 m、k、n 值。set_cube_tile_shapes 与 set_matrix_size 说明pypto.set_cube_tile_shapes的签名定义在 _controller.pydef set_cube_tile_shapes(m: List[int], k: List[int], n: List[int], enable_split_k: bool False):其中m、k、n各为一个长度为 2 的列表分别表示对应维度在 Cube 计算中的两级切分大小如 L1/L0 缓存层级enable_split_k表示是否启用多核切 K结果在 GM 中累加默认 False。由于文档明确scale / scale_tensor / bias_tensor / relu_type均不支持叠加多核切 k 功能因此在使用这些扩展参数时应保持enable_split_kFalse。ST 测试用例给出了典型的切分与循环调用模式可参考 test_scaled_mm_mxfp8.py在嵌套pypto.loop中通过pypto.view切出子矩阵每轮循环内调用pypto.set_cube_tile_shapes(*tile_shape, config.enable_ksplit)后执行pypto.scaled_mm最后用pypto.assemble将子结果写回输出张量。测试用例配置ScaledMMConfig与用例表SCALED_MM_TESTS位于 scaled_mm_mxfp8_test_case.py其中包含m_tile_shape、k_tile_shape、n_tile_shape、a_trans/b_trans、scale_a_trans/scale_b_trans、a_format/b_format/c_format、has_bias、enable_ksplit等字段是构造合法调用参数的重要参考。调用示例以下示例完整继承官方文档覆盖基本矩阵乘、Batch 广播、Bias 叠加与量化叠加 ReLU 四种场景。基本矩阵乘mat_a pypto.tensor([64, 128], pypto.DT_FP8E5M2, mat_a) mat_b pypto.tensor([128, 32], pypto.DT_FP8E5M2, mat_b) scale_a pypto.tensor([64, 2, 2], pypto.DT_FP8E8M0, scale_a) scale_b pypto.tensor([2, 32, 2], pypto.DT_FP8E8M0, scale_b) out pypto.scaled_mm(mat_a, mat_b, pypto.DT_BF16, scale_a, scale_b)这里K128Ks ceil(128/64) 2因此scale_a为[M64, Ks2, 2]scale_b为[Ks2, N32, 2]与表 1 的 shape 规则完全吻合。3 维单侧广播A/scale_a 作为整体沿 Batch 轴广播mat_a pypto.tensor([1, 128, 256], pypto.DT_FP8E5M2, mat_a) mat_b pypto.tensor([4, 256, 64], pypto.DT_FP8E5M2, mat_b) scale_a pypto.tensor([1, 128, 4, 2], pypto.DT_FP8E8M0, scale_a) scale_b pypto.tensor([4, 4, 64, 2], pypto.DT_FP8E8M0, scale_b) out pypto.scaled_mm(mat_a, mat_b, pypto.DT_BF16, scale_a, scale_b)3 维矩阵的 scale 比矩阵多 1 维scale_a的 Batch 轴[1]与mat_a的 Batch 轴[1]严格相等scale_b的 Batch 轴[4]与mat_b的 Batch 轴[4]严格相等且各自使用相同的广播索引因此mat_aBatch1可整体沿mat_b的 Batch4 广播。4 维交叉广播scale Batch 前缀分别与配对矩阵保持一致mat_a pypto.tensor([2, 1, 128, 256], pypto.DT_FP8E5M2, mat_a) mat_b pypto.tensor([1, 3, 256, 64], pypto.DT_FP8E5M2, mat_b) scale_a pypto.tensor([2, 1, 128, 4, 2], pypto.DT_FP8E8M0, scale_a) scale_b pypto.tensor([1, 3, 4, 64, 2], pypto.DT_FP8E8M0, scale_b) out pypto.scaled_mm(mat_a, mat_b, pypto.DT_BF16, scale_a, scale_b)4 维矩阵的 scale 为 5 维scale_a [2, 1, 128, 4, 2]的 Batch 前缀[2, 1]与mat_a的[2, 1]相等scale_b [1, 3, 4, 64, 2]的 Batch 前缀[1, 3]与mat_b的[1, 3]相等。交叉广播时每个 scale 只跟随其配对矩阵的 Batch 形状。叠加 Biasmat_a pypto.tensor([128, 64], pypto.DT_FP8E5M2, mat_a) mat_b pypto.tensor([32, 128], pypto.DT_FP8E5M2, mat_b) scale_a pypto.tensor([2, 64, 2], pypto.DT_FP8E8M0, scale_a) scale_b pypto.tensor([32, 2, 2], pypto.DT_FP8E8M0, scale_b) bias pypto.tensor((1, 32), pypto.DT_FP16, tensor_bias) extend_params {bias_tensor: bias} out pypto.scaled_mm(mat_a, mat_b, pypto.DT_BF16, scale_a, scale_b, scale_a_transTrue, scale_b_transTrue, extend_paramsextend_params)本例中mat_b实际排布为[N32, K128]因此b_trans语义通过scale_b_transTrue配合完成scale_a使用转置排布[Ks, M, 2]scale_b使用转置排布[N, Ks, 2]Bias 维度(1, 32)的 N32 与mat_b的 N 维相等。量化叠加 ReLUscale_cpu pypto.tensor((1, 32), pypto.DT_UINT64, tensor_scale) scale_tensor torch_npu.npu_trans_quant_param(scale_cpu.npu()) # 生成scale_tensor mat_a pypto.tensor([128, 64], pypto.DT_FP8E5M2, mat_a) mat_b pypto.tensor([32, 128], pypto.DT_FP8E5M2, mat_b) scale_a pypto.tensor([2, 64, 2], pypto.DT_FP8E8M0, scale_a) scale_b pypto.tensor([32, 2, 2], pypto.DT_FP8E8M0, scale_b) extend_params {scale_tensor: scale_tensor, relu_type: pypto.ReLuType.RELU} out pypto.scaled_mm(mat_a, mat_b, pypto.DT_BF16, scale_a, scale_b, scale_a_transTrue, scale_b_transTrue, extend_paramsextend_params)注意该示例中scale_tensor由torch_npu.npu_trans_quant_param生成在 DT_INT8 量化输出场景下需要预先将 float32 的 torch.tensor 转换为 int64/uint64 数据类型的scale_tensor文档要求输入固定为 uint64_t 或 int64_t 的 Tensor再以 perchannel 量化参数传入extend_params。仓库 ST 测试 test_matmul_quant.py 同样采用torch_npu.npu_trans_quant_param(golden_scale.to(fnpu:{device_id}))的生成方式可作为工程化参考。总结与工程建议pypto.scaled_mm在 Ascend 950 系列上为 MX 量化矩阵乘提供了统一的 Python 编程入口其核心要点可归纳为shape 规则先行scale 永远比矩阵多 1 维且最后一个维度固定为 2对应 MX 量化的 per-block 缩放语义Ks 由 K 维按 64 对齐后除以 64 得到Batch 维度必须与配对矩阵严格相等Format 与对齐纪律scale 只支持 NDFP8E5M2 输入不支持 NZ输出仅支持 NDc_matrix_nz当前只能为 FalseND 内轴 ≤ 65535NZ 场景需内轴 32 字节对齐、外轴 16 元素对齐扩展能力按需启用extend_params支持 pertensor 量化scale、perchannel 量化scale_tensor、Biasbias_tensor与 ReLUrelu_type但这些扩展均不支持多核切 K使用时应保持set_cube_tile_shapes的enable_split_kFalse运行前置条件调用前必须通过pypto.set_cube_tile_shapes完成 M/K/N 切分设置NZ reshape 场景还需配合pypto.set_matrix_size。如需更深入理解接口行为建议结合 matmul.py 的校验源码、scaled_mm_mxfp8_test_case.py 的用例配置以及 test_scaled_mm_mxfp8.py 的内核实现view 切分 set_cube_tile_shapes assemble 的完整流水对照阅读。【免费下载链接】pyptoPyPTO发音: pai p-t-oParallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价