已合并
perf(gmmaq): fit epilogue UB double buffer for target shape #11198
perf(gmmaq): fit epilogue UB double buffer for target shape #11198
已合并
zhoushaolong创建于 20 天前
2 个文件变更+52-0
@@ -70,6 +70,10 @@ constexpr uint32_t MXFP4_PACK_FACTOR = 2;
70// while its epilogue buffers hold at most 128 * 256 elements.70// while its epilogue buffers hold at most 128 * 256 elements.
71constexpr uint64_t MX_EPILOGUE_MAX_ELEMENTS_PER_AIV = 128UL * 256UL;71constexpr uint64_t MX_EPILOGUE_MAX_ELEMENTS_PER_AIV = 128UL * 256UL;
72constexpr uint64_t MX_EPILOGUE_AIV_COUNT = 2UL;72constexpr uint64_t MX_EPILOGUE_AIV_COUNT = 2UL;
73+constexpr uint32_t EPILOGUE_UB_SINGLE_BUFFER_COUNT = 1U;
74+constexpr uint32_t EPILOGUE_UB_DOUBLE_BUFFER_COUNT = GroupedMatmul::QUEUE_DOUBLE_BUFFER;
75+constexpr uint32_t EPILOGUE_SCALE_INTERMEDIATE_COUNT = 2U;
76+constexpr uint32_t EPILOGUE_BASE_M_SEARCH_STEP = GmmConstant::CUBE_BLOCK;
73constexpr uint32_t FP8_E4M3FN_VALUE = 36;77constexpr uint32_t FP8_E4M3FN_VALUE = 36;
74constexpr uint32_t FP8_E5M2_VALUE = 35;78constexpr uint32_t FP8_E5M2_VALUE = 35;
75constexpr uint32_t FP4_E2M1_VALUE = 40;79constexpr uint32_t FP4_E2M1_VALUE = 40;
@@ -469,6 +473,7 @@ ge::graphStatus GroupedMatmulActivationQuantTiling950::DoLibApiTiling()
469 OP_LOGE(context_->GetNodeName(), "baseN=%lu exceeds the MX epilogue UB capacity.", basicTiling_.baseN),473 OP_LOGE(context_->GetNodeName(), "baseN=%lu exceeds the MX epilogue UB capacity.", basicTiling_.baseN),
470 return ge::GRAPH_FAILED);474 return ge::GRAPH_FAILED);
471 basicTiling_.baseM = std::min(alignedBaseM, maxBaseMByUb);475 basicTiling_.baseM = std::min(alignedBaseM, maxBaseMByUb);
476+ AdjustBasicBlockForEpilogueDoubleBuffer();
472 OP_CHECK_IF(GroupedQmmBasicApiTiling::CalL1Tiling() != ge::GRAPH_SUCCESS,477 OP_CHECK_IF(GroupedQmmBasicApiTiling::CalL1Tiling() != ge::GRAPH_SUCCESS,
473 OP_LOGE(context_->GetNodeName(), "CalL1Tiling failed."), return ge::GRAPH_FAILED);478 OP_LOGE(context_->GetNodeName(), "CalL1Tiling failed."), return ge::GRAPH_FAILED);
474 tilingData_.mmTilingData.m = inputParams_.mSize;479 tilingData_.mmTilingData.m = inputParams_.mSize;
@@ -495,6 +500,50 @@ ge::graphStatus GroupedMatmulActivationQuantTiling950::DoLibApiTiling()
495 return ge::GRAPH_SUCCESS;500 return ge::GRAPH_SUCCESS;
496}501}
497 502 
503+uint64_t GroupedMatmulActivationQuantTiling950::CalcEpilogueUbBytes(uint64_t baseM, uint64_t baseN,
504+ uint32_t bufferCount) const
505+{
506+ const uint32_t effectiveBufferCount = std::max(bufferCount, EPILOGUE_UB_SINGLE_BUFFER_COUNT);
507+ const uint64_t mPerVector = GroupedMatmul::CeilDiv(baseM, GroupedMatmul::QUEUE_DOUBLE_BUFFER);
508+ const uint64_t maxBlockCount = mPerVector * baseN;
509+ const uint64_t maxScaleCount = GroupedMatmul::CeilDiv(maxBlockCount, static_cast<uint64_t>(AscendC::ONE_BLK_SIZE));
510+ const uint64_t afterIn = maxBlockCount * sizeof(float);
511+ const uint64_t scaleBlockBytes = mPerVector * AscendC::ONE_BLK_SIZE * sizeof(int8_t);
512+ const uint64_t singleBufferBytes =
513+ afterIn + maxBlockCount * sizeof(int8_t) + maxScaleCount * sizeof(int8_t) + maxBlockCount * sizeof(uint16_t) +
514+ maxScaleCount * sizeof(uint16_t) * EPILOGUE_SCALE_INTERMEDIATE_COUNT + scaleBlockBytes;
515+ const uint64_t extraBufferBytes = maxBlockCount * sizeof(int8_t) + scaleBlockBytes;
516+ return singleBufferBytes + (effectiveBufferCount - EPILOGUE_UB_SINGLE_BUFFER_COUNT) * extraBufferBytes;
517+}
518+ 
519+bool GroupedMatmulActivationQuantTiling950::CanEnableEpilogueDoubleBuffer(uint64_t baseM, uint64_t baseN) const
520+{
521+ return CalcEpilogueUbBytes(baseM, baseN, EPILOGUE_UB_DOUBLE_BUFFER_COUNT) <= aicoreParams_.ubSize;
522+}
523+ 
524+void GroupedMatmulActivationQuantTiling950::AdjustBasicBlockForEpilogueDoubleBuffer()
525+{
526+ if (basicTiling_.baseM == 0 || basicTiling_.baseN == 0 || inputParams_.groupNum == 0 ||
527+ CanEnableEpilogueDoubleBuffer(basicTiling_.baseM, basicTiling_.baseN)) {
528+ return;
529+ }
530+ 
531+ const uint64_t originalBaseM = basicTiling_.baseM;
532+ const uint64_t averageGroupM = GroupedMatmul::CeilDiv(inputParams_.mSize, inputParams_.groupNum);
533+ const uint64_t originalBlockCount = GroupedMatmul::CeilDiv(averageGroupM, originalBaseM);
534+ for (uint64_t candidateBaseM = originalBaseM; candidateBaseM >= EPILOGUE_BASE_M_SEARCH_STEP;
535+ candidateBaseM -= EPILOGUE_BASE_M_SEARCH_STEP) {
536+ if (!CanEnableEpilogueDoubleBuffer(candidateBaseM, basicTiling_.baseN)) {
537+ continue;
538+ }
539+ const uint64_t candidateBlockCount = GroupedMatmul::CeilDiv(averageGroupM, candidateBaseM);
540+ if (candidateBlockCount <= originalBlockCount) {
541+ basicTiling_.baseM = static_cast<uint32_t>(candidateBaseM);
542+ }
543+ return;
544+ }
545+}
546+ 
498uint64_t GroupedMatmulActivationQuantTiling950::GetTilingKey() const547uint64_t GroupedMatmulActivationQuantTiling950::GetTilingKey() const
499{548{
500 return GET_TPL_TILING_KEY(static_cast<uint64_t>(inputParams_.transB), static_cast<uint64_t>(inputParams_.transA));549 return GET_TPL_TILING_KEY(static_cast<uint64_t>(inputParams_.transB), static_cast<uint64_t>(inputParams_.transA));
@@ -99,6 +99,9 @@ private:
99 bool CheckWeightNzShape(const gert::Shape &wStorageShape) const;99 bool CheckWeightNzShape(const gert::Shape &wStorageShape) const;
100 bool IsFp4(ge::DataType dtype) const;100 bool IsFp4(ge::DataType dtype) const;
101 bool IsFp8(ge::DataType dtype) const;101 bool IsFp8(ge::DataType dtype) const;
102+ uint64_t CalcEpilogueUbBytes(uint64_t baseM, uint64_t baseN, uint32_t bufferCount) const;
103+ bool CanEnableEpilogueDoubleBuffer(uint64_t baseM, uint64_t baseN) const;
104+ void AdjustBasicBlockForEpilogueDoubleBuffer();
102 105 
103 GroupedMatmulActivationQuant::GMMActivationQuantTilingDataParams tilingData_;106 GroupedMatmulActivationQuant::GMMActivationQuantTilingDataParams tilingData_;
104 uint8_t roundMode_ = DEFAULT_ROUND_MODE_RINT;107 uint8_t roundMode_ = DEFAULT_ROUND_MODE_RINT;