* Copyright (c) Huawei Technologies Co., Ltd. 2026-2026. All rights reserved.
* MindIE is licensed under Mulan PSL v2.
* You can use this software according to the terms and conditions of the Mulan PSL v2.
* You may obtain a copy of Mulan PSL v2 at:
* http://license.coscl.org.cn/MulanPSL2
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND,
* EITHER EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT,
* MERCHANTABILITY OR FIT FOR A PARTICULAR PURPOSE.
* See the Mulan PSL v2 for more details.
*/
* \file split_core.cpp
* \brief
*/
#include "split_core.h"
#include <cstdio>
#include <math.h>
#include "log/log.h"
namespace optiling {
uint32_t GetS1SeqSize(uint32_t bIdx, const BaseInfo &baseInfo) {
if (baseInfo.actualSeqS1Size.empty()) {
return baseInfo.s1Size;
}
if (baseInfo.actualLenQDims == 1U) {
return static_cast<uint32_t>(baseInfo.actualSeqS1Size[0]);
}
if (!baseInfo.isAccumSeqS1) {
return static_cast<uint32_t>(baseInfo.actualSeqS1Size[bIdx]);
}
return (bIdx == 0) ? static_cast<uint32_t>(baseInfo.actualSeqS1Size[bIdx])
: static_cast<uint32_t>(baseInfo.actualSeqS1Size[bIdx] - baseInfo.actualSeqS1Size[bIdx - 1U]);
}
uint32_t GetS2SeqSize(uint32_t bIdx, const BaseInfo &baseInfo) {
uint32_t prefix = static_cast<uint32_t>(baseInfo.actualSeqPrefixSize);
if (baseInfo.actualSeqS2Size.empty()) {
return prefix + baseInfo.s2Size;
}
if (baseInfo.actualLenKvDims == 1U) {
return prefix + static_cast<uint32_t>(baseInfo.actualSeqS2Size[0]);
}
if (!baseInfo.isAccumSeqS2) {
return prefix + static_cast<uint32_t>(baseInfo.actualSeqS2Size[bIdx]);
}
return (bIdx == 0)
? prefix + static_cast<uint32_t>(baseInfo.actualSeqS2Size[bIdx])
: prefix + static_cast<uint32_t>(baseInfo.actualSeqS2Size[bIdx] - baseInfo.actualSeqS2Size[bIdx - 1U]);
}
int64_t CalcPreTokenLeftUp(uint32_t s1Size, uint32_t s2Size, const BaseInfo &baseInfo) {
auto mode = static_cast<SparseMode>(baseInfo.sparseMode);
if (mode == SparseMode::BAND) {
return static_cast<int64_t>(s1Size) - static_cast<int64_t>(s2Size) + baseInfo.preToken;
}
return baseInfo.preToken;
}
int64_t CalcNextTokenLeftUp(uint32_t s1Size, uint32_t s2Size, const BaseInfo &baseInfo) {
auto mode = static_cast<SparseMode>(baseInfo.sparseMode);
switch (mode) {
case SparseMode::DEFAULT_MASK:
case SparseMode::ALL_MASK:
case SparseMode::LEFT_UP_CAUSAL:
return baseInfo.nextToken;
case SparseMode::RIGHT_DOWN_CAUSAL:
return static_cast<int64_t>(s2Size) - static_cast<int64_t>(s1Size);
case SparseMode::BAND:
return static_cast<int64_t>(s2Size) - static_cast<int64_t>(s1Size) + baseInfo.nextToken;
case SparseMode::TREE:
return static_cast<int64_t>(s2Size) - static_cast<int64_t>(s1Size);
default:
return baseInfo.nextToken;
}
}
int64_t CalcCost(uint32_t basicM, uint32_t basicS2) {
uint32_t alignCoefM = 16U;
uint32_t alignCoefS2 = 64U;
uint32_t alignBasicM = (basicM + alignCoefM - 1U) >> 4U;
uint32_t alignBasicS2 = (basicS2 + alignCoefS2 - 1U) >> 6U;
return static_cast<int64_t>(6U * alignBasicM + 10U * alignBasicS2);
}
BlockCost<int64_t> CalcCostTable(
uint32_t s1NormalSize, uint32_t s2NormalSize, uint32_t s1GTailSize, uint32_t s2TailSize) {
BlockCost<int64_t> typeCost{};
typeCost[NORMAL_BLOCK][NORMAL_BLOCK] = CalcCost(s1NormalSize, s2NormalSize);
typeCost[TAIL_BLOCK][NORMAL_BLOCK] = (s1GTailSize == 0U) ? 0U : CalcCost(s1GTailSize, s2NormalSize);
typeCost[NORMAL_BLOCK][TAIL_BLOCK] = (s2TailSize == 0U) ? 0U : CalcCost(s1NormalSize, s2TailSize);
typeCost[TAIL_BLOCK][TAIL_BLOCK] = (s1GTailSize == 0U || s2TailSize == 0U) ? 0U : CalcCost(s1GTailSize, s2TailSize);
return typeCost;
}
Range<uint32_t> CalcS2Range(
uint32_t s1GIdx, const BaseInfo &baseInfo, const SplitParam &splitParam, const BatchCache &batchCache) {
uint32_t s2Start = 0U;
uint32_t s2End = 0U;
if (batchCache.s1Size == 0U || batchCache.s2Size == 0U) {
return std::make_pair(s2Start, s2End);
}
if (!baseInfo.attenMaskFlag) {
s2Start = 0U;
s2End = (batchCache.s2Size + splitParam.s2BaseSize - 1U) / splitParam.s2BaseSize;
return std::make_pair(s2Start, s2End);
}
int64_t s1GFirstToken = static_cast<int64_t>(s1GIdx) * static_cast<int64_t>(splitParam.mBaseSize);
int64_t s1GLastToken = std::min(s1GFirstToken + static_cast<int64_t>(splitParam.mBaseSize),
static_cast<int64_t>(batchCache.s1Size) * static_cast<int64_t>(baseInfo.gSize)) -
1;
int64_t s1FirstToken = 0;
int64_t s1LastToken = 0;
if (baseInfo.isS1G) {
s1FirstToken = s1GFirstToken / static_cast<int64_t>(baseInfo.gSize);
s1LastToken = s1GLastToken / static_cast<int64_t>(baseInfo.gSize);
} else {
if (s1GFirstToken / batchCache.s1Size == s1GLastToken / batchCache.s1Size) {
s1FirstToken = s1GFirstToken % static_cast<int64_t>(batchCache.s1Size);
s1LastToken = s1GLastToken % static_cast<int64_t>(batchCache.s1Size);
} else {
s1FirstToken = 0;
s1LastToken = batchCache.s1Size;
}
}
int64_t s2FirstToken = s1FirstToken - batchCache.preTokenLeftUp;
int64_t s2LastToken = s1LastToken + batchCache.nextTokenLeftUp;
if (s2FirstToken >= batchCache.s2Size || s2LastToken < 0 || s2LastToken < s2FirstToken) {
s2Start = 0U;
s2End = 0U;
return std::make_pair(s2Start, s2End);
}
s2FirstToken = Clip(s2FirstToken, static_cast<int64_t>(0), static_cast<int64_t>(batchCache.s2Size - 1U));
s2LastToken = Clip(s2LastToken, static_cast<int64_t>(0), static_cast<int64_t>(batchCache.s2Size - 1U));
s2Start = static_cast<uint32_t>(s2FirstToken) / splitParam.s2BaseSize;
s2End = static_cast<uint32_t>(s2LastToken) / splitParam.s2BaseSize +
1U;
return std::make_pair(s2Start, s2End);
}
void CalcSplitInfo(SplitContext &splitContext) {
const BaseInfo &baseInfo = splitContext.baseInfo;
const SplitParam &splitParam = splitContext.splitParam;
SplitInfo &splitInfo = splitContext.splitInfo;
for (uint32_t bIdx = 0; bIdx < baseInfo.bSize; bIdx++) {
uint32_t s1Size = GetS1SeqSize(bIdx, baseInfo);
uint32_t s2Size = GetS2SeqSize(bIdx, baseInfo);
splitInfo.s1GBaseNum[bIdx] = (s1Size * baseInfo.gSize + (splitParam.mBaseSize - 1U)) / splitParam.mBaseSize;
splitInfo.s1GTailSize[bIdx] = (s1Size * baseInfo.gSize) % splitParam.mBaseSize;
splitInfo.s2BaseNum[bIdx] = (s2Size + splitParam.s2BaseSize - 1U) / splitParam.s2BaseSize;
splitInfo.s2TailSize[bIdx] = s2Size % splitParam.s2BaseSize;
if (splitInfo.s1GBaseNum[bIdx] != 0U && splitInfo.s2BaseNum[bIdx] != 0U) {
splitInfo.isKvSeqAllZero = false;
}
}
OP_LOGI("SplitInfo", "========== SplitInfo ==========");
OP_LOGI("SplitInfo", "splitInfo.isKvSeqAllZero: %u", splitInfo.isKvSeqAllZero);
for (uint32_t i = 0; i < baseInfo.bSize; i++) {
OP_LOGI("SplitInfo", "splitInfo.s1GBaseNum[%u]: %d", i, splitInfo.s1GBaseNum[i]);
OP_LOGI("SplitInfo", "splitInfo.s1GTailSize[%u]: %d", i, splitInfo.s1GTailSize[i]);
OP_LOGI("SplitInfo", "splitInfo.s2BaseNum[%u]: %d", i, splitInfo.s2BaseNum[i]);
OP_LOGI("SplitInfo", "splitInfo.s2TailSize[%u]: %d", i, splitInfo.s2TailSize[i]);
}
}
void CalcBatchCache(uint32_t bIdx, const SplitContext &splitContext, BatchCache &batchCache) {
const BaseInfo &baseInfo = splitContext.baseInfo;
const SplitParam &splitParam = splitContext.splitParam;
const SplitInfo &splitInfo = splitContext.splitInfo;
batchCache.bIdx = bIdx;
batchCache.s1Size = GetS1SeqSize(bIdx, baseInfo);
batchCache.s2Size = GetS2SeqSize(bIdx, baseInfo);
batchCache.preTokenLeftUp = CalcPreTokenLeftUp(batchCache.s1Size, batchCache.s2Size, baseInfo);
batchCache.nextTokenLeftUp = CalcNextTokenLeftUp(batchCache.s1Size, batchCache.s2Size, baseInfo);
batchCache.typeCost = CalcCostTable(
splitParam.mBaseSize, splitParam.s2BaseSize, splitInfo.s1GTailSize[bIdx], splitInfo.s2TailSize[bIdx]);
}
void CalcS1GCache(uint32_t s1GIdx, const SplitContext &splitContext, const BatchCache &batchCache, S1GCache &s1GCache) {
const BaseInfo &baseInfo = splitContext.baseInfo;
const SplitParam &splitParam = splitContext.splitParam;
const SplitInfo &splitInfo = splitContext.splitInfo;
s1GCache.bIdx = batchCache.bIdx;
s1GCache.s1GIdx = s1GIdx;
auto s2Range = CalcS2Range(s1GIdx, baseInfo, splitParam, batchCache);
s1GCache.s2Start = s2Range.first;
s1GCache.s2End = s2Range.second;
if (s1GCache.s2Start >= s1GCache.s2End) {
s1GCache.s1GBlock = 0;
s1GCache.s1GCost = 0;
s1GCache.s1GLastBlockCost = 0;
s1GCache.s1GNormalBlockCost = 0;
return;
}
s1GCache.s1GBlock = s1GCache.s2End - s1GCache.s2Start;
uint32_t curTailS2Num =
(splitInfo.s2TailSize[batchCache.bIdx] != 0U && s1GCache.s2End == splitInfo.s2BaseNum[batchCache.bIdx]) ? 1U
: 0U;
uint32_t curNormalS2Num = s1GCache.s1GBlock - curTailS2Num;
if (splitInfo.s1GBaseNum[batchCache.bIdx] == 0) {
s1GCache.s1GCost = 0;
s1GCache.s1GLastBlockCost = 0;
s1GCache.s1GNormalBlockCost = 0;
} else if (s1GIdx == (splitInfo.s1GBaseNum[batchCache.bIdx] - 1U) && splitInfo.s1GTailSize[batchCache.bIdx] != 0U) {
s1GCache.s1GCost = batchCache.typeCost[TAIL_BLOCK][NORMAL_BLOCK] * curNormalS2Num +
batchCache.typeCost[TAIL_BLOCK][TAIL_BLOCK] * curTailS2Num;
s1GCache.s1GLastBlockCost = curTailS2Num > 0U ? batchCache.typeCost[TAIL_BLOCK][TAIL_BLOCK]
: batchCache.typeCost[TAIL_BLOCK][NORMAL_BLOCK];
s1GCache.s1GNormalBlockCost = batchCache.typeCost[TAIL_BLOCK][NORMAL_BLOCK];
} else {
s1GCache.s1GCost = batchCache.typeCost[NORMAL_BLOCK][NORMAL_BLOCK] * curNormalS2Num +
batchCache.typeCost[NORMAL_BLOCK][TAIL_BLOCK] * curTailS2Num;
s1GCache.s1GLastBlockCost = curTailS2Num > 0U ? batchCache.typeCost[NORMAL_BLOCK][TAIL_BLOCK]
: batchCache.typeCost[NORMAL_BLOCK][NORMAL_BLOCK];
s1GCache.s1GNormalBlockCost = batchCache.typeCost[NORMAL_BLOCK][NORMAL_BLOCK];
}
}
void CopyTmpResult(SplitResult &tmpRes, SplitResult &splitRes) {
uint64_t len = tmpRes.bN2End.size();
splitRes.usedCoreNum = tmpRes.usedCoreNum;
splitRes.maxCost = tmpRes.maxCost;
splitRes.numOfFdHead = tmpRes.numOfFdHead;
splitRes.maxS2SplitNum = tmpRes.maxS2SplitNum;
for (size_t i = 0; i < len; ++i) {
splitRes.bN2End[i] = tmpRes.bN2End[i];
splitRes.gS1End[i] = tmpRes.gS1End[i];
splitRes.s2End[i] = tmpRes.s2End[i];
splitRes.fdRes.bN2IdxOfFdHead[i] = tmpRes.fdRes.bN2IdxOfFdHead[i];
splitRes.fdRes.gS1IdxOfFdHead[i] = tmpRes.fdRes.gS1IdxOfFdHead[i];
splitRes.fdRes.s2SplitNumOfFdHead[i] = tmpRes.fdRes.s2SplitNumOfFdHead[i];
splitRes.fdRes.s2SplitStartIdxOfCore[i] = tmpRes.fdRes.s2SplitStartIdxOfCore[i];
splitRes.fdRes.gS1SplitNumOfFdHead[i] = tmpRes.fdRes.gS1SplitNumOfFdHead[i];
splitRes.fdRes.gS1LastPartSizeOfFdHead[i] = tmpRes.fdRes.gS1LastPartSizeOfFdHead[i];
}
}
void ClearTmpResult(SplitResult &tmpResult) {
uint64_t len = tmpResult.bN2End.size();
tmpResult.usedCoreNum = 0U;
tmpResult.maxCost = 0;
tmpResult.numOfFdHead = 0U;
tmpResult.maxS2SplitNum = 0U;
tmpResult.usedVecNumOfFd = 0U;
for (size_t i = 0; i < len; ++i) {
tmpResult.bN2End[i] = 0U;
tmpResult.gS1End[i] = 0U;
tmpResult.s2End[i] = 0U;
tmpResult.fdRes.bN2IdxOfFdHead[i] = 0U;
tmpResult.fdRes.gS1IdxOfFdHead[i] = 0U;
tmpResult.fdRes.s2SplitNumOfFdHead[i] = 0U;
tmpResult.fdRes.s2SplitStartIdxOfCore[i] = 0U;
tmpResult.fdRes.gS1SplitNumOfFdHead[i] = 0U;
tmpResult.fdRes.gS1LastPartSizeOfFdHead[i] = 0U;
}
}
void CalcBatchCost(uint32_t bIdx, const SplitContext &splitContext, CostInfo &costInfo) {
const BaseInfo &baseInfo = splitContext.baseInfo;
const SplitInfo &splitInfo = splitContext.splitInfo;
costInfo.bN2CostOfEachBatch[bIdx] = 0;
costInfo.bN2BlockOfEachBatch[bIdx] = 0U;
costInfo.bN2LastBlockCostOfEachBatch[bIdx] = 0U;
if (GetS1SeqSize(bIdx, baseInfo) == 0U || GetS2SeqSize(bIdx, baseInfo) == 0U) {
return;
}
BatchCache bCache;
S1GCache s1GCache;
CalcBatchCache(bIdx, splitContext, bCache);
for (uint32_t s1GIdx = 0; s1GIdx < splitInfo.s1GBaseNum[bIdx]; s1GIdx++) {
CalcS1GCache(s1GIdx, splitContext, bCache, s1GCache);
costInfo.bN2CostOfEachBatch[bIdx] += s1GCache.s1GCost;
costInfo.bN2BlockOfEachBatch[bIdx] += s1GCache.s1GBlock;
if (s1GCache.s1GBlock > 0) {
costInfo.bN2LastBlockCostOfEachBatch[bIdx] = s1GCache.s1GLastBlockCost;
}
}
}
void CalcCostInfo(SplitContext &splitContext) {
const BaseInfo &baseInfo = splitContext.baseInfo;
const SplitInfo &splitInfo = splitContext.splitInfo;
CostInfo &costInfo = splitContext.costInfo;
if (splitInfo.isKvSeqAllZero) {
costInfo.totalCost = 0;
costInfo.totalBlockNum = 0U;
return;
}
for (uint32_t bIdx = 0; bIdx < baseInfo.bSize; bIdx++) {
CalcBatchCost(bIdx, splitContext, costInfo);
costInfo.totalCost += costInfo.bN2CostOfEachBatch[bIdx] * baseInfo.n2Size;
costInfo.totalBlockNum += costInfo.bN2BlockOfEachBatch[bIdx] * baseInfo.n2Size;
}
OP_LOGI("CostInfo", "========== CostInfo ==========");
OP_LOGI("CostInfo", "costInfo.totalCost: %ld", costInfo.totalCost);
OP_LOGI("CostInfo", "costInfo.totalBlockNum: %u", costInfo.totalBlockNum);
for (uint32_t i = 0; i < baseInfo.bSize; i++) {
OP_LOGI("CostInfo", "costInfo.bN2CostOfEachBatch[%u]: %d", i, costInfo.bN2CostOfEachBatch[i]);
OP_LOGI("CostInfo", "costInfo.bN2LastBlockCostOfEachBatch[%u]: %d", i, costInfo.bN2LastBlockCostOfEachBatch[i]);
OP_LOGI("CostInfo", "costInfo.bN2BlockOfEachBatch[%u]: %d", i, costInfo.bN2BlockOfEachBatch[i]);
}
}
void UpdateCursor(const SplitContext &splitContext, AssignContext &assignContext) {
const BaseInfo &baseInfo = splitContext.baseInfo;
const SplitInfo &splitInfo = splitContext.splitInfo;
const CostInfo &costInfo = splitContext.costInfo;
bool UpdateS1G = false;
bool UpdateBatch = false;
if (assignContext.curS2Idx >= assignContext.s1GCache.s2End) {
assignContext.curS2Idx = 0U;
assignContext.curS1GIdx++;
UpdateS1G = true;
}
if (assignContext.curS1GIdx >= splitInfo.s1GBaseNum[assignContext.curBIdx]) {
assignContext.curS1GIdx = 0U;
assignContext.curBN2Idx++;
}
if (assignContext.curBN2Idx ==
baseInfo.bSize * baseInfo.n2Size) {
assignContext.curS1GIdx = 0U;
assignContext.curS2Idx = 0U;
assignContext.isFinished = true;
return;
}
if (assignContext.curBN2Idx / baseInfo.n2Size != assignContext.curBIdx) {
assignContext.curBIdx = assignContext.curBN2Idx / baseInfo.n2Size;
assignContext.curS1GIdx = 0U;
UpdateBatch = true;
UpdateS1G = true;
}
if (UpdateBatch) {
CalcBatchCache(assignContext.curBIdx, splitContext, assignContext.batchCache);
assignContext.bN2Cost = costInfo.bN2CostOfEachBatch[assignContext.curBIdx];
assignContext.bN2Block = costInfo.bN2BlockOfEachBatch[assignContext.curBIdx];
}
if (UpdateS1G) {
CalcS1GCache(assignContext.curS1GIdx, splitContext, assignContext.batchCache, assignContext.s1GCache);
assignContext.curS2Idx = assignContext.s1GCache.s2Start;
}
}
void AssignByBatch(const SplitContext &splitContext, AssignContext &assignContext) {
if (assignContext.isFinished) {
return;
}
const BaseInfo &baseInfo = splitContext.baseInfo;
const CostInfo &costInfo = splitContext.costInfo;
while (assignContext.bN2Cost == 0 ||
IsWithinTolerance(assignContext.coreCache.costLimit,
costInfo.bN2LastBlockCostOfEachBatch[assignContext.curBIdx] / FA_TOLERANCE_RATIO,
assignContext.coreCache.cost + assignContext.bN2Cost)) {
assignContext.coreCache.cost += assignContext.bN2Cost;
assignContext.coreCache.block += assignContext.bN2Block;
assignContext.curBN2Idx++;
if (assignContext.curBN2Idx == baseInfo.bSize * baseInfo.n2Size) {
assignContext.curS1GIdx = 0U;
assignContext.curS2Idx = 0U;
assignContext.isFinished = true;
return;
}
if (assignContext.curBN2Idx / baseInfo.n2Size != assignContext.curBIdx) {
assignContext.curBIdx = assignContext.curBN2Idx / baseInfo.n2Size;
CalcBatchCache(assignContext.curBIdx, splitContext, assignContext.batchCache);
}
assignContext.bN2Cost = costInfo.bN2CostOfEachBatch[assignContext.curBIdx];
assignContext.bN2Block = costInfo.bN2BlockOfEachBatch[assignContext.curBIdx];
assignContext.curS1GIdx = 0U;
CalcS1GCache(assignContext.curS1GIdx, splitContext, assignContext.batchCache, assignContext.s1GCache);
assignContext.curS2Idx = assignContext.s1GCache.s2Start;
}
}
void AssignByRow(const SplitContext &splitContext, AssignContext &assignContext) {
if (assignContext.isFinished) {
return;
}
while (IsWithinTolerance(assignContext.coreCache.costLimit,
assignContext.s1GCache.s1GLastBlockCost / FA_TOLERANCE_RATIO,
assignContext.coreCache.cost + assignContext.s1GCache.s1GCost)) {
assignContext.coreCache.cost += assignContext.s1GCache.s1GCost;
assignContext.coreCache.block += assignContext.s1GCache.s1GBlock;
assignContext.bN2Cost = assignContext.bN2Cost > assignContext.s1GCache.s1GCost
? assignContext.bN2Cost - assignContext.s1GCache.s1GCost
: 0;
assignContext.bN2Block = assignContext.bN2Block > assignContext.s1GCache.s1GBlock
? assignContext.bN2Block - assignContext.s1GCache.s1GBlock
: 0U;
do {
assignContext.curS1GIdx++;
CalcS1GCache(assignContext.curS1GIdx, splitContext, assignContext.batchCache, assignContext.s1GCache);
} while (assignContext.s1GCache.s1GBlock == 0);
assignContext.curS2Idx = assignContext.s1GCache.s2Start;
}
}
void AssignByBlock(AssignContext &assignContext) {
if (assignContext.isFinished) {
return;
}
int64_t curCost = assignContext.s1GCache.s1GNormalBlockCost;
if (assignContext.curS2Idx == (assignContext.s1GCache.s2End - 1U)) {
curCost = assignContext.s1GCache.s1GLastBlockCost;
}
while (IsWithinTolerance(assignContext.coreCache.costLimit, curCost / FA_TOLERANCE_RATIO,
assignContext.coreCache.cost +
curCost)) {
assignContext.coreCache.cost += curCost;
assignContext.coreCache.block++;
assignContext.curS2Idx++;
assignContext.bN2Cost = assignContext.bN2Cost - curCost;
assignContext.s1GCache.s1GCost = assignContext.s1GCache.s1GCost - curCost;
assignContext.bN2Block--;
assignContext.s1GCache.s1GBlock--;
}
}
void ForceAssign(const SplitContext &splitContext, AssignContext &assignContext) {
if (assignContext.isFinished) {
return;
}
int64_t curCost = assignContext.s1GCache.s1GNormalBlockCost;
if (assignContext.curS2Idx == (assignContext.s1GCache.s2End - 1U)) {
curCost = assignContext.s1GCache.s1GLastBlockCost;
}
assignContext.coreCache.cost += curCost;
assignContext.coreCache.block++;
assignContext.curS2Idx++;
assignContext.bN2Cost = assignContext.bN2Cost - curCost;
assignContext.bN2Block--;
assignContext.s1GCache.s1GCost = assignContext.s1GCache.s1GCost - curCost;
assignContext.s1GCache.s1GBlock--;
UpdateCursor(splitContext, assignContext);
}
bool IsNeedRecordFDInfo(const AssignContext &assignContext, const SplitResult &splitRes) {
if (assignContext.curCoreIdx == 0U) {
return false;
}
if (assignContext.curKvSplitPart <= 1U) {
return false;
}
if (assignContext.curBN2Idx == splitRes.bN2End[assignContext.curCoreIdx - 1U] &&
assignContext.curS1GIdx == splitRes.gS1End[assignContext.curCoreIdx - 1U]) {
return false;
}
return true;
}
void RecordFDInfo(const SplitContext &splitContext, const AssignContext &assignContext, SplitResult &result) {
const BaseInfo &baseInfo = splitContext.baseInfo;
const SplitParam &splitParam = splitContext.splitParam;
const SplitInfo &splitInfo = splitContext.splitInfo;
uint32_t splitBIdx = result.bN2End[assignContext.curCoreIdx - 1U] / baseInfo.n2Size;
uint32_t splitS1GIdx = result.gS1End[assignContext.curCoreIdx - 1U];
uint32_t s1Size = GetS1SeqSize(splitBIdx, baseInfo);
uint32_t curFdS1gSize = (splitS1GIdx == splitInfo.s1GBaseNum[splitBIdx] - 1U)
? (s1Size * baseInfo.gSize - splitS1GIdx * splitParam.mBaseSize)
: splitParam.mBaseSize;
uint32_t curFdS1gSplitPart = (curFdS1gSize + splitParam.gS1BaseSizeOfFd - 1U) / splitParam.gS1BaseSizeOfFd;
uint32_t curFdS1gLastPartSize = curFdS1gSize - (splitParam.gS1BaseSizeOfFd * (curFdS1gSplitPart - 1U));
result.maxS2SplitNum = std::max(result.maxS2SplitNum, assignContext.curKvSplitPart);
result.fdRes.bN2IdxOfFdHead[result.numOfFdHead] = result.bN2End[assignContext.curCoreIdx - 1U];
result.fdRes.gS1IdxOfFdHead[result.numOfFdHead] = result.gS1End[assignContext.curCoreIdx - 1U];
result.fdRes.s2SplitNumOfFdHead[result.numOfFdHead] = assignContext.curKvSplitPart;
result.fdRes.gS1SplitNumOfFdHead[result.numOfFdHead] = curFdS1gSplitPart;
result.fdRes.gS1LastPartSizeOfFdHead[result.numOfFdHead] = curFdS1gLastPartSize;
result.numOfFdHead++;
}
void LogAssignContext(const char *phase, const AssignContext &assignContext) {
OP_LOGD("assignContext",
"[%s] curCoreIdx: %u, unassignedCost: %ld, "
"bIdx: %u, bN2Idx: %u, s1GIdx: %u, s2Idx: %u, "
"bN2Cost: %ld, bN2Block: %u, "
"s1GBlock: %u, s2Start: %u, s2End: %u",
phase, assignContext.curCoreIdx, assignContext.unassignedCost, assignContext.curBIdx, assignContext.curBN2Idx,
assignContext.curS1GIdx, assignContext.curS2Idx, assignContext.bN2Cost, assignContext.bN2Block,
assignContext.s1GCache.s1GBlock, assignContext.s1GCache.s2Start, assignContext.s1GCache.s2End);
}
void CalcSplitPlan(uint32_t coreNum, int64_t costLimit, const SplitContext &splitContext, SplitResult &result) {
const CostInfo &costInfo = splitContext.costInfo;
if (coreNum == 0U) {
return;
}
result.maxCost = 0U;
result.usedCoreNum = 0U;
AssignContext assignContext{};
assignContext.curBIdx = 0U;
assignContext.curS1GIdx = 0U;
assignContext.unassignedCost = costInfo.totalCost;
assignContext.bN2Cost = costInfo.bN2CostOfEachBatch[assignContext.curBIdx];
assignContext.bN2Block = costInfo.bN2BlockOfEachBatch[assignContext.curBIdx];
CalcBatchCache(assignContext.curBIdx, splitContext, assignContext.batchCache);
CalcS1GCache(assignContext.curS1GIdx, splitContext, assignContext.batchCache, assignContext.s1GCache);
assignContext.curS2Idx = assignContext.s1GCache.s2Start;
for (uint32_t i = 0; i < coreNum; ++i) {
if (result.maxCost > costLimit) {
return;
}
if (assignContext.isFinished || assignContext.unassignedCost <= 0) {
break;
}
assignContext.curCoreIdx = i;
result.fdRes.s2SplitStartIdxOfCore[assignContext.curCoreIdx] = assignContext.curKvSplitPart - 1U;
assignContext.coreCache = {};
assignContext.coreCache.costLimit = assignContext.unassignedCost / (coreNum - assignContext.curCoreIdx);
LogAssignContext("START", assignContext);
AssignByBatch(splitContext, assignContext);
LogAssignContext("BATCH", assignContext);
AssignByRow(splitContext, assignContext);
LogAssignContext("ROW", assignContext);
AssignByBlock(assignContext);
LogAssignContext("BLOCK", assignContext);
if (assignContext.coreCache.block == 0) {
ForceAssign(splitContext, assignContext);
LogAssignContext("FORCE", assignContext);
}
result.bN2End[i] = assignContext.curBN2Idx;
result.gS1End[i] = assignContext.curS1GIdx;
result.s2End[i] = assignContext.curS2Idx;
result.maxCost = std::max(result.maxCost, assignContext.coreCache.cost);
assignContext.unassignedCost -= assignContext.coreCache.cost;
if (IsNeedRecordFDInfo(assignContext, result)) {
RecordFDInfo(splitContext, assignContext, result);
assignContext.curKvSplitPart = 1U;
}
if (assignContext.curS2Idx > assignContext.s1GCache.s2Start &&
assignContext.curS2Idx <= assignContext.s1GCache.s2End) {
assignContext.curKvSplitPart++;
}
}
result.usedCoreNum = assignContext.curCoreIdx + 1;
}
void SplitFD(SplitResult &result) {
uint32_t totalFDLoad = 0;
uint32_t totalFDHeadSplit = 0;
for (uint32_t i = 0; i < result.numOfFdHead; i++) {
totalFDLoad += result.fdRes.s2SplitNumOfFdHead[i] * result.fdRes.gS1SplitNumOfFdHead[i];
totalFDHeadSplit += result.fdRes.gS1SplitNumOfFdHead[i];
}
uint32_t maxVectorNum = std::min(totalFDHeadSplit, result.usedCoreNum * result.vecCubeRatio);
double loadThrOfVector =
static_cast<double>(totalFDLoad) / static_cast<double>(maxVectorNum);
int64_t loadOfCurVector = 0;
uint32_t curCoreIndex = 0;
uint32_t preTmpFDIndexEndOfFdHead = 0;
uint32_t preTmpFDIndexEndOfFdHeadSplit = 0;
for (uint32_t i = 0; i < result.numOfFdHead; i++) {
uint32_t fDKVSplitNum = result.fdRes.s2SplitNumOfFdHead[i];
for (uint32_t gS1SplitIdx = 0; gS1SplitIdx < result.fdRes.gS1SplitNumOfFdHead[i]; gS1SplitIdx++) {
double remainSpace = loadThrOfVector - static_cast<double>(loadOfCurVector);
if (fDKVSplitNum > remainSpace * FD_TOLERANCE_RATIO) {
result.fdRes.gS1IdxEndOfFdHead[curCoreIndex] = preTmpFDIndexEndOfFdHead;
result.fdRes.gS1IdxEndOfFdHeadSplit[curCoreIndex] = preTmpFDIndexEndOfFdHeadSplit;
curCoreIndex += 1U;
totalFDLoad -= static_cast<uint32_t>(loadOfCurVector);
loadThrOfVector = static_cast<double>(totalFDLoad) / static_cast<double>(maxVectorNum - curCoreIndex);
loadOfCurVector = 0;
}
loadOfCurVector += fDKVSplitNum;
preTmpFDIndexEndOfFdHead = i;
preTmpFDIndexEndOfFdHeadSplit = gS1SplitIdx;
}
}
result.fdRes.gS1IdxEndOfFdHead[curCoreIndex] = preTmpFDIndexEndOfFdHead;
result.fdRes.gS1IdxEndOfFdHeadSplit[curCoreIndex] = preTmpFDIndexEndOfFdHeadSplit;
result.usedVecNumOfFd = curCoreIndex + 1;
}
void LogSplitCoreInput(const BaseInfo &baseInfo, const SplitParam ¶m) {
OP_LOGI("BaseInfo", "========== BaseInfo ==========");
OP_LOGI("BaseInfo", "bSize: %u", baseInfo.bSize);
OP_LOGI("BaseInfo", "n2Size: %u", baseInfo.n2Size);
OP_LOGI("BaseInfo", "gSize: %u", baseInfo.gSize);
OP_LOGI("BaseInfo", "s1Size: %u", baseInfo.s1Size);
OP_LOGI("BaseInfo", "s2Size: %u", baseInfo.s2Size);
OP_LOGI("BaseInfo", "isS1G: %u", baseInfo.isS1G);
OP_LOGI("BaseInfo", "isAccumSeqS1: %u", baseInfo.isAccumSeqS1);
OP_LOGI("BaseInfo", "isAccumSeqS2: %u", baseInfo.isAccumSeqS2);
OP_LOGI("BaseInfo", "actualLenQDims: %u", baseInfo.actualLenQDims);
OP_LOGI("BaseInfo", "actualLenKvDims: %u", baseInfo.actualLenKvDims);
OP_LOGI("BaseInfo", "attenMaskFlag: %u", baseInfo.attenMaskFlag);
OP_LOGI("BaseInfo", "sparseMode: %u", baseInfo.sparseMode);
OP_LOGI("BaseInfo", "preToken: %ld", baseInfo.preToken);
OP_LOGI("BaseInfo", "nextToken: %ld", baseInfo.nextToken);
OP_LOGI("BaseInfo", "actualSeqPrefixSize: %ld", baseInfo.actualSeqPrefixSize);
OP_LOGI("SplitParam", "mBaseSize: %u", param.mBaseSize);
OP_LOGI("SplitParam", "s2BaseSize: %u", param.s2BaseSize);
}
void LogSplitCoreResult(const SplitResult &result) {
OP_LOGI("SplitRes", "splitRes.usedCoreNum: %d", result.usedCoreNum);
OP_LOGI("SplitRes", "splitRes.maxCost: %d", result.maxCost);
OP_LOGI("SplitRes", "splitRes.numOfFdHead: %d", result.numOfFdHead);
OP_LOGI("SplitRes", "splitRes.maxS2SplitNum: %d", result.maxS2SplitNum);
OP_LOGI("SplitRes", "splitRes.usedVecNumOfFd: %d", result.usedVecNumOfFd);
for (uint32_t i = 0; i < result.usedCoreNum; i++) {
OP_LOGI("SplitRes", "outerSplitParams.bN2End[%u]: %d", i, result.bN2End[i]);
OP_LOGI("SplitRes", "outerSplitParams.gS1End[%u]: %d", i, result.gS1End[i]);
OP_LOGI("SplitRes", "outerSplitParams.s2End[%u]: %d", i, result.s2End[i]);
OP_LOGI("fdRes", "fDParams.s2SplitStartIdxOfCore[%u]: %d", i, result.fdRes.s2SplitStartIdxOfCore[i]);
}
for (uint32_t i = 0; i < result.numOfFdHead; i++) {
OP_LOGI("fdRes", "fDParams.bN2IdxOfFdHead[%u]: %d", i, result.fdRes.bN2IdxOfFdHead[i]);
OP_LOGI("fdRes", "fDParams.gS1IdxOfFdHead[%u]: %d", i, result.fdRes.gS1IdxOfFdHead[i]);
OP_LOGI("fdRes", "fDParams.s2SplitNumOfFdHead[%u]: %d", i, result.fdRes.s2SplitNumOfFdHead[i]);
OP_LOGI("fdRes", "fDParams.gS1SplitNumOfFdHead[%u]: %d", i, result.fdRes.gS1SplitNumOfFdHead[i]);
OP_LOGI("fdRes", "fDParams.gS1LastPartSizeOfFdHead[%u]: %d", i, result.fdRes.gS1LastPartSizeOfFdHead[i]);
}
for (uint32_t i = 0; i < result.usedVecNumOfFd; i++) {
OP_LOGI("fdRes", "fDParams.gS1IdxEndOfFdHead[%u]: %d", i, result.fdRes.gS1IdxEndOfFdHead[i]);
OP_LOGI("fdRes", "fDParams.gS1IdxEndOfFdHeadSplit[%u]: %d", i, result.fdRes.gS1IdxEndOfFdHeadSplit[i]);
}
}
void SplitCore(uint32_t coreNum, const BaseInfo &baseInfo, const SplitParam ¶m, SplitResult &result) {
LogSplitCoreInput(baseInfo, param);
SplitContext splitContext(baseInfo, param);
CalcSplitInfo(splitContext);
if (splitContext.splitInfo.isKvSeqAllZero) {
result.usedCoreNum = 1U;
result.bN2End[0] = baseInfo.bSize * baseInfo.n2Size;
result.gS1End[0] = 0U;
result.s2End[0] = 0U;
return;
}
CalcCostInfo(splitContext);
uint32_t maxCore = std::min(coreNum, splitContext.costInfo.totalBlockNum);
uint32_t minCore =
static_cast<uint32_t>(std::sqrt(static_cast<float>(splitContext.costInfo.totalBlockNum) + 0.25f) + 0.5f);
minCore = std::min(minCore, maxCore);
result.maxCost = INT64_MAX;
result.usedCoreNum = 1U;
SplitResult tmpResult{coreNum, result.vecCubeRatio};
for (uint32_t i = minCore; i <= maxCore; ++i) {
OP_LOGD("SplitRes", "CoreNum: %u", i);
CalcSplitPlan(i, result.maxCost, splitContext, tmpResult);
if (tmpResult.maxCost < result.maxCost) {
CopyTmpResult(tmpResult, result);
}
ClearTmpResult(tmpResult);
}
if (result.numOfFdHead > 0U) {
SplitFD(result);
}
result.usedCoreNum = std::max(result.usedCoreNum, 1U);
LogSplitCoreResult(result);
}
}