| @@ -156,7 +156,11 @@ aclnnStatus aclnnTransposeBatchMatMul( | |||
| 156 | <td>permX2</td> | 156 | <td>permX2</td> |
| 157 | <td>输入</td> | 157 | <td>输入</td> |
| 158 | <td>表示矩阵乘的第二个矩阵的转置序列,host侧的aclIntArray。</td> | 158 | <td>表示矩阵乘的第二个矩阵的转置序列,host侧的aclIntArray。</td> |
| 159 | - <td>-</td> | 159 | + <td> |
| 160 | + <ul> | ||
| 161 | + <li> 支持[0, 1, 2]、[0, 2, 1]。</li> | ||
| 162 | + </ul> | ||
| 163 | + </td> | ||
| 160 | <td>INT64</td> | 164 | <td>INT64</td> |
| 161 | <td>-</td> | 165 | <td>-</td> |
| 162 | <td>1</td> | 166 | <td>1</td> |
| @@ -324,9 +328,8 @@ aclnnStatus aclnnTransposeBatchMatMul( | |||
| 324 | 328 | ||
| 325 | - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>: | 329 | - <term>Atlas A2 训练系列产品/Atlas A2 推理系列产品</term>、<term>Atlas A3 训练系列产品/Atlas A3 推理系列产品</term>: |
| 326 | - B的取值范围为[1, 65536),N的取值范围为[1, 65536)。 | 330 | - B的取值范围为[1, 65536),N的取值范围为[1, 65536)。 |
| 327 | - - 当x1的输入shape为(B, M, K)时,K <= 65535;当x1的输入shape为(M, B, K)时,B * K <= 65535。 | 331 | + - 当x1的输入shape为(B, M, K)时,需要K <= 65535; 当x1的输入shape为(M, B, K)且B * K > 65535时,不支持传入scale,并且batchSplitFactor只能等于1, permX1必须为[1, 0, 2]。 |
| 328 | - - x2的第二维或x2的第三维不能被16整除。 | 332 | + - 当permX2输入为[0, 2, 1]时,不支持传入scale,并且batchSplitFactor只能等于1, permX1必须为[1, 0, 2]。 |
| 329 | - - permX2仅支持输入[0, 1, 2]。 | ||
| 330 | - 当scale不为空时,batchSplitFactor只能等于1,B与N的乘积小于65536, 且仅支持输入为FLOAT16和输出为INT8的类型推导。 | 333 | - 当scale不为空时,batchSplitFactor只能等于1,B与N的乘积小于65536, 且仅支持输入为FLOAT16和输出为INT8的类型推导。 |
| 331 | - <term>Ascend 950PR/Ascend 950DT</term>: | 334 | - <term>Ascend 950PR/Ascend 950DT</term>: |
| 332 | - permX2支持输入[0, 1, 2]、[0, 2, 1]。 | 335 | - permX2支持输入[0, 1, 2]、[0, 2, 1]。 |
| @@ -133,28 +133,20 @@ inline static bool CheckScaleValid(const aclTensor* scale, int64_t batchN) | |||
| 133 | namespace { | 133 | namespace { |
| 134 | static bool CheckPermLimit(const aclIntArray* perm_x1, const aclIntArray* perm_x2, const aclIntArray* perm_y) | 134 | static bool CheckPermLimit(const aclIntArray* perm_x1, const aclIntArray* perm_x2, const aclIntArray* perm_y) |
| 135 | { | 135 | { |
| 136 | - auto x1_need_transpose = ((*perm_x1)[0] == 1 && (*perm_x1)[1] == 0 && (*perm_x1)[2] == 2) || | 136 | + auto x1_support_transpose = ((*perm_x1)[0] == 1 && (*perm_x1)[1] == 0 && (*perm_x1)[2] == 2) || |
| 137 | ((*perm_x1)[0] == 0 && (*perm_x1)[1] == 1 && (*perm_x1)[2] == 2); | 137 | ((*perm_x1)[0] == 0 && (*perm_x1)[1] == 1 && (*perm_x1)[2] == 2); |
| 138 | - auto x2_need_transpose = ((*perm_x2)[0] == 0 && (*perm_x2)[1] == 1 && (*perm_x2)[2] == 2); | 138 | + auto x2_support_transpose = ((*perm_x2)[0] == 0 && (*perm_x2)[1] == 1 && (*perm_x2)[2] == 2) || |
| 139 | - auto y_need_transpose = ((*perm_y)[0] == 1 && (*perm_y)[1] == 0 && (*perm_y)[2] == 2); | 139 | + ((*perm_x2)[0] == 0 && (*perm_x2)[1] == 2 && (*perm_x2)[2] == 1); |
| 140 | - std::string permX2ErrorInfo = "[0, 1, 2]."; | 140 | + auto y_support_transpose = ((*perm_y)[0] == 1 && (*perm_y)[1] == 0 && (*perm_y)[2] == 2); |
| 141 | - if (GetCurrentPlatformInfo().GetCurNpuArch() == NpuArch::DAV_3510) { | 141 | + if (!x1_support_transpose) { |
| 142 | - // For DAV-3510 architecture, the perm tensor for x2 operand only supports two patterns: | 142 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "the perm of x1 should be [0, 1, 2] or [1, 0, 2]."); |
| 143 | - // Pattern 1: [0, 1, 2] - No transpose (identity) | ||
| 144 | - // Pattern 2: [0, 2, 1] - Transpose last two dimensions (swap dim1 and dim2) | ||
| 145 | - // Any other permutation pattern will cause hardware compatibility issues. | ||
| 146 | - x2_need_transpose = x2_need_transpose || ((*perm_x2)[0] == 0 && (*perm_x2)[1] == 2 && (*perm_x2)[2] == 1); | ||
| 147 | - permX2ErrorInfo = "[0, 1, 2] or [0, 2, 1]."; | ||
| 148 | - } | ||
| 149 | - if (!x1_need_transpose) { | ||
| 150 | - OP_LOGE(ACLNN_ERR_PARAM_INVALID, "the perm of x1 should be [0,1,2] or [1, 0, 2]."); | ||
| 151 | return false; | 143 | return false; |
| 152 | } | 144 | } |
| 153 | - if (!x2_need_transpose) { | 145 | + if (!x2_support_transpose) { |
| 154 | - OP_LOGE(ACLNN_ERR_PARAM_INVALID, "the perm of x2 should be %s", permX2ErrorInfo.c_str()); | 146 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "the perm of x2 should be [0, 1, 2] or [0, 2, 1]."); |
| 155 | return false; | 147 | return false; |
| 156 | } | 148 | } |
| 157 | - if (!y_need_transpose) { | 149 | + if (!y_support_transpose) { |
| 158 | OP_LOGE(ACLNN_ERR_PARAM_INVALID, "the perm of y should be [1, 0, 2]."); | 150 | OP_LOGE(ACLNN_ERR_PARAM_INVALID, "the perm of y should be [1, 0, 2]."); |
| 159 | return false; | 151 | return false; |
| 160 | } | 152 | } |
| @@ -163,7 +155,7 @@ static bool CheckPermLimit(const aclIntArray* perm_x1, const aclIntArray* perm_x | |||
| 163 | } // namespace | 155 | } // namespace |
| 164 | 156 | ||
| 165 | static bool CheckShapeValid(const aclTensor* x1, const aclTensor* x2, const aclTensor* scale, | 157 | static bool CheckShapeValid(const aclTensor* x1, const aclTensor* x2, const aclTensor* scale, |
| 166 | - const aclIntArray* perm_x1, const aclIntArray* perm_x2, const aclIntArray* perm_y) | 158 | + const aclIntArray* perm_x1, const aclIntArray* perm_x2, const aclIntArray* perm_y, int32_t batch_split_factor) |
| 167 | { | 159 | { |
| 168 | if (!CheckPermLimit(perm_x1, perm_x2, perm_y)) { | 160 | if (!CheckPermLimit(perm_x1, perm_x2, perm_y)) { |
| 169 | OP_LOGE(ACLNN_ERR_PARAM_INVALID, "the perm is invalid."); | 161 | OP_LOGE(ACLNN_ERR_PARAM_INVALID, "the perm is invalid."); |
| @@ -181,23 +173,24 @@ static bool CheckShapeValid(const aclTensor* x1, const aclTensor* x2, const aclT | |||
| 181 | op::ToString(x1Shape).GetString(), op::ToString(x2Shape).GetString()); | 173 | op::ToString(x1Shape).GetString(), op::ToString(x2Shape).GetString()); |
| 182 | return false; | 174 | return false; |
| 183 | } | 175 | } |
| 176 | + | ||
🔴 [严重] 文件:
此处变量名为 如果意图是"当 x1 的 perm 为 [0,2,1] 时",则变量名 建议: 确认此处语义:
![]() ![]() | |||
| 184 | if (GetCurrentPlatformInfo().GetCurNpuArch() != NpuArch::DAV_3510) { | 177 | if (GetCurrentPlatformInfo().GetCurNpuArch() != NpuArch::DAV_3510) { |
| 185 | - auto x1_need_transpose = ((*perm_x1)[0] == 1 && (*perm_x1)[1] == 0 && (*perm_x1)[2] == 2); | 178 | + auto x1_need_transpose = ((*perm_x1)[0] == 1 && (*perm_x1)[1] == 0 && (*perm_x1)[2] == 2); |
| 186 | - if (x1_need_transpose && x1->GetViewShape().GetDim(1) * x1KDim >= SUPPORTED_INNER_AXIS) { | 179 | + auto x2_need_transpose = ((*perm_x2)[0] == 0 && (*perm_x2)[1] == 2 && (*perm_x2)[2] == 1); |
| 187 | - OP_LOGE(ACLNN_ERR_PARAM_INVALID, "batch mul k should be less than 65536."); | ||
| 188 | - return false; | ||
| 189 | - } | ||
| 190 | - if (!x1_need_transpose && x1KDim >= SUPPORTED_INNER_AXIS) { | ||
| 191 | - OP_LOGE(ACLNN_ERR_PARAM_INVALID, "K should be less than 65536."); | ||
| 192 | - return false; | ||
| 193 | - } | ||
| 194 | if (N >= SUPPORTED_INNER_AXIS || batchNum >= SUPPORTED_INNER_AXIS) { | 180 | if (N >= SUPPORTED_INNER_AXIS || batchNum >= SUPPORTED_INNER_AXIS) { |
| 195 | - OP_LOGE(ACLNN_ERR_PARAM_INVALID, "batch and n should be less than 65536."); | 181 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "Batch and N should be less than 65536."); |
| 196 | return false; | 182 | return false; |
| 197 | } | 183 | } |
| 198 | - if ((x2KDim % BLOCK_SIZE != 0) || (N % BLOCK_SIZE != 0)) { | 184 | + if (!x1_need_transpose && x1KDim >= SUPPORTED_INNER_AXIS) { |
| 199 | - OP_LOGE(ACLNN_ERR_PARAM_INVALID, | 185 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "When x1 is not transposed([B,M,K]), K should be less than 65536."); |
| 200 | - "The shape of the x2 is not supported, now they are %ld, %ld and %ld", batchNum, x2KDim, N); | 186 | + return false; |
| 187 | + } | ||
| 188 | + if (x1_need_transpose && batchNum * x1KDim >= SUPPORTED_INNER_AXIS && (scale != nullptr || batch_split_factor != 1)) { | ||
| 189 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "When x1 is transposed([M,B,K]) and Batch * K > 65535, input scale or batch_split_factor != 1 are not supported."); | ||
| 190 | + return false; | ||
| 191 | + } | ||
| 192 | + if (x2_need_transpose && (scale != nullptr || batch_split_factor != 1 || !x1_need_transpose)) { | ||
| 193 | + OP_LOGE(ACLNN_ERR_PARAM_INVALID, "When x2 is transposed, input scale or batch_split_factor != 1 are not supported, permX1 must be [1,0,2]."); | ||
| 201 | return false; | 194 | return false; |
| 202 | } | 195 | } |
| 203 | } | 196 | } |
| @@ -257,7 +250,7 @@ inline static aclnnStatus CheckParams(const aclTensor* x1, const aclTensor* x2, | |||
| 257 | } | 250 | } |
| 258 | 251 | ||
| 259 | CHECK_RET(CheckDtypeValid(x1, x2, scale, out), ACLNN_ERR_PARAM_INVALID); | 252 | CHECK_RET(CheckDtypeValid(x1, x2, scale, out), ACLNN_ERR_PARAM_INVALID); |
| 260 | - CHECK_RET(CheckShapeValid(x1, x2, scale, perm_x1, perm_x2, perm_y), ACLNN_ERR_PARAM_INVALID); | 253 | + CHECK_RET(CheckShapeValid(x1, x2, scale, perm_x1, perm_x2, perm_y, batch_split_factor), ACLNN_ERR_PARAM_INVALID); |
| 261 | 254 | ||
| 262 | if (batch_split_factor <= 0) { | 255 | if (batch_split_factor <= 0) { |
| 263 | OP_LOGE(ACLNN_ERR_PARAM_INVALID, "batch_split_factor[%d] should be greater than 0.", | 256 | OP_LOGE(ACLNN_ERR_PARAM_INVALID, "batch_split_factor[%d] should be greater than 0.", |
| @@ -125,9 +125,9 @@ inline bool IsExceedTilingLimit(uint64_t axes0, uint64_t priAxes0, | |||
| 125 | uint64_t basicBlockSize, const bool isPertokenArch20) | 125 | uint64_t basicBlockSize, const bool isPertokenArch20) |
| 126 | { | 126 | { |
| 127 | return (PRI_FLAG && axes0 > n0TilingLimit) || (!PRI_FLAG && priAxes0 > n0TilingLimit) || | 127 | return (PRI_FLAG && axes0 > n0TilingLimit) || (!PRI_FLAG && priAxes0 > n0TilingLimit) || |
| 128 | - (platformType == platform_ascendc::SocVersion::ASCEND910 && basicBlockSize > UB_LIMIT_SIZE_910A || | 128 | + ((platformType == platform_ascendc::SocVersion::ASCEND910 && basicBlockSize > UB_LIMIT_SIZE_910A) || |
🟡 [建议] 文件:
修复正确!原代码中 建议: 此类运算符优先级 bug 应添加单元测试覆盖,确保在不同 platformType 和 basicBlockSize 组合下 ![]() ![]() | |||
| 129 | - platformType == platform_ascendc::SocVersion::ASCEND310P && isPertokenArch20 == true && | 129 | + (platformType == platform_ascendc::SocVersion::ASCEND310P && isPertokenArch20 == true && |
| 130 | - basicBlockSize > UB_LIMIT_SIZE_PERTOKEN_ARCH20); | 130 | + basicBlockSize > UB_LIMIT_SIZE_PERTOKEN_ARCH20)); |
| 131 | } | 131 | } |
| 132 | 132 | ||
| 133 | template <bool PRI_FLAG, typename OpShareType> | 133 | template <bool PRI_FLAG, typename OpShareType> |
| @@ -615,7 +615,6 @@ ge::graphStatus TransposeBatchMatMulBaseTiling::GetShape() | |||
| 615 | uint64_t input_batch = batchInfo_.batchC; | 615 | uint64_t input_batch = batchInfo_.batchC; |
| 616 | uint64_t input_k = args_.kValue; | 616 | uint64_t input_k = args_.kValue; |
| 617 | uint64_t input_n = args_.nValue; | 617 | uint64_t input_n = args_.nValue; |
| 618 | - bool support_k_n = (input_k % BLOCK_CUBE == 0UL) && (input_n % BLOCK_CUBE == 0UL); | ||
| 619 | bool support_batch_m = true; | 618 | bool support_batch_m = true; |
| 620 | 619 | ||
| 621 | if (transA_ == 213UL) { //213 指的是 {1,0,2} 转置 | 620 | if (transA_ == 213UL) { //213 指的是 {1,0,2} 转置 |
| @@ -627,7 +626,7 @@ ge::graphStatus TransposeBatchMatMulBaseTiling::GetShape() | |||
| 627 | return ge::GRAPH_FAILED; | 626 | return ge::GRAPH_FAILED; |
| 628 | } | 627 | } |
| 629 | 628 | ||
| 630 | - OP_TILING_CHECK(!(support_k_n && support_batch_m), CUBE_INNER_ERR_REPORT(args_.opName, | 629 | + OP_TILING_CHECK(!support_batch_m, CUBE_INNER_ERR_REPORT(args_.opName, |
🔵 [优化] 文件:
移除 建议: 更新错误信息,使其准确反映当前的约束条件,例如 ![]() ![]() | |||
| 631 | "only support shape inner axis < 65536, input shape b, m, n, k = %lu, %lu, %lu, %lu.", | 630 | "only support shape inner axis < 65536, input shape b, m, n, k = %lu, %lu, %lu, %lu.", |
| 632 | input_batch, args_.mValue, input_n, input_k ), return ge::GRAPH_FAILED); | 631 | input_batch, args_.mValue, input_n, input_k ), return ge::GRAPH_FAILED); |
| 633 | return ge::GRAPH_SUCCESS; | 632 | return ge::GRAPH_SUCCESS; |
| @@ -31,13 +31,113 @@ | |||
| 31 | using namespace optiling::transpose_batch_mat_mul; | 31 | using namespace optiling::transpose_batch_mat_mul; |
| 32 | 32 | ||
| 33 | namespace { | 33 | namespace { |
| 34 | -constexpr uint64_t NO_BATCH_DIM_SUM = 2; | 34 | + |
| 35 | } // namespace | 35 | } // namespace |
| 36 | 36 | ||
| 37 | namespace optiling { | 37 | namespace optiling { |
🟡 [建议] 文件:
但 建议:
![]() ![]() | |||
| 38 | namespace transpose_batch_mat_mul { | 38 | namespace transpose_batch_mat_mul { |
| 39 | 39 | ||
| 40 | -void TransposeBatchMatMulEinsumTiling::DoTiling() | 40 | +inline static ge::graphStatus CheckInputArgs(gert::TilingContext* context){ |
| 41 | + constexpr size_t INDEX_X1 = 0; | ||
| 42 | + constexpr size_t INDEX_X2 = 1; | ||
| 43 | + constexpr size_t INDEX_BIAS = 2; | ||
| 44 | + constexpr size_t INDEX_SCALE = 3; | ||
| 45 | + constexpr size_t INDEX_ENABLE_HF32 = 3; | ||
| 46 | + constexpr size_t INDEX_BATCHSPLIT_FACTOR = 4; | ||
| 47 | + constexpr size_t ATTR_NUM = 5; | ||
| 48 | + constexpr bool EINSUM_SUPPORT_ENABLE_HF32 = false; | ||
| 49 | + constexpr int32_t EINSUM_SUPPORT_BATCHSPLIT_FACTOR = 1; | ||
| 50 | + OPS_CHECK_NULL_WITH_CONTEXT(context, context->GetInputDesc(INDEX_X1)); | ||
| 51 | + OPS_CHECK_NULL_WITH_CONTEXT(context, context->GetInputShape(INDEX_X1)); | ||
| 52 | + OPS_CHECK_NULL_WITH_CONTEXT(context, context->GetInputDesc(INDEX_X2)); | ||
| 53 | + OPS_CHECK_NULL_WITH_CONTEXT(context, context->GetInputShape(INDEX_X2)); | ||
| 54 | + OP_TILING_CHECK((context->GetOptionalInputShape(INDEX_BIAS)!= nullptr || | ||
| 55 | + context->GetOptionalInputShape(INDEX_SCALE)!= nullptr), | ||
| 56 | + OP_LOGI(context->GetNodeName(), "PpMatmulEinsum not support bias or scale."), | ||
| 57 | + return ge::GRAPH_FAILED); | ||
| 58 | + OPS_CHECK_NULL_WITH_CONTEXT(context, context->GetOutputDesc(0)); | ||
| 59 | + auto attrs = context->GetAttrs(); | ||
| 60 | + OPS_CHECK_NULL_WITH_CONTEXT(context, attrs); | ||
| 61 | + if (attrs->GetAttrNum() >= ATTR_NUM) { | ||
| 62 | + OP_TILING_CHECK( | ||
| 63 | + (*(attrs->GetAttrPointer<bool>(INDEX_ENABLE_HF32)) != EINSUM_SUPPORT_ENABLE_HF32), | ||
| 64 | + OP_LOGI(context->GetNodeName(), "PpMatmulEinsum only support ENABLE_HF32=false."), | ||
| 65 | + return ge::GRAPH_FAILED); | ||
| 66 | + OP_TILING_CHECK( | ||
| 67 | + (*(attrs->GetAttrPointer<int32_t>(INDEX_BATCHSPLIT_FACTOR)) != EINSUM_SUPPORT_BATCHSPLIT_FACTOR), | ||
| 68 | + OP_LOGI(context->GetNodeName(), "PpMatmulEinsum only support batch_split_factor=1."), | ||
| 69 | + return ge::GRAPH_FAILED); | ||
| 70 | + } | ||
| 71 | + return ge::GRAPH_SUCCESS; | ||
| 72 | +} | ||
| 73 | + | ||
| 74 | +ge::graphStatus IsPpMatmulEinsumMode(gert::TilingContext* context) | ||
| 75 | +{ | ||
| 76 | + constexpr size_t INDEX_PERM_X1 = 0; | ||
| 77 | + constexpr size_t INDEX_PERM_X2 = 1; | ||
| 78 | + constexpr size_t BATCH_IDX = 0; | ||
| 79 | + constexpr size_t K_IDX = 2; | ||
| 80 | + constexpr size_t ALLOW_DIM = 3; | ||
| 81 | + constexpr int64_t SUPPORTED_INNER_AXIS = 65536; | ||
| 82 | + if (CheckInputArgs(context) != ge::GRAPH_SUCCESS){ | ||
| 83 | + OP_LOGI(context->GetNodeName(), "Current scenario is not support PpMatmulEinsum."); | ||
| 84 | + return ge::GRAPH_FAILED; | ||
| 85 | + } | ||
| 86 | + auto attrs = context->GetAttrs(); | ||
| 87 | + auto x1PermList = attrs->GetAttrPointer<gert::ContinuousVector>(INDEX_PERM_X1); | ||
| 88 | + auto x2PermList = attrs->GetAttrPointer<gert::ContinuousVector>(INDEX_PERM_X2); | ||
| 89 | + const int64_t* perm_x1 = reinterpret_cast<const int64_t*>(x1PermList->GetData()); | ||
| 90 | + const int64_t* perm_x2 = reinterpret_cast<const int64_t*>(x2PermList->GetData()); | ||
| 91 | + OP_TILING_CHECK((PermDecode(perm_x1, x1PermList->GetSize()) != 213L), | ||
| 92 | + OP_LOGI(context->GetNodeName(), "PpMatmulEinsum only support permA={1,0,2}."), | ||
| 93 | + return ge::GRAPH_FAILED); | ||
| 94 | + const gert::Shape &x1Shape = context->GetInputShape(0)->GetOriginShape(); | ||
| 95 | + const gert::Shape &x2Shape = context->GetInputShape(1)->GetOriginShape(); | ||
| 96 | + int64_t Batch = x1Shape[perm_x1[BATCH_IDX]]; | ||
| 97 | + int64_t K = x1Shape[perm_x1[K_IDX]]; | ||
| 98 | + if (K >= SUPPORTED_INNER_AXIS || Batch * K >= SUPPORTED_INNER_AXIS) { | ||
🟡 [建议] 文件:
在 建议: 确认 ![]() ![]() | |||
| 99 | + OP_LOGI(context->GetNodeName(), "When K > 65535 or Batch*K > 65535, hit PpMatmulEinsum."); | ||
| 100 | + return ge::GRAPH_SUCCESS; | ||
🔴 [严重] 文件:
当 然而 但问题在于: 建议: 在 ![]() ![]() ww-blue 4月29日 评论: 4月29日 评论: | |||
| 101 | + } | ||
| 102 | + if (PermDecode(perm_x2, x2PermList->GetSize()) == 132L) { | ||
| 103 | + OP_LOGI(context->GetNodeName(), "When PermX2 = [0, 2, 1], hit PpMatmulEinsum."); | ||
| 104 | + return ge::GRAPH_SUCCESS; | ||
| 105 | + } | ||
| 106 | + const size_t x1DimNum = x1Shape.GetDimNum(); | ||
| 107 | + const size_t x2DimNum = x2Shape.GetDimNum(); | ||
| 108 | + if ((x1DimNum == ALLOW_DIM) && (x2DimNum == ALLOW_DIM)) { | ||
| 109 | + if ((context->GetInputDesc(0)->GetDataType() == ge::DT_FLOAT) && | ||
| 110 | + (context->GetInputDesc(1)->GetDataType() == ge::DT_FLOAT) && | ||
| 111 | + (context->GetOutputDesc(0)->GetDataType() == ge::DT_FLOAT)) { | ||
| 112 | + OP_LOGI(context->GetNodeName(), "When input and output are Fp32, hit PpMatmulEinsum."); | ||
| 113 | + return ge::GRAPH_SUCCESS; | ||
| 114 | + } | ||
| 115 | + } | ||
| 116 | + OP_LOGI(context->GetNodeName(), "Current scenario is not support PpMatmulEinsum."); | ||
| 117 | + return ge::GRAPH_FAILED; | ||
| 118 | +} | ||
| 119 | + | ||
| 120 | +bool GetCloseKShiftFlag(gert::TilingContext* context) | ||
| 121 | +{ | ||
| 122 | + if (context->GetDeterministicLevel() == INT32_MAX) { | ||
| 123 | + OP_LOGE(context->GetNodeName(), "GetDeterministicLevel() failed."); | ||
| 124 | + OP_LOGI(context->GetNodeName(), "GetCloseKShiftFlag: false"); | ||
| 125 | + return false; | ||
| 126 | + } | ||
| 127 | + if (context->GetDeterministic() == 1 && context->GetDeterministicLevel() > 1) { | ||
| 128 | + OP_LOGI(context->GetNodeName(), "GetCloseKShiftFlag: true"); | ||
| 129 | + return true; | ||
| 130 | + } | ||
| 131 | + const char* closeKShiftFlag = std::getenv("CLOSE_MATMUL_K_SHIFT"); | ||
| 132 | + if (closeKShiftFlag != nullptr && strcmp(closeKShiftFlag, "1") == 0) { | ||
| 133 | + OP_LOGI(context->GetNodeName(), "GetCloseKShiftFlag: true"); | ||
| 134 | + return true; | ||
| 135 | + } | ||
| 136 | + OP_LOGI(context->GetNodeName(), "GetCloseKShiftFlag: false"); | ||
| 137 | + return false; | ||
| 138 | +} | ||
| 139 | + | ||
| 140 | +ge::graphStatus TransposeBatchMatMulEinsumTiling::DoTiling() | ||
🟡 [建议] 文件:
通过环境变量控制算子行为存在以下问题:
建议:
![]() ![]() | |||
| 41 | { | 141 | { |
| 42 | auto inputDType = context_->GetInputDesc(0)->GetDataType(); | 142 | auto inputDType = context_->GetInputDesc(0)->GetDataType(); |
| 43 | matMulInfo_.sizeInDtype = ge::GetSizeByDataType(inputDType); | 143 | matMulInfo_.sizeInDtype = ge::GetSizeByDataType(inputDType); |
| @@ -47,11 +147,17 @@ void TransposeBatchMatMulEinsumTiling::DoTiling() | |||
| 47 | (void)GetMatMulInfo(); | 147 | (void)GetMatMulInfo(); |
| 48 | (void)GetTilingKey(); | 148 | (void)GetTilingKey(); |
| 49 | (void)GetMatMulTilingData(); | 149 | (void)GetMatMulTilingData(); |
| 150 | + ge::graphStatus ret = PostTiling(); | ||
| 151 | + if (ret != ge::GRAPH_SUCCESS) { | ||
| 152 | + return ret; | ||
| 153 | + } | ||
| 50 | PrintTiling(); | 154 | PrintTiling(); |
| 155 | + return ge::GRAPH_SUCCESS; | ||
| 51 | } | 156 | } |
| 52 | 157 | ||
| 53 | bool TransposeBatchMatMulEinsumTiling::GetMatMulInfo() | 158 | bool TransposeBatchMatMulEinsumTiling::GetMatMulInfo() |
| 54 | { | 159 | { |
| 160 | + constexpr uint64_t NO_BATCH_DIM_SUM = 2; | ||
| 55 | if (!isQuantBatchMatmulV3_) { | 161 | if (!isQuantBatchMatmulV3_) { |
| 56 | OP_TILING_CHECK(hardwareInfo_.socVersion != platform_ascendc::SocVersion::ASCEND910B, | 162 | OP_TILING_CHECK(hardwareInfo_.socVersion != platform_ascendc::SocVersion::ASCEND910B, |
| 57 | CUBE_INNER_ERR_REPORT(context_->GetNodeName(), "unsupported platform."), return false); | 163 | CUBE_INNER_ERR_REPORT(context_->GetNodeName(), "unsupported platform."), return false); |
| @@ -141,8 +247,14 @@ ge::graphStatus TransposeBatchMatMulEinsumTiling::PostTiling() | |||
| 141 | tilingData.blockDim = ppMatmulDefaultTilingData_.blockDim; | 247 | tilingData.blockDim = ppMatmulDefaultTilingData_.blockDim; |
| 142 | tilingData.swizzleDirect = ppMatmulDefaultTilingData_.swizzleDirect; | 248 | tilingData.swizzleDirect = ppMatmulDefaultTilingData_.swizzleDirect; |
| 143 | tilingData.splitk = ppMatmulDefaultTilingData_.splitk; | 249 | tilingData.splitk = ppMatmulDefaultTilingData_.splitk; |
| 144 | - tilingData.enShuffleK = ppMatmulDefaultTilingData_.enShuffleK; | 250 | + tilingData.enShuffleK = static_cast<uint32_t>(!GetCloseKShiftFlag(context_)); |
| 145 | - OPS_CHECK_NULL_WITH_CONTEXT(context_, context_->GetRawTilingData()); | 251 | + |
| 252 | + OP_TILING_CHECK(context_ == nullptr, | ||
| 253 | + CUBE_INNER_ERR_REPORT("TbmmEinsum", "context is null"), return ge::GRAPH_FAILED); | ||
| 254 | + size_t sysWorkspaceSize = static_cast<size_t>(24 * 1024 * 1024); // 24M same as ppmatmul tiling | ||
| 255 | + size_t* currentWorkSpace = context_->GetWorkspaceSizes(1); | ||
🟡 [建议] 文件:
workspace 大小硬编码为 24MB,注释仅说"same as ppmatmul tiling"。不同场景(B*K>65535、permX2=[0,2,1])可能需要不同的 workspace 大小。 建议:
![]() ![]() | |||
| 256 | + currentWorkSpace[0] = sysWorkspaceSize; | ||
| 257 | + | ||
| 146 | errno_t ret = memcpy_s(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity(), | 258 | errno_t ret = memcpy_s(context_->GetRawTilingData()->GetData(), context_->GetRawTilingData()->GetCapacity(), |
| 147 | reinterpret_cast<void *>(&tilingData), tilingDataSize); | 259 | reinterpret_cast<void *>(&tilingData), tilingDataSize); |
| 148 | if (ret != EOK){ | 260 | if (ret != EOK){ |
| @@ -23,6 +23,9 @@ | |||
| 23 | 23 | ||
| 24 | namespace optiling { | 24 | namespace optiling { |
| 25 | namespace transpose_batch_mat_mul { | 25 | namespace transpose_batch_mat_mul { |
| 26 | + | ||
| 27 | +ge::graphStatus IsPpMatmulEinsumMode(gert::TilingContext* context); | ||
| 28 | + | ||
| 26 | inline int64_t PermDecode(const int64_t* perm, size_t size) | 29 | inline int64_t PermDecode(const int64_t* perm, size_t size) |
| 27 | { | 30 | { |
| 28 | int64_t trans_ = 0L; | 31 | int64_t trans_ = 0L; |
| @@ -31,6 +34,7 @@ inline int64_t PermDecode(const int64_t* perm, size_t size) | |||
| 31 | } | 34 | } |
| 32 | return trans_; | 35 | return trans_; |
| 33 | } | 36 | } |
| 37 | + | ||
| 34 | class TransposeBatchMatMulEinsumTiling : public pp_matmul::PpMatMulDefault | 38 | class TransposeBatchMatMulEinsumTiling : public pp_matmul::PpMatMulDefault |
| 35 | { | 39 | { |
| 36 | public: | 40 | public: |
| @@ -43,7 +47,7 @@ public: | |||
| 43 | } | 47 | } |
| 44 | ~TransposeBatchMatMulEinsumTiling() override = default; | 48 | ~TransposeBatchMatMulEinsumTiling() override = default; |
| 45 | 49 | ||
| 46 | - void DoTiling(); | 50 | + ge::graphStatus DoTiling(); |
| 47 | bool GetMatMulInfo(); | 51 | bool GetMatMulInfo(); |
| 48 | bool GetTilingKey(); | 52 | bool GetTilingKey(); |
| 49 | ge::graphStatus PostTiling(); | 53 | ge::graphStatus PostTiling(); |
| @@ -36,81 +36,14 @@ namespace optiling { | |||
| 36 | 36 | ||
| 37 | REGISTER_TILING_TEMPLATE("TransposeBatchMatMul", TransposeBatchMatMulBaseTiling, 0); | 37 | REGISTER_TILING_TEMPLATE("TransposeBatchMatMul", TransposeBatchMatMulBaseTiling, 0); |
| 38 | 38 | ||
| 39 | -static ge::graphStatus IsEinsumMode(gert::TilingContext* context) | ||
| 40 | -{ | ||
| 41 | - constexpr size_t ALLOW_DIM = 3; | ||
| 42 | - constexpr size_t ATTR_NUM = 5; | ||
| 43 | - constexpr bool EINSUM_SUPPORT_ENABLE_HF32 = false; | ||
| 44 | - constexpr int32_t EINSUM_SUPPORT_BATCHSPLIT_FACTOR = 1; | ||
| 45 | - size_t idx = 0; | ||
| 46 | - OPS_CHECK_NULL_WITH_CONTEXT(context, context->GetInputDesc(idx)); | ||
| 47 | - OPS_CHECK_NULL_WITH_CONTEXT(context, context->GetInputShape(idx)); | ||
| 48 | - idx++; | ||
| 49 | - OPS_CHECK_NULL_WITH_CONTEXT(context, context->GetInputDesc(idx)); | ||
| 50 | - OPS_CHECK_NULL_WITH_CONTEXT(context, context->GetInputShape(idx)); | ||
| 51 | - idx++; | ||
| 52 | - OP_TILING_CHECK((context->GetOptionalInputShape(idx)!= nullptr || | ||
| 53 | - context->GetOptionalInputShape(idx + 1)!= nullptr), | ||
| 54 | - OP_LOGI(context->GetNodeName(), "Einsum mode not support bias or scale"), | ||
| 55 | - return ge::GRAPH_FAILED); | ||
| 56 | - OPS_CHECK_NULL_WITH_CONTEXT(context, context->GetOutputDesc(0)); | ||
| 57 | - auto attrs = context->GetAttrs(); | ||
| 58 | - OPS_CHECK_NULL_WITH_CONTEXT(context, attrs); | ||
| 59 | - idx = static_cast<size_t>(0); | ||
| 60 | - auto aPermList_ = attrs->GetAttrPointer<gert::ContinuousVector>(idx); | ||
| 61 | - const int64_t* perm_x1 = reinterpret_cast<const int64_t*>(aPermList_->GetData()); | ||
| 62 | - OP_TILING_CHECK((PermDecode(perm_x1, aPermList_->GetSize()) != 213L), | ||
| 63 | - OP_LOGI(context->GetNodeName(), "Einsum mode only support permA={1,0,2}"), | ||
| 64 | - return ge::GRAPH_FAILED); | ||
| 65 | - if (attrs->GetAttrNum() >= ATTR_NUM) { | ||
| 66 | - idx++; | ||
| 67 | - OP_TILING_CHECK( | ||
| 68 | - (*(attrs->GetAttrPointer<int32_t>(ATTR_NUM - idx)) != EINSUM_SUPPORT_BATCHSPLIT_FACTOR), | ||
| 69 | - OP_LOGI(context->GetNodeName(), "Einsum mode only support batch_split_factor=1"), | ||
| 70 | - return ge::GRAPH_FAILED); | ||
| 71 | - idx++; | ||
| 72 | - OP_TILING_CHECK( | ||
| 73 | - (*(attrs->GetAttrPointer<bool>(ATTR_NUM - idx)) != EINSUM_SUPPORT_ENABLE_HF32), | ||
| 74 | - OP_LOGI(context->GetNodeName(), "Einsum mode only support ENABLE_HF32=false"), | ||
| 75 | - return ge::GRAPH_FAILED); | ||
| 76 | - } | ||
| 77 | - const gert::Shape &aShape = context->GetInputShape(0)->GetOriginShape(); | ||
| 78 | - const gert::Shape &bShape = context->GetInputShape(1)->GetOriginShape(); | ||
| 79 | - const size_t aDimNum = aShape.GetDimNum(); | ||
| 80 | - const size_t bDimNum = bShape.GetDimNum(); | ||
| 81 | - if ((aDimNum == ALLOW_DIM) && (bDimNum == ALLOW_DIM)) { | ||
| 82 | - if ((context->GetInputDesc(0)->GetDataType() == ge::DT_FLOAT) && | ||
| 83 | - (context->GetInputDesc(1)->GetDataType() == ge::DT_FLOAT) && | ||
| 84 | - (context->GetOutputDesc(0)->GetDataType() == ge::DT_FLOAT)) { | ||
| 85 | - return ge::GRAPH_SUCCESS; | ||
| 86 | - } | ||
| 87 | - } | ||
| 88 | - return ge::GRAPH_FAILED; | ||
| 89 | -} | ||
| 90 | - | ||
| 91 | -static ge::graphStatus TbmmEinsumTilingFunc(gert::TilingContext* context) | ||
| 92 | -{ | ||
| 93 | - OP_TILING_CHECK(context == nullptr, | ||
| 94 | - CUBE_INNER_ERR_REPORT("TbmmEinsum", "context is null"), return ge::GRAPH_FAILED); | ||
| 95 | - size_t sysWorkspaceSize = static_cast<size_t>(24 * 1024 * 1024); // 24M same as ppmatmul tiling | ||
| 96 | - size_t* currentWorkSpace = context->GetWorkspaceSizes(1); | ||
| 97 | - currentWorkSpace[0] = sysWorkspaceSize; | ||
| 98 | - OP_LOGI(context->GetNodeName(), "TbmmEinsum Tiling start."); | ||
| 99 | - TransposeBatchMatMulEinsumTiling tbmmEinsumTiling(context); | ||
| 100 | - tbmmEinsumTiling.DoTiling(); | ||
| 101 | - auto res = tbmmEinsumTiling.PostTiling(); | ||
| 102 | - OP_LOGI(context->GetNodeName(), "TbmmEinsum Tiling end."); | ||
| 103 | - return res; | ||
| 104 | -} | ||
| 105 | - | ||
| 106 | static ge::graphStatus TransposeBatchMatMulTilingFunc(gert::TilingContext* context) { | 39 | static ge::graphStatus TransposeBatchMatMulTilingFunc(gert::TilingContext* context) { |
| 107 | OP_TILING_CHECK(context == nullptr, CUBE_INNER_ERR_REPORT("TransposeBatchMatMul", "context is null"), | 40 | OP_TILING_CHECK(context == nullptr, CUBE_INNER_ERR_REPORT("TransposeBatchMatMul", "context is null"), |
| 108 | return ge::GRAPH_FAILED); | 41 | return ge::GRAPH_FAILED); |
| 109 | if (IsAdvancedSocVersion(context)) { | 42 | if (IsAdvancedSocVersion(context)) { |
| 110 | return transpose_batch_mat_mul_advanced::TransposeBatchMatMulTiling(context).DoTiling(); | 43 | return transpose_batch_mat_mul_advanced::TransposeBatchMatMulTiling(context).DoTiling(); |
| 111 | } | 44 | } |
| 112 | - if (IsEinsumMode(context) == ge::GRAPH_SUCCESS) { | 45 | + if (IsPpMatmulEinsumMode(context) == ge::GRAPH_SUCCESS) { |
| 113 | - return TbmmEinsumTilingFunc(context); | 46 | + return TransposeBatchMatMulEinsumTiling(context).DoTiling(); |
| 114 | } | 47 | } |
| 115 | return TilingRegistry::GetInstance().DoTilingImpl(context); | 48 | return TilingRegistry::GetInstance().DoTilingImpl(context); |
| 116 | } | 49 | } |


🔵 [优化]
aclnnTransposeBatchMatMul.md文档中 950PR/950DT 平台的 permX2 支持范围与 A2/A3 不一致文件:
matmul/transpose_batch_mat_mul/docs/aclnnTransposeBatchMatMul.md文档中:
但代码中
CheckPermLimit已经移除了DAV_3510的平台判断,统一支持 [0,1,2] 和 [0,2,1]。这意味着文档中 950PR/950DT 的 permX2 支持描述与代码一致,但 A2/A3 的约束条件(不支持 scale、batchSplitFactor=1、permX1=[1,0,2])在 950PR/950DT 上是否同样适用?文档中未明确说明。建议: 在文档中明确 950PR/950DT 平台使用 permX2=[0,2,1] 时是否有额外约束,避免用户混淆。