已关闭
[Bug-Report|缺陷反馈]: nn仓torch新接口 mxfp4数据类型 不支持b矩阵非转置(0,1,2)场景 #4860
link164735创建于 8月17日关闭于 8月25日
8月17日 添加了label:bug-report
8月17日 将 Hu1L1 设为负责人
yuning_chen
8月17日 评论:
8月17日 评论:
您好,感谢反馈,问题已收到,当前 @Hu1L1 正在跟踪处理。


丛吉钰
8月18日 评论:
8月18日 评论:
/asign


丛吉钰
8月18日 评论:
8月18日 评论:
/assign


8月18日 将 cong-jiyu 设为负责人,移除负责人 Hu1L1
8月25日 issue状态由 进行中 改变为 已解决
8月25日 关闭了 issue
8月25日 添加了label:resolved
Thanks for sending an issue! Please fill in the following template to help quickly solve your problem.
Describe the current behavior / 问题描述 (Mandatory / 必填)
torch_extension 框架 fp4 布局约束导致 transpose_quant_batch_mat_mul perm_x2=[0,1,2] 不可用
现象
matmul/transpose_quant_batch_mat_mul的 torch 接口在 fp4(FLOAT4_E2M1)输入 +perm_x2=[0,1,2]参数组合下,aclnn 报:而
perm_x2=[0,2,1]路径在相同 fp4 输入下工作正常。README 与aclnnTransposeQuantBatchMatMul.md均声称 MX 模式permX2支持[0,1,2] 或 [0,2,1],因此表面上看是 aclnn 算子 bug,但实际根因在 torch_extension 框架层。根因定位
根因集中在
torch_extension/cann_ops_nn/common/aclnn_common.h,由 5 处协同实现一个隐式约定:fp4 tensor 的"压缩维"(字节数 = 逻辑元素数 / 2 的那一维)必须在物理最后一维。1. 常量定义(aclnn_common.h:78)
constexpr int64_t FP4_IN_INT8 = 2;1 字节 = 2 个 fp4 元素,即"压缩比"。
2. 4-bit dtype 识别(aclnn_common.h:181-184)
inline bool Is4BitDtype(const aclDataType acl_data_type) { return acl_data_type == ACL_FLOAT4_E2M1 || acl_data_type == ACL_FLOAT4_E1M2 || acl_data_type == ACL_INT4; }3. 核心约束:
CollectB4ShapeInfo(aclnn_common.h:210-239)inline void CollectB4ShapeInfo(const at::Tensor& at_tensor, c10::SmallVector<int64_t, MAX_DIM_NUM>& wrapperStride, c10::SmallVector<int64_t, MAX_DIM_NUM>& wrapperShape) { int64_t nDim = at_tensor.sizes().size(); if (nDim == 1) { wrapperShape[0] = wrapperShape[0] * FP4_IN_INT8; } else if (nDim > 1) { if (wrapperStride[nDim - 1] == 1 && wrapperStride[nDim - PENULTIMATE_DIM] == 1) { if (wrapperShape[nDim - PENULTIMATE_DIM] == 1) { wrapperStride[nDim - 1] *= FP4_IN_INT8; wrapperShape[nDim - PENULTIMATE_DIM] *= FP4_IN_INT8; } else if (wrapperShape[nDim - 1] == 1) { wrapperStride[nDim - PENULTIMATE_DIM] *= FP4_IN_INT8; wrapperShape[nDim - 1] *= FP4_IN_INT8; } } else if (wrapperStride[nDim - 1] == 1) { wrapperStride[nDim - PENULTIMATE_DIM] *= FP4_IN_INT8; wrapperShape[nDim - 1] *= FP4_IN_INT8; // ★ 连续 3D tensor 走这里 } else if (wrapperStride[nDim - PENULTIMATE_DIM] == 1) { wrapperStride[nDim - 1] *= FP4_IN_INT8; wrapperShape[nDim - PENULTIMATE_DIM] *= FP4_IN_INT8; } for (auto i = 0; i < nDim - PENULTIMATE_DIM; i++) { wrapperStride[i] = wrapperStride[i] * FP4_IN_INT8; } } }关键事实:
wrapperStride(物理布局连续性)决定,不参考算子语义或 perm 参数。wrapperStride[nDim-1] == 1),永远只展开wrapperShape[nDim-1](物理最后一维)。index 0..nDim-3)的wrapperShape永远不被 ×2。4. 展开触发点:
PrepareTensorMeta(aclnn_common.h:463-465)if (acl_data_type != ACL_STRING && Is4BitDtype(acl_data_type)) { CollectB4ShapeInfo(at_tensor, meta.wrapperStride, meta.wrapperShape); meta.storageDims.back() *= FP4_IN_INT8; }所有 4-bit dtype 的
at::Tensor → aclTensor转换前都会被调用,算子 csrc 无法绕过。5. csrc 触发入口:
ConvertType(TensorWrapper)(aclnn_common.h:633-653)inline aclTensor* ConvertType(const TensorWrapper& tensor_wrapper) { ... aclDataType acl_data_type = tensor_wrapper.dtype; auto meta = PrepareTensorMeta(at_tensor, acl_data_type); // ← 4-bit 在此被自动展开 auto acl_tensor = aclCreateTensor(meta.wrapperShape.data(), ...); return acl_tensor; }transpose_quant_batch_mat_mul.cpp:136-139用TensorWrapper{x, ACL_FLOAT4_E2M1}包装 x1/x2,框架在此自动展开物理最后一维 → csrc 无法控制展开哪一维。触发链路详解
按 README 物理 shape 约定(
B, M, K, N = 2, 16, 64, 32,fp4 时物理 K 维 = K/2 字节):x1(两种 perm 路径都正常)
perm_x1默认[1,0,2],x1 物理 shape[M, B, K/2] = [16, 2, 32]:CollectB4ShapeInfo展开物理最后一维 32 → 64x1.shape = [16, 2, 64]x1KDim = x1.view[permX1[2]] = x1.view[2] = 64✅ 逻辑 Kx2 + perm_x2=[0,1,2](问题路径)
x2 物理 shape
[B, K/2, N] = [2, 32, 32](K 压缩维在中间维 index=1):CollectB4ShapeInfo只展开物理最后一维 N:32 → 64x2.shape = [2, 32, 64]x2KDim = x2.view[permX2[1]] = x2.view[1] = 32❌ 仍是字节数 K/2x1KDim(64) != x2KDim(32)→ 报K-axis of x1, x2 must be equalx2 + perm_x2=[0,2,1](workaround 路径)
x2 物理 shape
[B, N, K/2] = [2, 32, 32](K 压缩维在最后维 index=2):CollectB4ShapeInfo展开物理最后一维 K/2:32 → 64x2.shape = [2, 32, 64]x2KDim = x2.view[permX2[1]] = x2.view[2] = 64✅ 逻辑 Kx1KDim(64) == x2KDim(64)→ 正常fp8 路径不受影响
fp8(
FLOAT8_E4M3FN)无压缩,1 字节 = 1 元素,Is4BitDtype返回 false,CollectB4ShapeInfo不触发,perm_x2=[0,1,2] 和 [0,2,1] 都正常。影响范围排查
在
torch_extension/csrc层涉及 fp4 的 5 个算子中,只有transpose_quant_batch_mat_mul受影响:perm_x2参数的 fp4 算子,perm_x2=[0,1,2]让 x2 压缩维落在物理中间维quant_matmul_activation_quanttransposeX1/X2硬编码 false(csrc L77-78),压缩维始终在物理最后维swiglu_group_quantflat_quantmx_to_block_mx_quant其他算子(
weight_quant_batch_matmul_v2、quant_batch_matmul_v3、rotate_quant等)虽 op_host 层支持 fp4,但没有torch_extension/csrc目录,不经过CollectB4ShapeInfo,不受影响。孤立问题:
transpose_quant_batch_mat_mul是唯一同时满足"fp4 输入 + 有 perm 参数可改变压缩维位置"的算子,所以不需要框架级改动。修复方案
方案 A(推荐,最小改动)
在 torch 接口层对
fp4 + perm_x2=[0,1,2]主动报清晰错误,避免用户看到模糊的 aclnn K-axis 报错;同时在文档中明确该约束。Python 前端(
transpose_quant_batch_mat_mul.py)在_check_mx_input后增加:def _check_fp4_perm_x2(x1_acl_dtype, x2_acl_dtype, perm_x2): # torch_extension 框架 CollectB4ShapeInfo 只展开物理最后一维, # fp4 压缩维必须在物理最后一维, 因此 perm_x2=[0,1,2] (K 在中间维) 不可用. fp4_acls = (_FP4_E2M1_ACL,) # 296 is_fp4 = (x1_acl_dtype in fp4_acls) or (x2_acl_dtype in fp4_acls) if is_fp4 and perm_x2 is not None and list(perm_x2) == [0, 1, 2]: raise NotImplementedError( "fp4 input requires perm_x2=[0,2,1] due to torch_extension framework " "CollectB4ShapeInfo constraint (fp4 compressed dim must be the physical " "last dim); perm_x2=[0,1,2] is only available for fp8 input" )csrc(
transpose_quant_batch_mat_mul.cpp)在 ACLNN_CMD 前增加对应TORCH_CHECK:const bool is_fp4 = (x1_acl == ACL_FLOAT4_E2M1) || (x2_acl == ACL_FLOAT4_E2M1); const bool is_perm_x2_012 = perm_x2.has_value() && perm_x2.value().size() == 3 && perm_x2.value()[0] == 0 && perm_x2.value()[1] == 1 && perm_x2.value()[2] == 2; TORCH_CHECK(!(is_fp4 && is_perm_x2_012), "fp4 input requires perm_x2=[0,2,1] (torch_extension framework fp4 layout constraint)");文档:把
FIX_dtype_check.md第 350-353 行的"已知遗留问题"重新归类为"框架约束"。方案 B(不推荐,框架级修复)
扩展
CollectB4ShapeInfo支持"任意维是压缩维"。需修改aclnn_common.h:210-239让展开逻辑参考算子语义/perm 参数,但CollectB4ShapeInfo的调用方(PrepareTensorMeta)无法获取 perm 信息,需要新增 API 透传。回归测试覆盖所有 4-bit dtype 算子,风险高、收益低(仅一个算子受益)。验证依据
环境:Ascend 950PR / CANN 9.2.0 / torch 2.7.1 / torch_npu 2.7.1.post5
[B, K/2, N] = [2, 32, 32][0,1,2][2, 32, 64](最后维 N 被 ×2,K/2 未展开)[B, N, K/2] = [2, 32, 32][0,2,1][2, 32, 64](最后维 K/2 被 ×2 = K)[B, K, N] = [2, 64, 32][0,1,2][2, 64, 32](无展开)[B, N, K] = [2, 32, 64][0,2,1][2, 32, 64](无展开)相关文件
torch_extension/cann_ops_nn/common/aclnn_common.h:78, 181-184, 210-239, 463-465, 633-653matmul/transpose_quant_batch_mat_mul/torch_extension/csrc/transpose_quant_batch_mat_mul.cpp:136-139matmul/transpose_quant_batch_mat_mul/torch_extension/transpose_quant_batch_mat_mul.pymatmul/transpose_quant_batch_mat_mul/torch_extension/FIX_dtype_check.md:350-353matmul/transpose_quant_batch_mat_mul/docs/aclnnTransposeQuantBatchMatMul.md:366Environment / 环境信息 (Mandatory / 必填)
Ascend950
Steps to reproduce the issue / 重现步骤 (Mandatory / 必填)
构建tqbmm mxfp4用例 mxfp4 + perm_x2=[0,1,2]场景,用例报错显示 k维度未对齐
Describe the expected behavior / 预期结果 (Mandatory / 必填)
用例pass
Related log / screenshot / 日志 / 截图 (Mandatory / 必填)
不涉及
Special notes for this issue/备注 (Optional / 选填)