已合并
feat: WeightQuantPreprocess A16MXF4 支持转置 weight ND 直拷;ErrMsg reason 句号整改 #5241
马琦钧创建于 8月31日
feat: WeightQuantPreprocess A16MXF4 支持转置 weight ND 直拷;ErrMsg reason 句号整改 #5241
已合并
Pull Request已成功合入, 合并人@CANN-robot
(感谢 马琦钧 的贡献)atomgit-bot
8月31日 评论:
8月31日 评论:
变更摘要
本 PR 为 WeightQuantPreprocess 的 A16F4(FP4,uint8 紧凑载体)数据流新增转置 weight ND 直拷支持:转置视图 {K, N}(末两维严格转置、沿 K 打包维须 K 偶)不再报错,与 A16S4 转置路径一致走 ND 直拷输出(outWeight 为 ND 格式、物理透传),pergroup 与 MX(MM_A16MXF4,E8M0 scale)两条数据流均支持;NZ_C0_16 出仍仅支持非转置,转置 + NZ 出返回 ACLNN_ERR_PARAM_INVALID。核心做法是将原有单一判定函数拆分为公共基判定 + 转置/非转置两个分支判定,并在 registry 中新增转置 ND 直拷条目,同时把 4-bit 直拷能力从 INT4 扩展到 FLOAT4_E2M1。
主要改动
- pergroup 判定拆分与转置 ND 条目:
IsMMA16F4PerGroupDataFlow拆分为IsMMA16F4PerGroupBase+IsMMA16F4PerGroupTransDataFlow(转置→ND 直拷)/IsMMA16F4PerGroupNonTransNzDataFlow(非转置→NZ),registry 新增 pergroup 转置条目,checks 组合INPUT_BASE_CHECKS+SCALE_DTYPE_F16_BF16_CHECK+PER_GROUP_SCALE_CHECKS+F4_BIAS_CHECKS+OUT_WEIGHT_ND_CHECKS+OUT_TAIL_CHECKS,processes 为 3 个直拷(ProcessWeightDirectCopy、ProcessWeightScaleDirectCopy、ProcessBiasDirectCopy)。 - MX 判定拆分与转置 ND 条目:
IsMMA16MXF4DataFlow拆分为IsMMA16MXF4Base+IsMMA16MXF4TransDataFlow/IsMMA16MXF4NonTransNzDataFlow,新增MX_SCALE_CHECKS(E8M0 dtype +CheckKGroupSizeMx固定 32 + 2D viewShape),registry 新增 MX 转置 ND 条目;转置 + NZ 出由转置条目的OUT_WEIGHT_ND_CHECKS中的CheckOutWeightFormatND拒绝。 - 新增
F4_BIAS_CHECKS:要求 offset 必须为 nullptr(CheckWeightOffsetOptionalNull)+ bias 直拷透传校验(CheckBiasOptionalNotEmpty/CheckBiasOptionalFormatND/CheckBiasOptionalViewShape/CheckBiasOptionalContiguous),两个转置条目共用。 - 4-bit 直拷扩展至
FLOAT4_E2M1:CheckWeightInt4DirectCopyView与ProcessWeightDirectCopy的 dtype 判断由仅DT_INT4扩展为同时支持DT_FLOAT4_E2M1,相关报错文案由 "INT4" 更新为 "4-bit"。 - UT 用例补充:新增 pergroup 转置 ND 正例 F-17(strides
[1, K]、scale{4, 256}FP16)、转置 + K 奇数负例 T-41(返回ACLNN_ERR_PARAM_INVALID),并将 T-26 注释更新为转置 + NZ 出由转置条目的CheckOutWeightFormatND拒绝。


不准确?
atomgit-bot
8月31日 评论:
8月31日 评论:
8月31日 添加了label:cann-cla/yes
CANN-robot
8月31日 评论:
8月31日 评论:
此处折叠了194条消息 查看更多
19 天前 添加了label:ci-pipeline-passed
19 天前 添加了label:lgtm
19 天前 关闭了关联的issue
19 天前 合入了pull request
描述
A16 MXFP4 数据流(FP4 uint8 紧凑载体,E8M0 scale
{K/32, N},kGroupSize 固定 32)新增转置 weight 支持,并顺带完成 ErrMsg reason 句号整改(合并原 https://gitcode.com/cann/ops-math/pull/5126 ,5126 已关闭)。IsMMA16MXF4DataFlow拆分为 base + 转置/非转置 judge。转置 weight(末两维严格转置 stride[1,K],打包维沿 K 须 K 为偶数)走 ND 直拷透传(ProcessWeightDirectCopy,按打包维建 UINT8 视图物理透传),outWeight 为 ND 格式、viewShape 与 weight 相同,可直接衔接 wqbmmv2 MX kernel(主干已支持 MX FP4 ND 转置输入);非转置维持原 NZ_C0_16 分形转换。pergroup(A16F4 per-group)不支持转置,转置仍返回 ACLNN_ERR_PARAM_INVALID(pergroup FP4 ND 无下游 wqbmmv2 支持)。CheckOutSameBase(空指针对称 + out 非 empty + viewShape 一致)与CheckOutSameAsInput(Base + 连续性 + format/dtype/storageShape 全 ==)两层封装;scale/offset/bias 为直拷透传参数,删除其输入侧校验(empty/format/dtype/viewShape/contiguous 及 kGroupSize 一致性,共 20 个检查函数),约束下放 wqbmmv2。CheckOutWeightDtypeSame——真实场景 NZ 出 dtype 恒等于输入,不一致会在 process ViewCopy 处失败,提前在 check 阶段拒绝。保留CheckWeightOffsetOptionalNull(A8W4-MX/F4 无 offset 直拷 process,非空 offset 会被静默丢弃)。关联的Issue
https://gitcode.com/cann/ops-math/issues/3197
测试
Ascend950 实测:
文档更新
无
类型标签