已合并
extendconvtranspose support A16W8 #3472
sunlesheng创建于 4月2日
extendconvtranspose support A16W8 #3472
已合并
sunlesheng创建于 4月2日
11 个文件变更+265-103
@@ -317,54 +317,69 @@ bool Conv3DDXV2InnerProductTiling::CheckVecTransEnable(
317 return true;317 return true;
318}318}
319 319 
320-void Conv3DDXV2InnerProductTiling::SetTilingCondition(320+uint32_t Conv3DDXV2InnerProductTiling::GetLoadB1Condition()
321- const CoreTilingParams& coreParams, const L1TilingParams& l1Params, const L0TilingParams& l0Params)
322{321{
323- tilingRunInfo_.enableVecTransFlag = CheckVecTransEnable(coreParams, l1Params, l0Params);
324- loadB1Condition_ = 0;
325 if (tilingRunInfo_.enableC04Flag) {322 if (tilingRunInfo_.enableC04Flag) {
326- loadB1Condition_ = ENABLE_C04;323+ return ENABLE_C04;
327 } else if (tilingRunInfo_.tilingHkWkMode == TILING_HK) {324 } else if (tilingRunInfo_.tilingHkWkMode == TILING_HK) {
328- loadB1Condition_ = ENABLE_TILING_HK; // 表示load2b1时只加载wk325+ return ENABLE_TILING_HK; // 表示load2b1时只加载wk
329 } else if (tilingRunInfo_.tilingHkWkMode == TILING_HK_WK) {326 } else if (tilingRunInfo_.tilingHkWkMode == TILING_HK_WK) {
330- loadB1Condition_ = ENABLE_TILING_HK_WK; // 表示load2b1时hk wk均不加载,每次只加载hkwk=1的数据327+ return ENABLE_TILING_HK_WK; // 表示load2b1时hk wk均不加载,每次只加载hkwk=1的数据
328+ }
329+ return 0;
330+}
331+ 
332+uint32_t Conv3DDXV2InnerProductTiling::GetLoadB2Condition(
333+ const L1TilingParams& l1Params, const L0TilingParams& l0Params)
334+{
335+ if (IsSocVersionFuse(context_) && runInfo_.filterFormat == ge::FORMAT_FRACTAL_Z && runInfo_.groups == 1) {
336+ return B2_NO_TRANSPOSE_NO_REVERSE; // fractal_z格式不转置不逆序,通过fusspass做
331 }337 }
332 338 
333 a1DbFlag_ = l1Params.al1Pbuffer == DB_ON;339 a1DbFlag_ = l1Params.al1Pbuffer == DB_ON;
334 b1DbFlag_ = l1Params.bl1Pbuffer == DB_ON;340 b1DbFlag_ = l1Params.bl1Pbuffer == DB_ON;
335 if (tilingRunInfo_.enableC04Flag) {341 if (tilingRunInfo_.enableC04Flag) {
336- loadB2Condition_ = B2_NO_TRANSPOSE_NO_REVERSE; // 功能约束, 不转置不逆序342+ return B2_NO_TRANSPOSE_NO_REVERSE; // 功能约束, 不转置不逆序
337- return;
338 }343 }
339 344 
340 if (runInfo_.filterFormat == ge::FORMAT_DHWCN) {345 if (runInfo_.filterFormat == ge::FORMAT_DHWCN) {
341- loadB2Condition_ = B2_REVERSE_ONLY; // DHWCN只逆序不转置346+ return B2_REVERSE_ONLY; // DHWCN只逆序不转置
342- return;
343 }347 }
344 348 
345 if (groupConvMode_ == TILING_GROUP_MODE_ENLARGE || tilingRunInfo_.enableVecTransFlag ||349 if (groupConvMode_ == TILING_GROUP_MODE_ENLARGE || tilingRunInfo_.enableVecTransFlag ||
346 static_cast<int32_t>(dtypeByteL0b_) == ge::GetSizeByDataType(ge::DT_FLOAT) ||350 static_cast<int32_t>(dtypeByteL0b_) == ge::GetSizeByDataType(ge::DT_FLOAT) ||
347 static_cast<int32_t>(dtypeByteL0b_) == ge::GetSizeByDataType(ge::DT_HIFLOAT8)) {351 static_cast<int32_t>(dtypeByteL0b_) == ge::GetSizeByDataType(ge::DT_HIFLOAT8)) {
348- loadB2Condition_ = B2_REVERSE_ONLY; // 功能约束, 只逆序不转置352+ return B2_REVERSE_ONLY; // 功能约束, 只逆序不转置
349- return;
350 }353 }
351 354 
355+ return GetLoadB2ConditionByFormatAndKernel(l0Params);
356+}
357+ 
358+uint32_t Conv3DDXV2InnerProductTiling::GetLoadB2ConditionByFormatAndKernel(const L0TilingParams& l0Params)
359+{
352 uint32_t kernelHW = runInfo_.kernel_h * runInfo_.kernel_w;360 uint32_t kernelHW = runInfo_.kernel_h * runInfo_.kernel_w;
353 uint64_t kernelDHW = static_cast<uint64_t>(runInfo_.kernel_d) * kernelHW;361 uint64_t kernelDHW = static_cast<uint64_t>(runInfo_.kernel_d) * kernelHW;
354 if (kernelDHW == ONE_U64 && tilingRunInfo_.tilingHkWkMode == NO_TILING_HWK && groupConvMode_ == TILING_GROUP_MODE_ORIGIN) {362 if (kernelDHW == ONE_U64 && tilingRunInfo_.tilingHkWkMode == NO_TILING_HWK && groupConvMode_ == TILING_GROUP_MODE_ORIGIN) {
355- loadB2Condition_ = B2_TRANSPOSE_ONLY; // kernel为1, 不需要逆序,DHWCN除外363+ return B2_TRANSPOSE_ONLY; // kernel为1, 不需要逆序,DHWCN除外
356- return;
357 }364 }
358 365 
359 if (runInfo_.filterFormat == ge::FORMAT_NDHWC && l0Params.baseN >= kernelHW) {366 if (runInfo_.filterFormat == ge::FORMAT_NDHWC && l0Params.baseN >= kernelHW) {
360- loadB2Condition_ = B2_REVERSE_ONLY; // 性能优化分支,加快格式转换效率367+ return B2_REVERSE_ONLY; // 性能优化分支,加快格式转换效率
361 } else if (runInfo_.filterFormat == ge::FORMAT_NCDHW && kernelDHW * runInfo_.dedx_cin * dtypeByteL0b_ <= BYTE_64) {368 } else if (runInfo_.filterFormat == ge::FORMAT_NCDHW && kernelDHW * runInfo_.dedx_cin * dtypeByteL0b_ <= BYTE_64) {
362- loadB2Condition_ = B2_REVERSE_ONLY; // 性能优化分支,加快逆序效率369+ return B2_REVERSE_ONLY; // 性能优化分支,加快逆序效率
363 } else {370 } else {
364- loadB2Condition_ = B2_TRANSPOSE_AND_REVERSE;371+ return B2_TRANSPOSE_AND_REVERSE;
365 }372 }
366}373}
367 374 
375+void Conv3DDXV2InnerProductTiling::SetTilingCondition(
376+ const CoreTilingParams& coreParams, const L1TilingParams& l1Params, const L0TilingParams& l0Params)
377+{
378+ tilingRunInfo_.enableVecTransFlag = CheckVecTransEnable(coreParams, l1Params, l0Params);
379+ loadB1Condition_ = GetLoadB1Condition();
380+ loadB2Condition_ = GetLoadB2Condition(l1Params, l0Params);
381+}
382+ 
368void Conv3DDXV2InnerProductTiling::SetCommonTilingData(383void Conv3DDXV2InnerProductTiling::SetCommonTilingData(
369 const CoreTilingParams& coreParams, const L1TilingParams& l1Params, const L0TilingParams& l0Params)384 const CoreTilingParams& coreParams, const L1TilingParams& l1Params, const L0TilingParams& l0Params)
370{385{
@@ -105,6 +105,9 @@ private:
105 void AdjustBaseNWhenSmallM(uint32_t& baseN, uint32_t baseM, const L0TilingParams& l0Params, const TilingRunInfo& tilingRunInfo);105 void AdjustBaseNWhenSmallM(uint32_t& baseN, uint32_t baseM, const L0TilingParams& l0Params, const TilingRunInfo& tilingRunInfo);
106 uint32_t CalculateOptimalBaseK(uint32_t baseM, uint32_t baseN, const L0TilingParams& l0Params, const TilingRunInfo& tilingRunInfo);106 uint32_t CalculateOptimalBaseK(uint32_t baseM, uint32_t baseN, const L0TilingParams& l0Params, const TilingRunInfo& tilingRunInfo);
107 void UpdateL0CBufferMode(L0TilingParams& l0Params);107 void UpdateL0CBufferMode(L0TilingParams& l0Params);
108+ uint32_t GetLoadB1Condition();
109+ uint32_t GetLoadB2Condition(const L1TilingParams& l1Params, const L0TilingParams& l0Params);
110+ uint32_t GetLoadB2ConditionByFormatAndKernel(const L0TilingParams& l0Params);
108};111};
109 112 
110} // namespace Conv113} // namespace Conv
@@ -144,29 +144,29 @@ void Conv3DDXV2KernelSplitTiling::SetParamForKernelSplit(bool isKernelSplitOnlyH
144bool Conv3DDXV2KernelSplitTiling::CheckKernelSplitHW11Enable()144bool Conv3DDXV2KernelSplitTiling::CheckKernelSplitHW11Enable()
145{145{
146 uint64_t mValueForCheck = static_cast<uint64_t>(runInfo_.dedx_h) * runInfo_.dedx_w;146 uint64_t mValueForCheck = static_cast<uint64_t>(runInfo_.dedx_h) * runInfo_.dedx_w;
147- uint64_t nValueForCheck = static_cast<uint64_t>(runInfo_.dedx_cin1_g) * BLOCK_CUBE;147+ uint64_t nValueForCheck = static_cast<uint64_t>(runInfo_.dedx_cin1_g) * BLOCK_CUBE;
148- uint64_t kValueForCheck = runInfo_.dedy_cout1_g * BASIC_BLOCK_SIZE_32 / dtypeByteL0b_; 148+ uint64_t kValueForCheck = runInfo_.dedy_cout1_g * BASIC_BLOCK_SIZE_32 / dtypeByteL0b_;
149- if (runInfo_.dedx_cin <= BASIC_BLOCK_SIZE_32 || runInfo_.dedy_cout <= BASIC_BLOCK_SIZE_32) { // 小shape不准入149+ if (runInfo_.dedx_cin <= BASIC_BLOCK_SIZE_32 || runInfo_.dedy_cout <= BASIC_BLOCK_SIZE_32) { // 小shape不准入
150- return false;150+ return false;
151- }151+ }
152- // 判断是否能走进fullload tiling模板152+ // 判断是否能走进fullload tiling模板
153 bool fullLoadCondition = (runInfo_.kernel_d <= 1) &&153 bool fullLoadCondition = (runInfo_.kernel_d <= 1) &&
154 (mValueForCheck > nValueForCheck) &&154 (mValueForCheck > nValueForCheck) &&
155- (runInfo_.dedx_w <= static_cast<int32_t>(BASIC_BLOCK_SIZE_512));155+ (runInfo_.dedx_w <= static_cast<int32_t>(BASIC_BLOCK_SIZE_512));
156- if (fullLoadCondition) {156+ if (fullLoadCondition) {
157 bool isMTE2BoundThreshold = (mValueForCheck * nValueForCheck) < (mValueForCheck + nValueForCheck) * BASIC_BLOCK_SIZE_512;157 bool isMTE2BoundThreshold = (mValueForCheck * nValueForCheck) < (mValueForCheck + nValueForCheck) * BASIC_BLOCK_SIZE_512;
158- bool isFixpBoundThreshold = kValueForCheck <= BASIC_BLOCK_SIZE_128;158+ bool isFixpBoundThreshold = kValueForCheck <= BASIC_BLOCK_SIZE_128;
159- if (isMTE2BoundThreshold || isFixpBoundThreshold) {159+ if (isMTE2BoundThreshold || isFixpBoundThreshold) {
160- uint64_t bestBaseN = BASIC_BLOCK_SIZE_256;160+ uint64_t bestBaseN = BASIC_BLOCK_SIZE_256;
161- bestBaseN = std::min(bestBaseN, nValueForCheck);161+ bestBaseN = std::min(bestBaseN, nValueForCheck);
162- if (kValueForCheck * bestBaseN * dtypeByteL0b_ * TWO <= static_cast<uint64_t>(platformInfo_.l1_size)) {162+ if (kValueForCheck * bestBaseN * dtypeByteL0b_ * TWO <= static_cast<uint64_t>(platformInfo_.l1_size)) {
163- return false; // 能走进fullload tiling模板的用例163+ return false; // 能走进fullload tiling模板的用例
164- }
165 }164 }
166 }165 }
167- if (runInfo_.dedx_h > runInfo_.dedx_w || runInfo_.dedx_cin > runInfo_.dedy_cout) {166+ }
167+ if (runInfo_.dedx_h > runInfo_.dedx_w || runInfo_.dedx_cin > runInfo_.dedy_cout) {
168 return false; // MTE2 bound性能恶化 经验判断公式168 return false; // MTE2 bound性能恶化 经验判断公式
169- }169+ }
170 return true;170 return true;
171}171}
172 172 
@@ -182,14 +182,14 @@ bool Conv3DDXV2KernelSplitTiling::CheckBestBlockEnable(uint64_t nValue, uint64_t
182bool Conv3DDXV2KernelSplitTiling::CheckShapeConditions()182bool Conv3DDXV2KernelSplitTiling::CheckShapeConditions()
183{183{
184 if (!IsSocVersionFuse(context_) && (runInfo_.filterFormat == ge::FORMAT_NDHWC && // CV耦合架构,kernel拆分省scalar,性能有收益184 if (!IsSocVersionFuse(context_) && (runInfo_.filterFormat == ge::FORMAT_NDHWC && // CV耦合架构,kernel拆分省scalar,性能有收益
185- (kSCoutFullLoad_ || runInfo_.dedx_cin == 1))) { // cin较小,则转为NDHWC性能较差185+ (kSCoutFullLoad_ || runInfo_.dedx_cin == 1))) { // cin较小,则转为NDHWC性能较差
186 return false;186 return false;
187 }187 }
188 188 
189 if (runInfo_.kernel_h == 1 && runInfo_.kernel_w == 1) {189 if (runInfo_.kernel_h == 1 && runInfo_.kernel_w == 1) {
190 if (!CheckKernelSplitHW11Enable()) {190 if (!CheckKernelSplitHW11Enable()) {
191 return false;191 return false;
192- } 192+ }
193 kSCoutFullLoad_ = false;193 kSCoutFullLoad_ = false;
194 }194 }
195 // 12 经验值,wi较小时kernel拆分性能较差195 // 12 经验值,wi较小时kernel拆分性能较差
@@ -293,8 +293,25 @@ bool Conv3DDXV2KernelSplitTiling::CheckKernelSplitHEnable(const uint32_t bestBas
293 return true;293 return true;
294}294}
295 295 
296-// kernel拆分判断296+bool Conv3DDXV2KernelSplitTiling::CheckDtypeCompatibility()
297-bool Conv3DDXV2KernelSplitTiling::CheckKernelSplitEnable()297+{
298+ if (!IsSocVersionFuse(context_)) {
299+ return true;
300+ }
301+ 
302+ size_t filterIndex = FILTER_INDEX;
303+ size_t outputBackpropIndex = OUTPUT_BP_INDEX;
304+ if (opType_ == optiling::OpTypeV2::kConv3DTransposeV2) {
305+ outputBackpropIndex = FILTER_INDEX;
306+ filterIndex = OUTPUT_BP_INDEX;
307+ }
308+ 
309+ ge::DataType filterDtype = context_->GetInputDesc(filterIndex)->GetDataType();
310+ ge::DataType outputBackpropDtype = context_->GetInputDesc(outputBackpropIndex)->GetDataType();
311+ return !(outputBackpropDtype == ge::DT_FLOAT16 && filterDtype == ge::DT_INT8);
312+}
313+ 
314+bool Conv3DDXV2KernelSplitTiling::CheckBasicConstraints()
298{315{
299 if (runInfo_.groups > 1 || runInfo_.dilation_h != 1 || runInfo_.dilation_w != 1) {316 if (runInfo_.groups > 1 || runInfo_.dilation_h != 1 || runInfo_.dilation_w != 1) {
300 return false;317 return false;
@@ -306,35 +323,74 @@ bool Conv3DDXV2KernelSplitTiling::CheckKernelSplitEnable()
306 }323 }
307 324 
308 bool isPadAlign = (runInfo_.pad_u == runInfo_.pad_d && runInfo_.pad_l == runInfo_.pad_r);325 bool isPadAlign = (runInfo_.pad_u == runInfo_.pad_d && runInfo_.pad_l == runInfo_.pad_r);
309- if (!isPadAlign) {326+ return isPadAlign;
310- return false;327+}
311- }
312 328 
329+bool Conv3DDXV2KernelSplitTiling::CheckShapeValue() const
330+{
313 constexpr uint32_t bestBaseMN = 256; // kernel拆分M N最优基本块331 constexpr uint32_t bestBaseMN = 256; // kernel拆分M N最优基本块
314 uint64_t mValue = static_cast<uint64_t>(runInfo_.dedx_w) * runInfo_.dedx_h;332 uint64_t mValue = static_cast<uint64_t>(runInfo_.dedx_w) * runInfo_.dedx_h;
315- if (mValue < bestBaseMN && runInfo_.dedy_cout1_g == 1) {333+ return !(mValue < bestBaseMN && runInfo_.dedy_cout1_g == 1); // 当输入shape较少时,拆分后还多了几次数据搬运,性能可能没有正收益
316- return false; // 当输入shape较少时,拆分后还多了几次数据搬运,性能可能没有正收益334+}
317- }
318 335 
336+bool Conv3DDXV2KernelSplitTiling::CheckKernelSizeStrideMatch(
337+ int32_t kernelSplitStrideVal, uint32_t splitKernelSize1, uint32_t splitKernelSize2, uint32_t splitKernelSize3,
338+ uint32_t splitKernelSize4)
339+{
340+ return (runInfo_.kernel_w == splitKernelSize1 && runInfo_.kernel_h == splitKernelSize1) ||
341+ (runInfo_.kernel_w == splitKernelSize2 && runInfo_.kernel_h == splitKernelSize2) ||
342+ (runInfo_.kernel_w == splitKernelSize3 && runInfo_.kernel_h == splitKernelSize3) ||
343+ (runInfo_.kernel_w == splitKernelSize4 && runInfo_.kernel_h == splitKernelSize4);
344+}
345+ 
346+bool Conv3DDXV2KernelSplitTiling::TryKernelSplitHW(uint32_t bestBaseMN)
347+{
319 constexpr int32_t kernelSplitStrideVal = 2; // 2:当前仅stride等于2的kernel拆分348 constexpr int32_t kernelSplitStrideVal = 2; // 2:当前仅stride等于2的kernel拆分
320 constexpr uint32_t splitKernelSize1 = 1; // 1:当前支持kernel等于1, stride等于2的kernel拆分349 constexpr uint32_t splitKernelSize1 = 1; // 1:当前支持kernel等于1, stride等于2的kernel拆分
321 constexpr uint32_t splitKernelSize2 = 2; // 2:当前支持kernel等于2, stride等于2的kernel拆分350 constexpr uint32_t splitKernelSize2 = 2; // 2:当前支持kernel等于2, stride等于2的kernel拆分
322 constexpr uint32_t splitKernelSize3 = 3; // 3:当前支持kernel等于3, stride等于2的kernel拆分351 constexpr uint32_t splitKernelSize3 = 3; // 3:当前支持kernel等于3, stride等于2的kernel拆分
323 constexpr uint32_t splitKernelSize4 = 4; // 4:当前支持kernel等于4, stride等于2的kernel拆分352 constexpr uint32_t splitKernelSize4 = 4; // 4:当前支持kernel等于4, stride等于2的kernel拆分
324- bool isEnableKernelSplitFlag1 = (runInfo_.kernel_w == splitKernelSize1 && runInfo_.kernel_h == splitKernelSize1);353+ 
325- bool isEnableKernelSplitFlag2 = (runInfo_.kernel_w == splitKernelSize2 && runInfo_.kernel_h == splitKernelSize2);354+ bool strideMatch = (runInfo_.stride_w == kernelSplitStrideVal && runInfo_.stride_h == kernelSplitStrideVal);
326- bool isEnableKernelSplitFlag3 = (runInfo_.kernel_w == splitKernelSize3 && runInfo_.kernel_h == splitKernelSize3);355+ if (!strideMatch) {
327- bool isEnableKernelSplitFlag4 = (runInfo_.kernel_w == splitKernelSize4 && runInfo_.kernel_h == splitKernelSize4);356+ return false;
328- if (runInfo_.stride_w == kernelSplitStrideVal && runInfo_.stride_h == kernelSplitStrideVal &&
329- (isEnableKernelSplitFlag1 || isEnableKernelSplitFlag2 || isEnableKernelSplitFlag3 || isEnableKernelSplitFlag4)) {
330- SetParamForKernelSplit(false);
331- return CheckKernelSplitHWEnable(isEnableKernelSplitFlag2, kernelSplitStrideVal, bestBaseMN);
332- } else if (runInfo_.stride_h >= kernelSplitStrideVal && runInfo_.kernel_h >= kernelSplitStrideVal) {
333- SetParamForKernelSplit();
334- return CheckKernelSplitHEnable(bestBaseMN);
335 }357 }
336 358 
337- return false;359+ bool kernelSizeMatch = CheckKernelSizeStrideMatch(
360+ kernelSplitStrideVal, splitKernelSize1, splitKernelSize2, splitKernelSize3, splitKernelSize4);
361+ if (!kernelSizeMatch) {
362+ return false;
363+ }
364+ 
365+ bool isEnableKernelSplitFlag2 = (runInfo_.kernel_w == splitKernelSize2 && runInfo_.kernel_h == splitKernelSize2);
366+ SetParamForKernelSplit(false);
367+ return CheckKernelSplitHWEnable(isEnableKernelSplitFlag2, kernelSplitStrideVal, bestBaseMN);
368+}
369+ 
370+bool Conv3DDXV2KernelSplitTiling::TryKernelSplitH(uint32_t bestBaseMN)
371+{
372+ constexpr int32_t kernelSplitStrideVal = 2;
373+ if (runInfo_.stride_h < kernelSplitStrideVal || runInfo_.kernel_h < kernelSplitStrideVal) {
374+ return false;
375+ }
376+ 
377+ SetParamForKernelSplit();
378+ return CheckKernelSplitHEnable(bestBaseMN);
379+}
380+ 
381+bool Conv3DDXV2KernelSplitTiling::CheckKernelSplitEnable()
382+{
383+ if (!CheckDtypeCompatibility() || !CheckBasicConstraints() || !CheckShapeValue()) {
384+ return false;
385+ }
386+ 
387+ constexpr uint32_t bestBaseMN = 256;
388+ 
389+ if (TryKernelSplitHW(bestBaseMN)) {
390+ return true;
391+ }
392+ 
393+ return TryKernelSplitH(bestBaseMN);
338}394}
339 395 
340void Conv3DDXV2KernelSplitTiling::UpdateWorkSpaceSize(L0TilingParams& l0Params)396void Conv3DDXV2KernelSplitTiling::UpdateWorkSpaceSize(L0TilingParams& l0Params)
@@ -72,6 +72,14 @@ protected:
72 bool ShrinkBaseMN(L1TilingParams& l1Params, L0TilingParams& l0Params) override;72 bool ShrinkBaseMN(L1TilingParams& l1Params, L0TilingParams& l0Params) override;
73 73 
74 bool CheckKernelSplitEnable();74 bool CheckKernelSplitEnable();
75+ bool CheckDtypeCompatibility();
76+ bool CheckBasicConstraints();
77+ bool CheckShapeValue() const;
78+ bool CheckKernelSizeStrideMatch(
79+ int32_t kernelSplitStrideVal, uint32_t splitKernelSize1, uint32_t splitKernelSize2, uint32_t splitKernelSize3,
80+ uint32_t splitKernelSize4);
81+ bool TryKernelSplitHW(uint32_t bestBaseMN);
82+ bool TryKernelSplitH(uint32_t bestBaseMN);
75 void UpdateWorkSpaceSize(L0TilingParams& l0Params);83 void UpdateWorkSpaceSize(L0TilingParams& l0Params);
76 84 
77 // for kernel split param85 // for kernel split param
@@ -97,7 +105,7 @@ private:
97 void UpdateBaseKParams(L1TilingParams& l1Params, L0TilingParams& l0Params, uint32_t coutA1, uint32_t coutB1);105 void UpdateBaseKParams(L1TilingParams& l1Params, L0TilingParams& l0Params, uint32_t coutA1, uint32_t coutB1);
98 void ShrinkBaseKForKernelSplit(106 void ShrinkBaseKForKernelSplit(
99 L1TilingParams& l1Params, L0TilingParams& l0Params, uint32_t coutA1, uint32_t coutB1);107 L1TilingParams& l1Params, L0TilingParams& l0Params, uint32_t coutA1, uint32_t coutB1);
100- 108+ 
101 bool CheckBestBlockEnable(uint64_t nValue, uint64_t bestBlockCnt);109 bool CheckBestBlockEnable(uint64_t nValue, uint64_t bestBlockCnt);
102 bool CheckShapeConditions();110 bool CheckShapeConditions();
103};111};
@@ -52,6 +52,8 @@ __aicore__ inline constexpr Convolution3DBackprop::CubeFormat GetFormat(int form
52 return Convolution3DBackprop::CubeFormat::NDHWC;52 return Convolution3DBackprop::CubeFormat::NDHWC;
53 } else if (format == FORMAT_DHWCN || format == FORMAT_HWCN) {53 } else if (format == FORMAT_DHWCN || format == FORMAT_HWCN) {
54 return Convolution3DBackprop::CubeFormat::DHWCN;54 return Convolution3DBackprop::CubeFormat::DHWCN;
55+ } else if (format == FORMAT_FRACTAL_Z) {
56+ return Convolution3DBackprop::CubeFormat::FRACTALZ;
55 } else {57 } else {
56 return Convolution3DBackprop::CubeFormat::NCDHW;58 return Convolution3DBackprop::CubeFormat::NCDHW;
57 }59 }
@@ -280,6 +282,8 @@ protected:
280 cinStrideB_ = static_cast<uint64_t>(tiling_->dk) * tiling_->hk * tiling_->wk;282 cinStrideB_ = static_cast<uint64_t>(tiling_->dk) * tiling_->hk * tiling_->wk;
281 } else if constexpr (filterCubeFormat == Convolution3DBackprop::CubeFormat::NDHWC) {283 } else if constexpr (filterCubeFormat == Convolution3DBackprop::CubeFormat::NDHWC) {
282 cinStrideB_ = 1;284 cinStrideB_ = 1;
285+ } else if constexpr (filterCubeFormat == Convolution3DBackprop::CubeFormat::FRACTALZ) {
286+ cinStrideB_ = static_cast<uint64_t>(tiling_->c0);
283 } else { // DHWCN287 } else { // DHWCN
284 cinStrideB_ = static_cast<uint64_t>(tiling_->cout);288 cinStrideB_ = static_cast<uint64_t>(tiling_->cout);
285 }289 }
@@ -308,6 +312,8 @@ protected:
308 groupStrideB_ = static_cast<uint64_t>(tiling_->coutG) * tiling_->dk * tiling_->hk * tiling_->wk *312 groupStrideB_ = static_cast<uint64_t>(tiling_->coutG) * tiling_->dk * tiling_->hk * tiling_->wk *
309 tiling_->cinG;313 tiling_->cinG;
310 }314 }
315+ } else if constexpr (filterCubeFormat == Convolution3DBackprop::CubeFormat::FRACTALZ) {
316+ groupStrideB_ = static_cast<uint64_t>(tiling_->coutG) * tiling_->hk * tiling_->wk * tiling_->cinG;
311 } else { // DHWCN317 } else { // DHWCN
312 groupStrideB_ = static_cast<uint64_t>(tiling_->coutG);318 groupStrideB_ = static_cast<uint64_t>(tiling_->coutG);
313 }319 }
@@ -380,10 +386,10 @@ protected:
380 }386 }
381 387 
382#if (__NPU_ARCH__ == 5102)388#if (__NPU_ARCH__ == 5102)
383- __aicore__ inline void CalcBiasOffset()389+ __aicore__ inline void CalcBiasOffset(uint32_t groupIdx)
384 {390 {
385 if constexpr (biasFormat != FORMAT_MAX) {391 if constexpr (biasFormat != FORMAT_MAX) {
386- offsetBias_ = static_cast<uint64_t>(nCoreIdx_) * tiling_->singleCoreCin;392+ offsetBias_ = static_cast<uint64_t>(nCoreIdx_) * tiling_->singleCoreCin + groupIdx * tiling_->cinG;
387 }393 }
388 }394 }
389#endif395#endif
@@ -391,7 +397,11 @@ protected:
391 __aicore__ inline void CalcScaleOffset()397 __aicore__ inline void CalcScaleOffset()
392 {398 {
393 if constexpr (GetScaleFormat<filterType>(scaleFormat) != Convolution3DBackprop::CubeFormat::UNSUPPORT) {399 if constexpr (GetScaleFormat<filterType>(scaleFormat) != Convolution3DBackprop::CubeFormat::UNSUPPORT) {
394- offsetScale_ = static_cast<uint64_t>(nCoreIdx_) * tiling_->singleCoreCin;400+ if (tiling_->quantMode == static_cast<uint8_t>(Convolution3DBackprop::QuantMode::VECTOR_QUANT)) {
401+ offsetScale_ = static_cast<uint64_t>(nCoreIdx_) * tiling_->singleCoreCin;
402+ } else if (tiling_->quantMode == static_cast<uint8_t>(Convolution3DBackprop::QuantMode::SCALAR_QUANT)) {
403+ offsetScale_ = 0;
404+ }
395 }405 }
396 }406 }
397 __aicore__ inline void CalcBlockOffset(uint32_t batchIdx, uint32_t groupIdx)407 __aicore__ inline void CalcBlockOffset(uint32_t batchIdx, uint32_t groupIdx)
@@ -401,7 +411,7 @@ protected:
401 CalcBlockOffsetC(batchIdx);411 CalcBlockOffsetC(batchIdx);
402 CalcGroupBlockOffset(groupIdx);412 CalcGroupBlockOffset(groupIdx);
403#if (__NPU_ARCH__ == 5102)413#if (__NPU_ARCH__ == 5102)
404- CalcBiasOffset();414+ CalcBiasOffset(groupIdx);
405#endif415#endif
406 CalcScaleOffset();416 CalcScaleOffset();
407 }417 }
@@ -31,6 +31,7 @@ enum class CubeFormat : uint8_t {
31 FRACTALZ_3D,31 FRACTALZ_3D,
32 ND,32 ND,
33 UNSUPPORT,33 UNSUPPORT,
34+ FRACTALZ,
34};35};
35 36 
36enum class QuantMode : std::uint8_t {37enum class QuantMode : std::uint8_t {
@@ -336,6 +336,24 @@ __aicore__ inline void LoadGmDataToB1ForNd2Nz(Intf *self, uint32_t curCinSize, u
336 }336 }
337}337}
338 338 
339+template <class Intf>
340+__aicore__ inline void LoadGmDataToB1ForFz(Intf *self, uint32_t curCinSize, uint32_t curCoutSize,
341+ uint64_t out2B1SrcAddrOffset, const LocalTensor<typename Intf::SrcBT> &useB1Buf)
342+{
343+ DataCopyPadExtParams<typename Intf::SrcBT> padParams;
344+ DataCopyExtParams dataCopyParams;
345+ if (self->ctx.tiling_->cinG == curCinSize) {
346+ dataCopyParams.blockCount = 1;
347+ dataCopyParams.blockLen = AlignUp16(curCinSize) * AlignUp(curCoutSize, self->ctx.tiling_->c0) * self->ctx.hkWk_ * sizeof(typename Intf::SrcBT);
348+ dataCopyParams.srcStride = 0;
349+ } else {
350+ dataCopyParams.blockCount = DivCeil(AlignUp(curCoutSize, self->ctx.tiling_->c0) * self->ctx.hkWk_, self->ctx.tiling_->c0);
351+ dataCopyParams.blockLen = static_cast<uint64_t>(AlignUp16(curCinSize)) * self->ctx.tiling_->c0 * sizeof(typename Intf::SrcBT);
352+ dataCopyParams.srcStride = (self->ctx.tiling_->cinG - AlignUp16(curCinSize)) * self->ctx.tiling_->c0 * sizeof(typename Intf::SrcBT);
353+ }
354+ DataCopyPad<typename Intf::SrcBT>(useB1Buf, self->ctx.weightGlobal_[out2B1SrcAddrOffset], dataCopyParams, padParams);
355+}
356+ 
339template <class Intf>357template <class Intf>
340__aicore__ inline void LoadGmDataToB1DHWCN2NzTranspose(Intf *self, uint32_t curCinSize, uint32_t curCoutSize,358__aicore__ inline void LoadGmDataToB1DHWCN2NzTranspose(Intf *self, uint32_t curCinSize, uint32_t curCoutSize,
341 uint64_t out2B1SrcAddrOffset, const LocalTensor<typename Intf::SrcBT> &useB1Buf)359 uint64_t out2B1SrcAddrOffset, const LocalTensor<typename Intf::SrcBT> &useB1Buf)
@@ -395,6 +413,11 @@ __aicore__ inline void LoadGmDataToB1(Intf *self, uint32_t kIdx, uint32_t curDkI
395 (curDkIdx * self->ctx.hkWk_ + self->ctx.curHkIdx_ * self->ctx.tiling_->wk + self->ctx.curWkIdx_) *413 (curDkIdx * self->ctx.hkWk_ + self->ctx.curHkIdx_ * self->ctx.tiling_->wk + self->ctx.curWkIdx_) *
396 self->ctx.tiling_->cinG * self->ctx.tiling_->cout;414 self->ctx.tiling_->cinG * self->ctx.tiling_->cout;
397 LoadGmDataToB1DHWCN2NzTranspose(self, curCinSize, curCoutSize, out2B1SrcAddrOffset, useB1Buf);415 LoadGmDataToB1DHWCN2NzTranspose(self, curCinSize, curCoutSize, out2B1SrcAddrOffset, useB1Buf);
416+ } else if constexpr (Intf::Config::xType::format == Convolution3DBackprop::CubeFormat::FRACTALZ) {
417+ uint64_t out2B1SrcAddrOffset = static_cast<uint64_t>(curCoutIdx) * self->ctx.hkWk_ * AlignUp16(self->ctx.tiling_->cinG) +
418+ (self->ctx.curHkIdx_ * self->ctx.tiling_->wk + self->ctx.curWkIdx_) * self->ctx.tiling_->cinG * self->ctx.tiling_->c0 +
419+ curCinIdx * self->ctx.tiling_->c0;
420+ LoadGmDataToB1ForFz(self, curCinSize, curCoutSize, out2B1SrcAddrOffset, useB1Buf);
398 } else { // NDHWC421 } else { // NDHWC
399 uint64_t out2B1SrcAddrOffset = static_cast<uint64_t>(curCoutIdx) * self->ctx.tiling_->cinG * self->ctx.dkHkWk_ +422 uint64_t out2B1SrcAddrOffset = static_cast<uint64_t>(curCoutIdx) * self->ctx.tiling_->cinG * self->ctx.dkHkWk_ +
400 curCinIdx + (curDkIdx * self->ctx.hkWk_ + self->ctx.curHkIdx_ * self->ctx.tiling_->wk + self->ctx.curWkIdx_) *423 curCinIdx + (curDkIdx * self->ctx.hkWk_ + self->ctx.curHkIdx_ * self->ctx.tiling_->wk + self->ctx.curWkIdx_) *
@@ -99,6 +99,13 @@ ASCENDC_TPL_SEL(
99 ASCENDC_TPL_BOOL_SEL(isBasicBlockTiling, 1),99 ASCENDC_TPL_BOOL_SEL(isBasicBlockTiling, 1),
100 ASCENDC_TPL_UINT_SEL(loadB1Condition, ASCENDC_TPL_UI_LIST, TPL_VEC_TO_L1_C04)100 ASCENDC_TPL_UINT_SEL(loadB1Condition, ASCENDC_TPL_UI_LIST, TPL_VEC_TO_L1_C04)
101 ),101 ),
102+ ASCENDC_TPL_ARGS_SEL(
103+ ASCENDC_TPL_UINT_SEL(loadB2Condition, ASCENDC_TPL_UI_LIST, TPL_NO_TRANSPOSE_NO_REVERSE),
104+ ASCENDC_TPL_UINT_SEL(kernelSplitMode, ASCENDC_TPL_UI_LIST, TPL_NO_SPLIT_KERNEL),
105+ ASCENDC_TPL_UINT_SEL(groupConvMode, ASCENDC_TPL_UI_LIST, TPL_GROUP_MODE_ORIGIN),
106+ ASCENDC_TPL_BOOL_SEL(isBasicBlockTiling, 1),
107+ ASCENDC_TPL_UINT_SEL(loadB1Condition, ASCENDC_TPL_UI_LIST, TPL_GM_TO_L1)
108+ ),
102 ASCENDC_TPL_ARGS_SEL(109 ASCENDC_TPL_ARGS_SEL(
103 ASCENDC_TPL_UINT_SEL(loadB2Condition, ASCENDC_TPL_UI_LIST, TPL_TRANSPOSE_AND_REVERSE),110 ASCENDC_TPL_UINT_SEL(loadB2Condition, ASCENDC_TPL_UI_LIST, TPL_TRANSPOSE_AND_REVERSE),
104 ASCENDC_TPL_UINT_SEL(kernelSplitMode, ASCENDC_TPL_UI_LIST, TPL_NO_SPLIT_KERNEL),111 ASCENDC_TPL_UINT_SEL(kernelSplitMode, ASCENDC_TPL_UI_LIST, TPL_NO_SPLIT_KERNEL),
@@ -23,51 +23,51 @@ public:
23 this->Input("input_size")23 this->Input("input_size")
24 .ParamType(REQUIRED)24 .ParamType(REQUIRED)
25 .DataType({ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,25 .DataType({ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
26- ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT32, ge::DT_INT32})26+ ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32})
27 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,27 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
28- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})28+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
29 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,29 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
30- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});30+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
31 this->Input("x")31 this->Input("x")
32 .ParamType(REQUIRED)32 .ParamType(REQUIRED)
33 .DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT16, ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT16,33 .DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT16, ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT16,
34- ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT16, ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16})34+ ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT16, ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16})
35 .Format({ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW,35 .Format({ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW,
36- ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW})36+ ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW})
37 .UnknownShapeFormat({ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW,37 .UnknownShapeFormat({ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW,
38- ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW});38+ ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW});
39 this->Input("filter")39 this->Input("filter")
40 .ParamType(REQUIRED)40 .ParamType(REQUIRED)
41 .DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT16, ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT16,41 .DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT16, ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT16,
42- ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT16, ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT16, ge::DT_INT8, ge::DT_INT8})42+ ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT16, ge::DT_INT8, ge::DT_INT8, ge::DT_FLOAT16, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8})
43 .Format({ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NDHWC, ge::FORMAT_NDHWC, ge::FORMAT_NDHWC,43 .Format({ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NDHWC, ge::FORMAT_NDHWC, ge::FORMAT_NDHWC,
44- ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NDHWC, ge::FORMAT_NDHWC, ge::FORMAT_NDHWC, ge::FORMAT_NCDHW, ge::FORMAT_NDHWC})44+ ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NDHWC, ge::FORMAT_NDHWC, ge::FORMAT_NDHWC, ge::FORMAT_NCDHW, ge::FORMAT_NDHWC, ge::FORMAT_FRACTAL_Z})
45 .UnknownShapeFormat({ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NDHWC, ge::FORMAT_NDHWC, ge::FORMAT_NDHWC,45 .UnknownShapeFormat({ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NDHWC, ge::FORMAT_NDHWC, ge::FORMAT_NDHWC,
46- ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NDHWC, ge::FORMAT_NDHWC, ge::FORMAT_NDHWC, ge::FORMAT_NCDHW, ge::FORMAT_NDHWC});46+ ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NDHWC, ge::FORMAT_NDHWC, ge::FORMAT_NDHWC, ge::FORMAT_NCDHW, ge::FORMAT_NDHWC, ge::FORMAT_FRACTAL_Z});
47 this->Input("bias")47 this->Input("bias")
48 .ParamType(OPTIONAL)48 .ParamType(OPTIONAL)
49 .DataType({ge::DT_INT32, ge::DT_INT32, ge::DT_FLOAT16, ge::DT_INT32, ge::DT_INT32, ge::DT_FLOAT16,49 .DataType({ge::DT_INT32, ge::DT_INT32, ge::DT_FLOAT16, ge::DT_INT32, ge::DT_INT32, ge::DT_FLOAT16,
50- ge::DT_INT32, ge::DT_INT32, ge::DT_FLOAT16, ge::DT_INT32, ge::DT_INT32, ge::DT_FLOAT16, ge::DT_INT32, ge::DT_INT32})50+ ge::DT_INT32, ge::DT_INT32, ge::DT_FLOAT16, ge::DT_INT32, ge::DT_INT32, ge::DT_FLOAT16, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32})
51 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,51 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
52- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})52+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
53 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,53 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
54- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});54+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
55 this->Input("scale")55 this->Input("scale")
56 .ParamType(OPTIONAL)56 .ParamType(OPTIONAL)
57 .DataType({ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64,57 .DataType({ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64,
58- ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64})58+ ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64, ge::DT_UINT64})
59 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,59 .Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
60- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})60+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
61 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,61 .UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
62- ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});62+ ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
63 this->Output("y")63 this->Output("y")
64 .ParamType(REQUIRED)64 .ParamType(REQUIRED)
65 .DataType({ge::DT_INT8, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_INT8, ge::DT_FLOAT16, ge::DT_FLOAT16,65 .DataType({ge::DT_INT8, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_INT8, ge::DT_FLOAT16, ge::DT_FLOAT16,
66- ge::DT_INT8, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_INT8, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16})66+ ge::DT_INT8, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_INT8, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16})
67 .Format({ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW,67 .Format({ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW,
68- ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW})68+ ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW})
69 .UnknownShapeFormat({ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW,69 .UnknownShapeFormat({ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW,
70- ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW});70+ ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW, ge::FORMAT_NCDHW});
71 71 
72 this->Attr("strides").AttrType(REQUIRED).ListInt();72 this->Attr("strides").AttrType(REQUIRED).ListInt();
73 this->Attr("pads").AttrType(REQUIRED).ListInt();73 this->Attr("pads").AttrType(REQUIRED).ListInt();
@@ -867,8 +867,8 @@ ExtendConvTransposeTilingTestParam cases_params_fuse[] = {
867 true,867 true,
868 true,868 true,
869 8,869 8,
870- 16777474,870+ 16777218,
871- "1 1 1 1 1 1 8 0 2 2 2 1 1 1 32 4 5 1 0 0 33554433 1 1 256 1 256 8 1 16 1 1 36 64 1 72 128 1 2 2 1 1 1 2 2 0 0 0 0 0 0 0 1 1 1 1 1 1 1 1 256 16 1 288 32 16 16 16 1 0 1 0 320 0 0 0 1 0 "},871+ "1 1 1 1 1 1 8 0 1 1 2 2 2 1 32 4 5 1 0 0 33554433 1 1 256 1 256 8 1 16 1 1 36 64 1 72 128 1 2 2 1 1 1 2 2 0 0 0 0 0 0 0 1 1 1 1 1 1 1 1 256 16 1 384 64 16 16 16 1 0 1 0 384 0 0 0 0 0 "},
872 872 
873 {"net_ndhwc_a16w8_2_fp16_pertensor_quant_mode",873 {"net_ndhwc_a16w8_2_fp16_pertensor_quant_mode",
874 "SOC_L1_1024",874 "SOC_L1_1024",
@@ -1019,8 +1019,8 @@ ExtendConvTransposeTilingTestParam cases_params_fuse[] = {
1019 true,1019 true,
1020 true,1020 true,
1021 8,1021 8,
1022- 16777474,1022+ 16777218,
1023- "1 1 1 1 1 1 8 0 1 1 2 2 2 1 32 4 5 1 0 0 33554433 1 64 128 64 128 4 4 8 4 1 288 112 1 576 224 1 2 2 1 1 1 2 2 0 0 0 0 0 0 0 1 1 1 1 1 1 1 1 128 64 1 448 64 64 2 2 1 0 1 0 448 0 0 0 0 0 "},1023+ "1 1 1 1 1 1 8 0 1 1 1 2 1 1 32 4 5 1 0 0 33619969 1 64 128 64 128 4 4 8 4 1 288 112 1 576 224 1 2 2 1 1 1 2 2 0 0 0 0 0 0 0 1 1 1 1 1 1 1 1 128 64 1 1024 32 64 16 32 1 0 1 0 16128 0 0 0 0 0 "},
1024 1024 
1025 {"net_ndhwc_a16w16_large_input_stride_4",1025 {"net_ndhwc_a16w16_large_input_stride_4",
1026 "SOC_L1_1024",1026 "SOC_L1_1024",
@@ -1095,8 +1095,8 @@ ExtendConvTransposeTilingTestParam cases_params_fuse[] = {
1095 true,1095 true,
1096 true,1096 true,
1097 8,1097 8,
1098- 16777730,1098+ 16777218,
1099- "1 1 1 1 1 1 8 0 1 1 2 2 2 1 32 4 5 1 0 0 33554433 1 64 128 64 128 4 4 8 4 1 112 56 1 448 224 1 4 4 1 1 1 4 4 0 0 0 0 0 0 0 3 3 3 3 1 1 1 1 128 64 1 512 64 64 8 8 1 0 1 0 512 0 0 0 0 0 "},1099+ "1 1 1 1 1 1 8 0 1 1 1 2 1 1 32 4 5 1 0 0 33619969 1 64 128 64 128 4 4 8 4 1 112 56 1 448 224 1 4 4 1 1 1 4 4 0 0 0 0 0 0 0 3 3 3 3 1 1 1 1 128 64 1 1024 32 64 32 128 1 0 1 0 12544 0 0 0 0 0 "},
1100 1100 
1101 {"net_ndhwc_a16w16_multi_batch",1101 {"net_ndhwc_a16w16_multi_batch",
1102 "SOC_L1_1024",1102 "SOC_L1_1024",
@@ -1171,8 +1171,8 @@ ExtendConvTransposeTilingTestParam cases_params_fuse[] = {
1171 true,1171 true,
1172 true,1172 true,
1173 8,1173 8,
1174- 16777474,1174+ 16777218,
1175- "1 1 1 1 1 1 8 0 1 1 2 2 2 1 32 4 5 1 0 0 33554433 4 64 256 64 256 8 4 16 4 1 40 32 1 80 64 1 2 2 1 1 1 2 2 0 0 0 0 0 0 0 1 1 1 1 1 1 1 1 256 64 1 512 64 64 4 4 1 0 1 0 512 0 0 0 0 1 "},1175+ "1 1 1 1 1 1 8 0 1 1 1 2 1 1 32 4 5 1 0 0 33619969 4 64 256 64 256 8 4 16 4 1 40 32 1 80 64 1 2 2 1 1 1 2 2 0 0 0 0 0 0 0 1 1 1 1 1 1 1 1 256 64 1 1024 32 64 16 64 1 0 1 0 2560 0 0 0 0 0 "},
1176};1176};
1177 1177 
1178INSTANTIATE_TEST_CASE_P(1178INSTANTIATE_TEST_CASE_P(