已合并
添加bmmv3 vector kernel拦截条件 #7264
zhang-junming21创建于 7月9日
添加bmmv3 vector kernel拦截条件 #7264
已合并
共 3 个文件变更+26-11
| @@ -621,8 +621,8 @@ bool BatchMatmulV3BaseTiling::DoMultiBatchOutTiling() | |||
| 621 | void BatchMatmulV3BaseTiling::DoMultiBatchTiling() | 621 | void BatchMatmulV3BaseTiling::DoMultiBatchTiling() |
| 622 | { | 622 | { |
| 623 | bool isEqualBatch = batchInfo_.batchA0 == batchInfo_.batchB0 && batchInfo_.batchA1 == batchInfo_.batchB1 && | 623 | bool isEqualBatch = batchInfo_.batchA0 == batchInfo_.batchB0 && batchInfo_.batchA1 == batchInfo_.batchB1 && |
| 624 | - batchInfo_.batchA2 == batchInfo_.batchB2 && batchInfo_.batchA3 == batchInfo_.batchB3; //广播 | 624 | + batchInfo_.batchA2 == batchInfo_.batchB2 && batchInfo_.batchA3 == batchInfo_.batchB3; // 广播 |
| 625 | - if (!isEqualBatch || (args_.hasBias && !batchInfo_.biasWithBatch)) { //不支持broadcast\非多batch bias | 625 | + if (!isEqualBatch || (args_.hasBias && !batchInfo_.biasWithBatch)) { // 不支持broadcast\非多batch bias |
| 626 | return; | 626 | return; |
| 627 | } | 627 | } |
| 628 | uint64_t shapeM = ops::CeilAlign(static_cast<uint64_t>(bmmTilingData_.matmulTiling.matmulTiling.M), BLOCK_CUBE); | 628 | uint64_t shapeM = ops::CeilAlign(static_cast<uint64_t>(bmmTilingData_.matmulTiling.matmulTiling.M), BLOCK_CUBE); |
| @@ -735,7 +735,7 @@ void BatchMatmulV3BaseTiling::DoMultiBatchL1FullLoadTilingImpl() | |||
| 735 | void BatchMatmulV3BaseTiling::DoMultiBatchL1FullLoadTiling() | 735 | void BatchMatmulV3BaseTiling::DoMultiBatchL1FullLoadTiling() |
| 736 | { | 736 | { |
| 737 | bool isEqualBatch = batchInfo_.batchA0 == batchInfo_.batchB0 && batchInfo_.batchA1 == batchInfo_.batchB1 && | 737 | bool isEqualBatch = batchInfo_.batchA0 == batchInfo_.batchB0 && batchInfo_.batchA1 == batchInfo_.batchB1 && |
| 738 | - batchInfo_.batchA2 == batchInfo_.batchB2 && batchInfo_.batchA3 == batchInfo_.batchB3; //广播 | 738 | + batchInfo_.batchA2 == batchInfo_.batchB2 && batchInfo_.batchA3 == batchInfo_.batchB3; // 广播 |
| 739 | if (!isEqualBatch || args_.hasBias) { // 暂时不支持bias | 739 | if (!isEqualBatch || args_.hasBias) { // 暂时不支持bias |
| 740 | return; | 740 | return; |
| 741 | } | 741 | } |
| @@ -1159,6 +1159,16 @@ bool BatchMatmulV3BaseTiling::CheckVectorNpuArch() | |||
| 1159 | return true; | 1159 | return true; |
| 1160 | } | 1160 | } |
| 1161 | 1161 | ||
| 1162 | +bool BatchMatmulV3BaseTiling::CheckVectorTranspose() | ||
| 1163 | +{ | ||
| 1164 | + if (args_.isATrans || !args_.isBTrans) { | ||
| 1165 | + OP_LOGD(args_.opName, "BatchMatmulV3BaseTiling: transA should be false, transB should be true. " | ||
| 1166 | + "Bmm vector opt version not supported."); | ||
| 1167 | + return false; | ||
| 1168 | + } | ||
| 1169 | + return true; | ||
| 1170 | +} | ||
| 1171 | + | ||
| 1162 | bool BatchMatmulV3BaseTiling::CheckVectorShapeDims() | 1172 | bool BatchMatmulV3BaseTiling::CheckVectorShapeDims() |
| 1163 | { | 1173 | { |
| 1164 | auto aShape = context_->GetInputShape(0)->GetOriginShape(); | 1174 | auto aShape = context_->GetInputShape(0)->GetOriginShape(); |
| @@ -1251,6 +1261,9 @@ bool BatchMatmulV3BaseTiling::CheckVectorComputationCondition() | |||
| 1251 | if (!CheckVectorNpuArch()) { | 1261 | if (!CheckVectorNpuArch()) { |
| 1252 | return false; | 1262 | return false; |
| 1253 | } | 1263 | } |
| 1264 | + if (!CheckVectorTranspose()) { | ||
| 1265 | + return false; | ||
| 1266 | + } | ||
| 1254 | if (!CheckVectorShapeDims()) { | 1267 | if (!CheckVectorShapeDims()) { |
| 1255 | return false; | 1268 | return false; |
| 1256 | } | 1269 | } |
| @@ -1328,4 +1341,4 @@ void BatchMatmulV3BaseTiling::DoTilingKeyCustom() | |||
| 1328 | } | 1341 | } |
| 1329 | 1342 | ||
| 1330 | } // namespace optiling | 1343 | } // namespace optiling |
| 1331 | -} | 1344 | +} |
| @@ -40,7 +40,7 @@ struct BatchShapeInfo { | |||
| 40 | bool biasWithBatch = false; | 40 | bool biasWithBatch = false; |
| 41 | }; | 41 | }; |
| 42 | 42 | ||
| 43 | -enum class TilingCalcSelect //选择不同的计算Tiling的方法 | 43 | +enum class TilingCalcSelect // 选择不同的计算Tiling的方法 |
| 44 | { | 44 | { |
| 45 | ALL = 0, | 45 | ALL = 0, |
| 46 | COMMON = 1, | 46 | COMMON = 1, |
| @@ -51,14 +51,14 @@ enum class TilingEnableMultiBatchL1FullLoad : int32_t // 互斥flag, 对应不 | |||
| 51 | { | 51 | { |
| 52 | IS_FALSE = 0, | 52 | IS_FALSE = 0, |
| 53 | IS_TRUE = 1, | 53 | IS_TRUE = 1, |
| 54 | - MAX = 10 //模板类别不能超过10个 | 54 | + MAX = 10 // 模板类别不能超过10个 |
| 55 | }; | 55 | }; |
| 56 | 56 | ||
| 57 | enum class TilingEnableMultiBatch : int32_t // 互斥flag, 对应不同全载模板选择 | 57 | enum class TilingEnableMultiBatch : int32_t // 互斥flag, 对应不同全载模板选择 |
| 58 | { | 58 | { |
| 59 | IS_FALSE = 0, | 59 | IS_FALSE = 0, |
| 60 | IS_TRUE = 1, | 60 | IS_TRUE = 1, |
| 61 | - MAX = 10 //模板类别不能超过10个 | 61 | + MAX = 10 // 模板类别不能超过10个 |
| 62 | }; | 62 | }; |
| 63 | 63 | ||
| 64 | enum class TilingEnableLoadMode : int32_t // 互斥flag, 对应不同全载模板选择 | 64 | enum class TilingEnableLoadMode : int32_t // 互斥flag, 对应不同全载模板选择 |
| @@ -67,21 +67,21 @@ enum class TilingEnableLoadMode : int32_t // 互斥flag, 对应不同全载模 | |||
| 67 | AL1_FULL_LOAD = 1, | 67 | AL1_FULL_LOAD = 1, |
| 68 | BL1_FULL_LOAD = 2, | 68 | BL1_FULL_LOAD = 2, |
| 69 | VECTOR_FULL_LOAD = 3, | 69 | VECTOR_FULL_LOAD = 3, |
| 70 | - MAX = 10 //模板类别不能超过10个 | 70 | + MAX = 10 // 模板类别不能超过10个 |
| 71 | }; | 71 | }; |
| 72 | 72 | ||
| 73 | enum class TilingEnableMultiBatchOut : int32_t // 互斥flag, 对应不同全载模板选择 | 73 | enum class TilingEnableMultiBatchOut : int32_t // 互斥flag, 对应不同全载模板选择 |
| 74 | { | 74 | { |
| 75 | IS_FALSE = 0, | 75 | IS_FALSE = 0, |
| 76 | IS_TRUE = 1, | 76 | IS_TRUE = 1, |
| 77 | - MAX = 10 //模板类别不能超过10个 | 77 | + MAX = 10 // 模板类别不能超过10个 |
| 78 | }; | 78 | }; |
| 79 | 79 | ||
| 80 | enum class TilingEnableMixNd2Nz : int32_t // 互斥flag, 对应不同全载模板选择 | 80 | enum class TilingEnableMixNd2Nz : int32_t // 互斥flag, 对应不同全载模板选择 |
| 81 | { | 81 | { |
| 82 | IS_TRUE = 0, | 82 | IS_TRUE = 0, |
| 83 | IS_FALSE = 1, | 83 | IS_FALSE = 1, |
| 84 | - MAX = 10 //模板类别不能超过10个 | 84 | + MAX = 10 // 模板类别不能超过10个 |
| 85 | }; | 85 | }; |
| 86 | 86 | ||
| 87 | struct TilingEnable { | 87 | struct TilingEnable { |
| @@ -157,6 +157,7 @@ protected: | |||
| 157 | private: | 157 | private: |
| 158 | // CheckVectorComputationCondition的子检查函数 | 158 | // CheckVectorComputationCondition的子检查函数 |
| 159 | bool CheckVectorNpuArch(); | 159 | bool CheckVectorNpuArch(); |
| 160 | + bool CheckVectorTranspose(); | ||
| 160 | bool CheckVectorShapeDims(); | 161 | bool CheckVectorShapeDims(); |
| 161 | bool CheckVectorDtypeAndKAxis(); | 162 | bool CheckVectorDtypeAndKAxis(); |
| 162 | bool CheckVectorBatchBroadcast(); | 163 | bool CheckVectorBatchBroadcast(); |
| @@ -249,7 +249,8 @@ static bool CheckAscendCScenario(const aclTensor* x1, const aclTensor* x2, const | |||
| 249 | int64_t mDim = x1->GetViewShape().GetDim(x1DimNum - 2); | 249 | int64_t mDim = x1->GetViewShape().GetDim(x1DimNum - 2); |
| 250 | int64_t kDim = x1->GetViewShape().GetDim(x1DimNum - 1); | 250 | int64_t kDim = x1->GetViewShape().GetDim(x1DimNum - 1); |
| 251 | int64_t MAX_MK_DIM = 8; | 251 | int64_t MAX_MK_DIM = 8; |
| 252 | - if (npuArch == NpuArch::DAV_2201 && adjX2 == 1 && nDim == 1 && mDim <= MAX_MK_DIM && kDim <= MAX_MK_DIM) { | 252 | + if (npuArch == NpuArch::DAV_2201 && adjX1 == 0 && adjX2 == 1 && nDim == 1 && mDim <= MAX_MK_DIM && |
| 253 | + kDim <= MAX_MK_DIM) { | ||
| 253 | OP_LOGI("Hit batch_mat_mul_v3 vector kernel scenario: effective N dimension is 1."); | 254 | OP_LOGI("Hit batch_mat_mul_v3 vector kernel scenario: effective N dimension is 1."); |
| 254 | return true; | 255 | return true; |
| 255 | } | 256 | } |