已合并
extendconvtranspose support A16W8 #3472
sunlesheng创建于 4月2日
extendconvtranspose support A16W8 #3472
已合并
共 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时只加载wk | 325 | + 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 | + | ||
| 368 | void Conv3DDXV2InnerProductTiling::SetCommonTilingData( | 383 | void 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 Conv | 113 | } // namespace Conv |
| @@ -144,29 +144,29 @@ void Conv3DDXV2KernelSplitTiling::SetParamForKernelSplit(bool isKernelSplitOnlyH | |||
| 144 | bool Conv3DDXV2KernelSplitTiling::CheckKernelSplitHW11Enable() | 144 | bool 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 | |||
| 182 | bool Conv3DDXV2KernelSplitTiling::CheckShapeConditions() | 182 | bool 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 | ||
| 340 | void Conv3DDXV2KernelSplitTiling::UpdateWorkSpaceSize(L0TilingParams& l0Params) | 396 | void 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 param | 85 | // 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 | }; |
Mconv/conv3d_backprop_input_v2/op_kernel/arch35/conv3d_backprop_input_v2/conv3d_backprop_input_v2.h+14-4
| @@ -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 { // DHWCN | 287 | } 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 { // DHWCN | 317 | } 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 | 388 | ||
| 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 | 395 | ||
| @@ -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 | 413 | ||
| 404 | - CalcBiasOffset(); | 414 | + CalcBiasOffset(groupIdx); |
| 405 | 415 | ||
| 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 | ||
| 36 | enum class QuantMode : std::uint8_t { | 37 | enum 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 | + | ||
| 339 | template <class Intf> | 357 | template <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 { // NDHWC | 421 | } 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 | ||
| 1178 | INSTANTIATE_TEST_CASE_P( | 1178 | INSTANTIATE_TEST_CASE_P( |