已合并
添加bmmv3 vector kernel拦截条件 #7264
zhang-junming21创建于 7月9日
添加bmmv3 vector kernel拦截条件 #7264
已合并
zhang-junming21创建于 7月9日
3 个文件变更+26-11
@@ -621,8 +621,8 @@ bool BatchMatmulV3BaseTiling::DoMultiBatchOutTiling()
621void BatchMatmulV3BaseTiling::DoMultiBatchTiling()621void 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 bias625+ 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()
735void BatchMatmulV3BaseTiling::DoMultiBatchL1FullLoadTiling()735void 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) { // 暂时不支持bias739 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+ 
1162bool BatchMatmulV3BaseTiling::CheckVectorShapeDims()1172bool 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 optiling1343} // 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 
57enum class TilingEnableMultiBatch : int32_t // 互斥flag, 对应不同全载模板选择57enum 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 
64enum class TilingEnableLoadMode : int32_t // 互斥flag, 对应不同全载模板选择64enum 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 
73enum class TilingEnableMultiBatchOut : int32_t // 互斥flag, 对应不同全载模板选择73enum 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 
80enum class TilingEnableMixNd2Nz : int32_t // 互斥flag, 对应不同全载模板选择80enum 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 
87struct TilingEnable {87struct TilingEnable {
@@ -157,6 +157,7 @@ protected:
157private:157private:
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 }