已合并
aclnnTransposeBatchMatMul支持B、K非对齐,PermX2[0,2,1],B*K>65535等场景 #4240
ww-blue创建于 4月25日
aclnnTransposeBatchMatMul支持B、K非对齐,PermX2[0,2,1],B*K>65535等场景 #4240
已合并
ww-blue创建于 4月25日
共 7 个文件变更+159-115
@@ -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]。
ztj0209
ztj0209ztj02094月29日

🔵 [优化] aclnnTransposeBatchMatMul.md 文档中 950PR/950DT 平台的 permX2 支持范围与 A2/A3 不一致

文件: matmul/transpose_batch_mat_mul/docs/aclnnTransposeBatchMatMul.md

文档中:

  • A2/A3 平台:permX2 支持 [0,1,2]、[0,2,1](本次 PR 新增)
  • 950PR/950DT 平台:permX2 支持 [0,1,2]、[0,2,1](原有支持)

但代码中 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] 时是否有额外约束,避免用户混淆。

likedislike
ww-blue
ww-blue
4月29日 评论:
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)
133namespace {133namespace {
134static bool CheckPermLimit(const aclIntArray* perm_x1, const aclIntArray* perm_x2, const aclIntArray* perm_y)134static 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) ||
唐
唐唐超4月29日

为啥不用驼峰?代码风格都不太对的,后需要要统一整改

likedislike
ww-blue
ww-blue
4月29日 评论:
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} // namespace155} // namespace
164 156 
165static bool CheckShapeValid(const aclTensor* x1, const aclTensor* x2, const aclTensor* scale,157static 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+ 
ztj0209
ztj0209ztj02094月29日

🔴 [严重] CheckShapeValid 中 x2_need_transpose 使用了 perm_x1 而非 perm_x2

文件: op_api/aclnn_transpose_batch_mat_mul.cpp

auto x2_need_transpose = ((*perm_x1)[0] == 0 && (*perm_x1)[1] == 2 && (*perm_x1)[2] == 1);

此处变量名为 x2_need_transpose,但实际使用的是 perm_x1 而非 perm_x2。如果这是笔误,则 permX2=[0,2,1] 的校验将完全失效,x2 的转置判断实际上检查的是 x1 的 perm 值。

如果意图是"当 x1 的 perm 为 [0,2,1] 时",则变量名 x2_need_transpose 严重误导,应改为 x1_perm_021 或类似名称。

建议: 确认此处语义:

  • 如果是笔误,修正为 (*perm_x2):
auto x2_need_transpose = ((*perm_x2)[0] == 0 && (*perm_x2)[1] == 2 && (*perm_x2)[2] == 1);
  • 如果是有意为之(即"当 x1 perm 为 [0,2,1] 时触发 x2 转置限制"),则必须更改变量名避免混淆。
likedislike
ww-blue
ww-blue
4月29日 评论:
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) ||
ztj0209
ztj0209ztj02094月29日

🟡 [建议] pp_matmul_common_tiling.h 运算符优先级修复正确,但建议增加单元测试

文件: op_host/op_tiling/pp_matmul_common_tiling.h

// 修改前(运算符优先级错误):
platformType == platform_ascendc::SocVersion::ASCEND910 && basicBlockSize > UB_LIMIT_SIZE_910A ||
platformType == platform_ascendc::SocVersion::ASCEND310P && isPertokenArch20 == true && 
basicBlockSize > UB_LIMIT_SIZE_PERTOKEN_ARCH20

// 修改后(正确添加括号):
(platformType == platform_ascendc::SocVersion::ASCEND910 && basicBlockSize > UB_LIMIT_SIZE_910A) ||
(platformType == platform_ascendc::SocVersion::ASCEND310P && isPertokenArch20 == true && 
basicBlockSize > UB_LIMIT_SIZE_PERTOKEN_ARCH20)

修复正确!原代码中 && 和 || 的优先级导致逻辑错误:A && B || C && D 被解析为 (A && B) || (C && D),恰好与预期一致,但如果 isPertokenArch20 条件更复杂则可能出错。添加括号使意图更明确。

建议: 此类运算符优先级 bug 应添加单元测试覆盖,确保在不同 platformType 和 basicBlockSize 组合下 IsExceedTilingLimit 返回正确结果。

likedislike
ww-blue
ww-blue
4月29日 评论:
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 
133template <bool PRI_FLAG, typename OpShareType>133template <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,
ztj0209
ztj0209ztj02094月29日

🔵 [优化] transpose_batch_mat_mul_base_tiling.cpp 移除 support_k_n 校验后,错误信息可能不够准确

文件: op_host/op_tiling/transpose_batch_mat_mul_base_tiling.cpp

// 修改前:
OP_TILING_CHECK(!(support_k_n && support_batch_m), CUBE_INNER_ERR_REPORT(args_.opName,
    "only support shape inner axis < 65536, input shape b, m, n, k = %lu, %lu, %lu, %lu.",
    input_batch, args_.mValue, input_n, input_k), return ge::GRAPH_FAILED);

// 修改后:
OP_TILING_CHECK(!support_batch_m, CUBE_INNER_ERR_REPORT(args_.opName,
    "only support shape inner axis < 65536, input shape b, m, n, k = %lu, %lu, %lu, %lu.",
    input_batch, args_.mValue, input_n, input_k), return ge::GRAPH_FAILED);

移除 support_k_n 校验(即 K/N 不再要求 16 对齐)是本次 PR 的核心改动之一。但错误信息 "only support shape inner axis < 65536" 仍然保留,可能让用户误以为 K/N 对齐仍是限制条件。

建议: 更新错误信息,使其准确反映当前的约束条件,例如 "only support batch * m < 65536 for transposed input".

likedislike
ww-blue
ww-blue
4月29日 评论:
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 @@
31using namespace optiling::transpose_batch_mat_mul;31using namespace optiling::transpose_batch_mat_mul;
32 32 
33namespace {33namespace {
34-constexpr uint64_t NO_BATCH_DIM_SUM = 2;34+ 
35} // namespace35} // namespace
36 36 
37namespace optiling {37namespace optiling {
ztj0209
ztj0209ztj02094月29日

🟡 [建议] CheckInputArgs 与 IsPpMatmulEinsumMode 存在重复的 attrs 校验逻辑

文件: op_host/op_tiling/transpose_batch_mat_mul_einsum_tiling.cpp

CheckInputArgs 函数校验了:

  • bias/scale 不存在
  • enableHF32=false
  • batchSplitFactor=1

IsPpMatmulEinsumMode 在调用 CheckInputArgs 之后,又独立校验了 permX1=[1,0,2]。

但 CheckInputArgs 中的 bias/scale 校验与 IsPpMatmulEinsumMode 中 B*K>65535 场景的限制(不支持 scale、batchSplitFactor=1、permX1=[1,0,2])语义重叠但不完全一致。

建议:

  • 将 CheckInputArgs 合并到 IsPpMatmulEinsumMode 中,或将其重命名为更明确的名称如 CheckEinsumBasicConstraints
  • 在 IsPpMatmulEinsumMode 中增加注释,说明不同分支(B*K>65535 vs permX2=[0,2,1] vs FP32)各自需要满足的约束条件
likedislike
ww-blue
ww-blue
4月29日 评论:
38namespace transpose_batch_mat_mul {38namespace 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) {
ztj0209
ztj0209ztj02094月29日

🟡 [建议] IsPpMatmulEinsumMode 中 K >= SUPPORTED_INNER_AXIS 条件判断与文档不一致

文件: op_host/op_tiling/transpose_batch_mat_mul_einsum_tiling.cpp

if (K >= SUPPORTED_INNER_AXIS || Batch * K >= SUPPORTED_INNER_AXIS) {

SUPPORTED_INNER_AXIS = 65536,所以 K >= 65536 等价于 K > 65535。但文档中写的是"K <= 65535"(即 K 最大为 65535),而 CheckShapeValid 中是 x1KDim >= SUPPORTED_INNER_AXIS(即 K >= 65536 时报错)。

在 IsPpMatmulEinsumMode 中,K >= SUPPORTED_INNER_AXIS 意味着当 K >= 65536 时路由到 Einsum 模式。但根据文档,K > 65535 时应该报错而非路由到 Einsum。

建议: 确认 K >= 65536 是否为合法场景。如果 K 的上限确实是 65535,则 IsPpMatmulEinsumMode 中 K >= SUPPORTED_INNER_AXIS 这个条件永远不会为真(因为 CheckShapeValid 已经拦截),可以移除或添加注释说明这是防御性编程。

likedislike
ww-blue
ww-blue
4月29日 评论:
99+ OP_LOGI(context->GetNodeName(), "When K > 65535 or Batch*K > 65535, hit PpMatmulEinsum.");
100+ return ge::GRAPH_SUCCESS;
ztj0209
ztj0209ztj02094月29日

🔴 [严重] IsPpMatmulEinsumMode 中 B*K>65535 时直接路由到 Einsum,但未校验 scale/batchSplitFactor 限制

文件: op_host/op_tiling/transpose_batch_mat_mul_einsum_tiling.cpp

if (K >= SUPPORTED_INNER_AXIS || Batch * K >= SUPPORTED_INNER_AXIS) {
    OP_LOGI(context->GetNodeName(), "When K > 65535 or Batch*K > 65535, hit PpMatmulEinsum.");
    return ge::GRAPH_SUCCESS;  // 直接进入 Einsum 模式
}

当 Batch * K >= 65536 时,IsPpMatmulEinsumMode 直接返回 GRAPH_SUCCESS 路由到 Einsum tiling。但根据 PR 描述和 CheckShapeValid 中的约束,当 B*K > 65535 时,不支持传入 scale,且 batchSplitFactor 只能等于 1,permX1 必须为 [1,0,2]。

然而 IsPpMatmulEinsumMode 在此分支中没有校验 scale 和 batchSplitFactor。虽然 CheckInputArgs 中校验了 bias/scale 为空和 batchSplitFactor=1,但 CheckInputArgs 是在 IsPpMatmulEinsumMode 内部被调用的——如果 CheckInputArgs 返回 GRAPH_FAILED,则整个 IsPpMatmulEinsumMode 也返回 GRAPH_FAILED,不会进入 Einsum 模式。

但问题在于:CheckInputArgs 中对 scale 的校验是 GetOptionalInputShape(INDEX_SCALE) != nullptr,这只检查了 scale 是否存在,没有检查 permX1 是否为 [1,0,2]。如果用户传入 B*K>65535 + scale=nullptr + permX1=[0,1,2],CheckInputArgs 会通过,但 CheckShapeValid 会拒绝——两个路径的校验逻辑不一致。

建议: 在 IsPpMatmulEinsumMode 的 B*K>65535 分支中,增加 permX1=[1,0,2] 的校验,确保 tiling 路由和 API 层校验一致。

likedislike
ww-blue
ww-blue
4月29日 评论:
ww-blue
ww-blue
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()
ztj0209
ztj0209ztj02094月29日

🟡 [建议] GetCloseKShiftFlag 使用环境变量 CLOSE_MATMUL_K_SHIFT 控制功能开关,存在安全风险

文件: op_host/op_tiling/transpose_batch_mat_mul_einsum_tiling.cpp

const char* closeKShiftFlag = std::getenv("CLOSE_MATMUL_K_SHIFT");
if (closeKShiftFlag != nullptr && strcmp(closeKShiftFlag, "1") == 0) {
    return true;
}

通过环境变量控制算子行为存在以下问题:

  1. 不可追踪: 环境变量是运行时外部状态,难以在日志或调试中追踪
  2. 并发不安全: 多个算子实例共享同一环境变量,无法对不同算子设置不同行为
  3. 安全风险: 恶意用户可通过设置环境变量改变算子计算行为

建议:

  • 如果此开关仅用于调试/测试,添加注释说明不应在生产环境使用
  • 如果需要长期保留,考虑通过 op_attr 或 context 传入而非环境变量
  • 至少添加 OP_LOGW 警告日志,提示环境变量已被设置
likedislike
ww-blue
ww-blue
4月29日 评论:
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 
53bool TransposeBatchMatMulEinsumTiling::GetMatMulInfo()158bool 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);
ztj0209
ztj0209ztj02094月29日

🟡 [建议] PostTiling 中 workspace 大小硬编码为 24MB,缺乏依据说明

文件: op_host/op_tiling/transpose_batch_mat_mul_einsum_tiling.cpp

size_t sysWorkspaceSize = static_cast<size_t>(24 * 1024 * 1024);  // 24M same as ppmatmul tiling

workspace 大小硬编码为 24MB,注释仅说"same as ppmatmul tiling"。不同场景(B*K>65535、permX2=[0,2,1])可能需要不同的 workspace 大小。

建议:

  • 将 24 * 1024 * 1024 提取为命名常量,如 constexpr size_t EINSUM_DEFAULT_WORKSPACE_SIZE = 24UL * 1024 * 1024;
  • 补充注释说明 24MB 的计算依据(是否与最大 batch/K 组合相关)
  • 考虑是否需要根据实际输入 shape 动态计算 workspace 大小
likedislike
ww-blue
ww-blue
4月29日 评论:
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 
24namespace optiling {24namespace optiling {
25namespace transpose_batch_mat_mul {25namespace transpose_batch_mat_mul {
26+ 
27+ge::graphStatus IsPpMatmulEinsumMode(gert::TilingContext* context);
28+ 
26inline int64_t PermDecode(const int64_t* perm, size_t size)29inline 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+ 
34class TransposeBatchMatMulEinsumTiling : public pp_matmul::PpMatMulDefault38class TransposeBatchMatMulEinsumTiling : public pp_matmul::PpMatMulDefault
35{39{
36public:40public:
@@ -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 
37REGISTER_TILING_TEMPLATE("TransposeBatchMatMul", TransposeBatchMatMulBaseTiling, 0);37REGISTER_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- 
106static ge::graphStatus TransposeBatchMatMulTilingFunc(gert::TilingContext* context) {39static 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}