已合并
增加aclnnBaddbmmToFmm拦截条件 #7390
wuyufei创建于 7月10日
增加aclnnBaddbmmToFmm拦截条件 #7390
已合并
共 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); | ||
| 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 | } |
🟡 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)。