已关闭
[Bug-Report|缺陷反馈]: nn仓torch新接口 mxfp4数据类型 不支持b矩阵非转置(0,1,2)场景 #4860
link164735创建于  8月17日关闭于  8月25日
link164735
link164735成员
8月17日 创建

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 报:

Invalid_Argument_Tensor_Shape(EZ0010): Parameters x1, x2 of
aclnnTransposeQuantBatchMatMulGetWorkspaceSize have incorrect shapes
[16, 2, 64], [2, 32, 64]. Reason: K-axis of x1, x2 must be equal.

而 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 参数。
  • 对连续 3D tensor(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 → 64
  • aclnn 看到 x1.shape = [16, 2, 64]
  • aclnn 取 x1KDim = x1.view[permX1[2]] = x1.view[2] = 64 ✅ 逻辑 K

x2 + perm_x2=[0,1,2](问题路径)

x2 物理 shape [B, K/2, N] = [2, 32, 32](K 压缩维在中间维 index=1):

  • 框架 CollectB4ShapeInfo 只展开物理最后一维 N:32 → 64
  • 中间维 K/2=32 未被展开
  • aclnn 看到 x2.shape = [2, 32, 64]
  • aclnn 取 x2KDim = x2.view[permX2[1]] = x2.view[1] = 32 ❌ 仍是字节数 K/2
  • x1KDim(64) != x2KDim(32) → 报 K-axis of x1, x2 must be equal

x2 + perm_x2=[0,2,1](workaround 路径)

x2 物理 shape [B, N, K/2] = [2, 32, 32](K 压缩维在最后维 index=2):

  • 框架 CollectB4ShapeInfo 展开物理最后一维 K/2:32 → 64
  • aclnn 看到 x2.shape = [2, 32, 64]
  • aclnn 取 x2KDim = x2.view[permX2[1]] = x2.view[2] = 64 ✅ 逻辑 K
  • x1KDim(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 受影响:

算子 分类 关键依据
transpose_quant_batch_mat_mul 有风险 唯一有 perm_x2 参数的 fp4 算子,perm_x2=[0,1,2] 让 x2 压缩维落在物理中间维
quant_matmul_activation_quant 安全 fp4 仅在 output,transposeX1/X2 硬编码 false(csrc L77-78),压缩维始终在物理最后维
swiglu_group_quant 安全 fp4 仅在 output,csrc L94-95 显式 ceil/2 物理最后维,无 perm 参数
flat_quant 安全 fp4 仅在 output(2 维 tensor),压缩维即最后维,无 perm 参数
mx_to_block_mx_quant 安全 fp4 仅在 input x,csrc L162-164 显式 ×2 物理最后维,无 perm 参数

其他算子(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

用例 x2 物理 shape perm_x2 aclnn 看到 x2 shape 结果
fp4 + perm_x2=[0,1,2] [B, K/2, N] = [2, 32, 32] [0,1,2] [2, 32, 64](最后维 N 被 ×2,K/2 未展开) ❌ K-axis 不匹配
fp4 + perm_x2=[0,2,1] [B, N, K/2] = [2, 32, 32] [0,2,1] [2, 32, 64](最后维 K/2 被 ×2 = K) ✅ 正常
fp8 + perm_x2=[0,1,2] [B, K, N] = [2, 64, 32] [0,1,2] [2, 64, 32](无展开) ✅ 正常
fp8 + perm_x2=[0,2,1] [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-653
  • 问题算子 csrc:matmul/transpose_quant_batch_mat_mul/torch_extension/csrc/transpose_quant_batch_mat_mul.cpp:136-139
  • 问题算子 Python:matmul/transpose_quant_batch_mat_mul/torch_extension/transpose_quant_batch_mat_mul.py
  • 上一轮修复记录:matmul/transpose_quant_batch_mat_mul/torch_extension/FIX_dtype_check.md:350-353
  • aclnn 算子约束文档:matmul/transpose_quant_batch_mat_mul/docs/aclnnTransposeQuantBatchMatMul.md:366

Environment / 环境信息 (Mandatory / 必填)

Ascend950

Steps to reproduce the issue / 重现步骤 (Mandatory / 必填)

构建tqbmm mxfp4用例 mxfp4 + perm_x2=[0,1,2]场景,用例报错显示 k维度未对齐

Describe the expected behavior / 预期结果 (Mandatory / 必填)

用例pass

不涉及

Special notes for this issue/备注 (Optional / 选填)

likedislike
link164735link164735成员
8月17日 添加了label:bug-report
yuning_chenyuning_chen成员
8月17日 将 Hu1L1 设为负责人
yuning_chen
yuning_chen成员
8月17日 评论:

您好,感谢反馈,问题已收到,当前 @Hu1L1 正在跟踪处理。

likedislike
丛吉钰
丛吉钰成员
8月18日 评论:

/asign

likedislike
丛吉钰
丛吉钰成员
8月18日 评论:

/assign

likedislike
CANN-robotCANN-robot成员
8月18日 将 cong-jiyu 设为负责人,移除负责人 Hu1L1
丛吉钰丛吉钰成员
8月25日 issue状态由 进行中 改变为 已解决
丛吉钰丛吉钰成员
8月25日 关闭了 issue
CANN-robotCANN-robot成员
8月25日 添加了label:resolved