已合并
增加aclnnBaddbmmToFmm拦截条件 #7390
wuyufei创建于 7月10日
增加aclnnBaddbmmToFmm拦截条件 #7390
已合并
wuyufei创建于 7月10日
共 1 个文件变更+16-0
@@ -1315,6 +1315,22 @@ bool checkFusedmm(
1315 if (!Ops::NN::IsTransposeNonContiguous(mat2, isNeedSwapInnerTwoDim)) {1315 if (!Ops::NN::IsTransposeNonContiguous(mat2, isNeedSwapInnerTwoDim)) {
1316 return false;1316 return false;
1317 }1317 }
1318+ // shape校验
1319+ const auto& selfShape = self->GetViewShape();
1320+ const auto& mat2Shape = mat2->GetViewShape();
1321+ const auto& biasShape = bias->GetViewShape();
1322+ int64_t aM = selfShape.GetDim(selfShape.GetDimNum() - NUM_TWO);
1323+ int64_t bN = mat2Shape.GetDim(mat2Shape.GetDimNum() - 1);
atomgit-bot
atomgit-botatomgit-bot7月10日

🟡 Medium Priority

第 1323 行,bN 被定义为 mat2Shape.GetDim(mat2Shape.GetDimNum() - 1),即始终取 mat2 的最后一维。但根据第 1314-1316 行的注释和逻辑,mat2 支持两种形式:(n,b,k) 和 (k,b,n)。当 isNeedSwapInnerTwoDim = true 时,mat2 的原始 shape 为 [B, N, K](最后两维需要交换),此时语义上的 N 维度位于倒数第二维(GetDimNum() - NUM_TWO),而非最后一维(GetDimNum() - 1)。

ExecFusedmmOp 函数(第 1346-1348 行)正是通过 SwapLastTwoDimValue 将最后两维交换来得到正确的 [B, K, N] 形式。但 checkFusedmm 中的 shape 校验未考虑这一点,导致当 isNeedSwapInnerTwoDim = true 时,bN 实际取到的是 K 维度,与 bias 的 N 维度比较会失败,从而错误地拒绝合法的配置(false negative)。

对于 isNeedSwapInnerTwoDim = false 的情况(mat2 为 [B, K, N]),最后一维恰好是 N,此时 bN 取值正确。

建议:bN 应根据 isNeedSwapInnerTwoDim 决定从哪个维度读取:当需要交换内部两维时,N 在倒数第二维(GetDimNum() - NUM_TWO);否则 N 在最后一维(GetDimNum() - 1)。

likedislike
不准确?
wuyufei
7月11日 评论:
1324+ if (biasShape.GetDimNum() == NUM_TWO) {
1325+ if (aM != biasShape.GetDim(0) || bN != biasShape.GetDim(1)) {
1326+ return false;
1327+ }
1328+ } else {
1329+ if ((biasShape.GetDim(0) != 1 && biasShape.GetDim(0) != selfShape.GetDim(0)) || aM != biasShape.GetDim(1) ||
1330+ bN != biasShape.GetDim(NUM_TWO)) {
1331+ return false;
1332+ }
1333+ }
1318 OP_LOGI("Check fusedmm success.");1334 OP_LOGI("Check fusedmm success.");
1319 return true;1335 return true;
1320}1336}