已合并
perf(gmmaq): fit epilogue UB double buffer for target shape #11198
zhoushaolong创建于 20 天前
perf(gmmaq): fit epilogue UB double buffer for target shape #11198
已合并
共 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. |
| 71 | constexpr uint64_t MX_EPILOGUE_MAX_ELEMENTS_PER_AIV = 128UL * 256UL; | 71 | constexpr uint64_t MX_EPILOGUE_MAX_ELEMENTS_PER_AIV = 128UL * 256UL; |
| 72 | constexpr uint64_t MX_EPILOGUE_AIV_COUNT = 2UL; | 72 | constexpr 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; | ||
| 73 | constexpr uint32_t FP8_E4M3FN_VALUE = 36; | 77 | constexpr uint32_t FP8_E4M3FN_VALUE = 36; |
| 74 | constexpr uint32_t FP8_E5M2_VALUE = 35; | 78 | constexpr uint32_t FP8_E5M2_VALUE = 35; |
| 75 | constexpr uint32_t FP4_E2M1_VALUE = 40; | 79 | constexpr 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 | + | ||
| 498 | uint64_t GroupedMatmulActivationQuantTiling950::GetTilingKey() const | 547 | uint64_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; |