文件最后提交记录最后更新时间
1 个月前
4 个月前
2 个月前
7 个月前
1 个月前
1 个月前
2 个月前
7 个月前
2 个月前
README

TransposeBatchMatMul

产品支持情况

产品 是否支持
Ascend 950PR/Ascend 950DT
Atlas A3 训练系列产品/Atlas A3 推理系列产品
Atlas A2 训练系列产品/Atlas A2 推理系列产品
Atlas 200I/500 A2 推理产品 ×
Atlas 推理系列产品 ×
Atlas 训练系列产品 ×
Kirin X90 处理器系列产品
Kirin 9030 处理器系列产品

功能说明

  • 算子功能:完成张量x1与张量x2的矩阵乘计算。仅支持三维的Tensor传入。Tensor支持转置,转置序列根据传入的序列进行变更。permX1代表张量x1的转置序列,支持[0,1,2]、[1,0,2],permX2代表张量x2的转置序列[0,1,2],permY表示矩阵乘输出矩阵的转置序列,当前仅支持[1,0,2],序列值为0的是batch维度,其余两个维度做矩阵乘法。scale表示输出矩阵的量化系数,可在输入为FLOAT16且输出为INT8时开启,详细约束条件可见约束说明或者aclnnTransposeBatchMatMul调用说明文档。

  • 示例:

    • x1的shape是(B, M, K),x2的shape是(B, K, N),scale为None,batchSplitFactor等于1时,计算输出out的shape是(M, B, N)。
    • x1的shape是(B, M, K),x2的shape是(B, K, N),scale不为None,batchSplitFactor等于1时,计算输出out的shape是(M, 1, B * N)。
    • x1的shape是(B, M, K),x2的shape是(B, K, N),scale为None,batchSplitFactor大于1时,计算输出out的shape是(batchSplitFactor, M, B * N / batchSplitFactor)。

参数说明

参数名 输入/输出/属性 描述 数据类型 数据格式
x1 输入 矩阵乘运算中的左矩阵。 FLOAT32, FLOAT16, BF16 ND
x2 输入 矩阵乘运算中的右矩阵。 FLOAT32, FLOAT16, BF16 ND
bias 输入 矩阵乘运算后累加的偏置。 FLOAT32, FLOAT16, BF16 ND
scale 输入 量化参数的缩放因子。 INT64, UINT64 ND
permX1 输入 表示矩阵乘的第一个矩阵的转置序列。 INT64 -
permX2 输入 表示矩阵乘的第二个矩阵的转置序列。 INT64 -
permY 输入 表示矩阵乘输出矩阵的转置序列。 INT64 -
cubeMathType 输入 指定Cube单元的计算逻辑。 INT8 -
batchSplitFactor 输入 用于指定矩阵乘输出矩阵中B维的切分大小。 INT32 -
y 输出 矩阵乘运算的计算结果。 FLOAT32, FLOAT16, BF16, INT8 ND
  • Kirin X90/Kirin 9030 处理器系列产品:不支持BFLOAT16。
  • Atlas A2 训练系列产品/Atlas A2 推理系列产品:只有输入x2支持FRACTAL_NZ格式。
  • Atlas A3 训练系列产品/Atlas A3 推理系列产品:只有输入x2支持FRACTAL_NZ格式。
  • Ascend 950PR/Ascend 950DT:只有输入x2支持FRACTAL_NZ格式。

约束说明

  • Atlas A2 训练系列产品/Atlas A2 推理系列产品、Atlas A3 训练系列产品/Atlas A3 推理系列产品:
    • 不支持空tensor。
    • 支持非连续tensor。
    • B的取值范围为[1, 65536),N的取值范围为[1, 65536)。
    • 当x1的输入shape为(B, M, K)时,K <= 65535;当x1的输入shape为(M, B, K)时,B * K <= 65535。
    • 当scale不为空时,batchSplitFactor只能等于1,B与N的乘积小于65536,且仅支持输入为FLOAT16和输出为INT8的类型推导。
  • Ascend 950PR/Ascend 950DT:
    • 当scale不为空时,batchSplitFactor只能等于1,且仅支持输入为FLOAT16和输出为INT8的类型推导。
    • bias为预留参数,当前暂不支持。

调用说明

调用方式 样例代码 说明
aclnn接口 test_aclnn_transpose_batch_mat_mul 通过
- aclnnTransposeBatchMatMul
- aclnnTransposeBatchMatMulWeightNz
等方式调用TransposeBatchMatMul算子。