Pull Request已成功合入, 合并人@ascend-robot
(感谢 马琦钧 的贡献)变更摘要
此 PR 主要修复 4-bit 数据类型格式处理中 A8/A16 场景的区分问题,为 Ascend950 平台新增 A16W4 量化路径支持(涵盖 INT4、FP4 和 MX 三种变体),同时修正了已有 A8W4 MX 配置中的 NzC0 参数并对 op_api_common.cpp 中 4-bit 数据类型的格式转换逻辑做了简化。
主要改动
- 新增三个 A16W4 判断函数:在
WeightQuantPreprocessKernelNpuOpApi.cpp中新增judge_mm_a16w4_int4、judge_mm_a16w4_fp4和judge_mm_mx_a16w4,分别对 A16W4 场景下的 INT4 权重、FP4(ACL_FLOAT4_E2M1)权重以及 MX 微缩放格式的激活/权重组合进行条件匹配,并设置ctx.is_weight_trans标记。 - 新增 per-channel 权重 scale 准备函数:引入模板函数
prepare_out_weight_scale_per_channel<IsGmm>,用于按通道维度准备输出权重 scale 张量,替代 MX 场景专用的 scale 准备逻辑。 - 修正 A8W4 MX 配置的 NzC0 参数:将
SOC_DATA_FLOW_CONFIG_MAP中judge_mm_mx_a8w4和judge_gmm_mx_a8w4的prepare_out_weight_nz模板参数从NZ_C0_16改为NZ_C0_32,输出格式从ACL_FORMAT_FRACTAL_NZ_C0_16改为ACL_FORMAT_FRACTAL_NZ_C0_32,同时新增常量NZ_C0_8和NZ_C0_32。 - 新增 A16W4 三条数据流配置:在
SOC_DATA_FLOW_CONFIG_MAP中为judge_mm_a16w4_int4、judge_mm_a16w4_fp4和judge_mm_mx_a16w4分别注册对应的预处理流程,均使用NZ_C0_8和ACL_FORMAT_FRACTAL_NZ_C0_16格式。 - 简化 4-bit 格式转换逻辑:在
op_api_common.cpp的ConvertType和ConvertTypeV2中,移除了 4-bit dtype 分支内对FORMAT_FAKE_TO_REAL映射表的查找与格式替换逻辑,并将ConvertTypeV2中CollectB4ShapeInfo的调用参数从static_cast<int64_t>(dimNum)改为直接传入at_tensor。


代码审查
审查总结
本次 diff 涉及 2 个文件的变更,现逐文件审查结果汇总如下:
审查结果
| 优先级 | 数量 |
|---|---|
| P0 | 1 |
| P1 | 0 |
| P2 | 0 |
| P3 | 1 |
各文件审查结论
-
op_plugin/ops/opapi/WeightQuantPreprocessKernelNpuOpApi.cpp:新增了 3 个 judge 函数(judge_mm_a16w4_int4、judge_mm_a16w4_fp4、judge_mm_mx_a16w4)和 1 个 prepare 函数(prepare_out_weight_scale_per_channel),更新了SOC_DATA_FLOW_CONFIG_MAP配置以区分 A8/A16 的 4-bit dtype 处理。发现 1 个问题:prepare_out_weight_scale_per_channel的模板参数IsGmm未被使用(P3)。 -
op_plugin/utils/op_api_common.cpp:在ConvertType(TensorWrapper)和ConvertTypeV2两个重载中移除了FORMAT_FAKE_TO_REAL格式映射逻辑。发现 1 个问题:ConvertTypeV2中CollectB4ShapeInfo(at_tensor, ...)调用无法通过编译,因为at_tensor类型为TensorStructPtr(std::shared_ptr<TensorStruct>),与CollectB4ShapeInfo的任意重载均不匹配(P0)。
整体风险判断
高风险:存在 1 个 P0 编译错误,会导致 ConvertTypeV2 函数在构建时失败,阻塞合入。建议优先修复该编译错误后再合入。
| 类型 | 数量 |
|---|---|
| 🔴 阻塞 | 1 |
| 🟡 建议 | 0 |
⛔ 需要修改


Pull Request 已合并或已关闭。
If you want to solve this problem, you can click here to do it in the FAQs.


Pull Request 已合并或已关闭。
If you want to solve this problem, you can click here to do it in the FAQs.


Pull Request 已合并或已关闭。
If you want to solve this problem, you can click here to do it in the FAQs.


WeightQuantPreprocess A16W4(INT4/FP4/MXFP4)数据流支持
概述
为
npu_weight_quant_preprocess新增 A16W4 紧凑排布(4-bit,uint8 载体)数据流:A16S4 INT4(per-tensor/per-channel/per-group)与 A16F4 FP4(per-group/MX)。torch 接口保持原 9 参形式,无新增参数。ND/NZ 内部判定
{N}/{1,N})/ per-group({G,N}G>1)judge_mm_a16s4_per_tensor/per_channel/per_group,prepare_out_weight_a16s4按is_weight_trans内部分支npu_weight_quant_batchmatmul(eager/图模式均已支持 ND uint8 载体)主要修改
op_api_common.cpp):NZ_C0_8→NZ_C0_16(A16 4-bit),保留 NZ_C0_16→NZ_C0_32(A8)WeightQuantPreprocessKernelNpuOpApi.cpp:A16S4 三个 judge + prepare 内部分支;A16F4judge_mm_a16f4_nz_pergroup/judge_mm_a16f4_mx;prepare_out_weight_nz_a16s4(N-first NZ_C0_8 构造);prepare_out_weight_nd/prepare_out_weight_scale_per_channel保持输入 strides(转置透传)npu_weight_quant_batchmatmul:uint8 载体 4-bit ND weight 按打包方向还原逻辑 K/N;is_weight_nz_4bit_compact识别 NZ_C0_8/NZ_C0_16weight_dtype时按 scale 末维推导逻辑 N验证
test_torch_nz_perchannel、test_torch_nd_pg(per-tensor×4 + ND pc/pg 转置×4 + NZ pg×2)、test_eager_wqbmmv2_nd_direct(ND 直入 4 流)、test_torch_a16s4_variants(29 条:正向 + 转置自动 ND + 负向)test_torch_graph_a16s4_all_flows8 流×16 用例;A16F4 全套关联Issue #444