* 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.h
* \brief
*/
#ifndef SPLIT_CORE_H
#define SPLIT_CORE_H
#include <cstdint>
#include <vector>
#include <array>
#include <algorithm>
namespace optiling {
constexpr int64_t FA_TOLERANCE_RATIO = 2;
constexpr uint32_t FD_TOLERANCE_RATIO = 2U;
enum BlockType : uint32_t { NORMAL_BLOCK = 0, TAIL_BLOCK, BLOCK_MAX_TYPE };
enum class SparseMode : uint8_t {
DEFAULT_MASK = 0,
ALL_MASK,
LEFT_UP_CAUSAL,
RIGHT_DOWN_CAUSAL,
BAND,
SPARSE_BUTT,
TREE = 9,
};
template <class T> using Range = std::pair<T, T>;
template <class T>
using BlockCost = std::array<std::array<T, static_cast<size_t>(BLOCK_MAX_TYPE)>, static_cast<size_t>(BLOCK_MAX_TYPE)>;
template <typename T> T Clip(T value, T minValue, T maxValue) {
if (value < minValue) {
return minValue;
}
if (value > maxValue) {
return maxValue;
}
return value;
}
template <typename T> inline bool IsWithinTolerance(T limit, T tolerance, T value) {
return limit + tolerance >= value;
}
struct BaseInfo {
uint32_t bSize{0U};
uint32_t n2Size{0U};
uint32_t gSize{0U};
uint32_t s1Size{0U};
uint32_t s2Size{0U};
bool isS1G{true};
bool isAccumSeqS1{false};
bool isAccumSeqS2{false};
std::vector<int64_t> actualSeqS1Size{};
std::vector<int64_t> actualSeqS2Size{};
uint32_t actualLenQDims{0U};
uint32_t actualLenKvDims{0U};
bool attenMaskFlag{false};
int32_t sparseMode{0U};
int64_t preToken{0};
int64_t nextToken{0};
int64_t actualSeqPrefixSize{0};
};
struct SplitParam {
uint32_t mBaseSize{1U};
uint32_t s2BaseSize{1U};
uint32_t gS1BaseSizeOfFd{8U};
};
struct FlashDecodeResult {
std::vector<uint32_t> bN2IdxOfFdHead{};
std::vector<uint32_t> gS1IdxOfFdHead{};
std::vector<uint32_t> s2SplitNumOfFdHead{};
std::vector<uint32_t> gS1SplitNumOfFdHead{};
std::vector<uint32_t> gS1LastPartSizeOfFdHead{};
std::vector<uint32_t>
gS1IdxEndOfFdHead{};
std::vector<uint32_t>
gS1IdxEndOfFdHeadSplit{};
std::vector<uint32_t> s2SplitStartIdxOfCore{};
FlashDecodeResult(uint32_t coreNum, uint32_t vecCubeRatio)
: bN2IdxOfFdHead(coreNum), gS1IdxOfFdHead(coreNum), s2SplitNumOfFdHead(coreNum), gS1SplitNumOfFdHead(coreNum),
gS1LastPartSizeOfFdHead(coreNum), gS1IdxEndOfFdHead(coreNum * vecCubeRatio),
gS1IdxEndOfFdHeadSplit(coreNum * vecCubeRatio), s2SplitStartIdxOfCore(coreNum) {}
};
struct SplitResult {
uint32_t usedCoreNum{0U};
uint32_t vecCubeRatio{0U};
std::vector<uint32_t> bN2End{};
std::vector<uint32_t> gS1End{};
std::vector<uint32_t> s2End{};
int64_t maxCost{0};
uint32_t numOfFdHead{0U};
uint32_t maxS2SplitNum{0U};
uint32_t usedVecNumOfFd{0U};
FlashDecodeResult fdRes{0U, 0U};
SplitResult(uint32_t coreNum, uint32_t ratio)
: bN2End(coreNum), vecCubeRatio(ratio), gS1End(coreNum), s2End(coreNum), fdRes(coreNum, ratio) {};
};
struct SplitInfo {
std::vector<uint32_t> s1GBaseNum{};
std::vector<uint32_t> s2BaseNum{};
std::vector<uint32_t> s1GTailSize{};
std::vector<uint32_t> s2TailSize{};
bool isKvSeqAllZero{true};
explicit SplitInfo(uint32_t batchSize)
: s1GBaseNum(batchSize), s2BaseNum(batchSize), s1GTailSize(batchSize), s2TailSize(batchSize) {}
};
struct CostInfo {
std::vector<int64_t> bN2CostOfEachBatch{};
std::vector<uint32_t> bN2BlockOfEachBatch{};
std::vector<int64_t> bN2LastBlockCostOfEachBatch{};
uint32_t totalBlockNum{0U};
int64_t totalCost{0};
explicit CostInfo(uint32_t batchSize)
: bN2CostOfEachBatch(batchSize), bN2BlockOfEachBatch(batchSize), bN2LastBlockCostOfEachBatch(batchSize) {}
};
struct SplitContext {
const BaseInfo &baseInfo;
const SplitParam &splitParam;
SplitInfo splitInfo{0U};
CostInfo costInfo{0U};
explicit SplitContext(const BaseInfo &info, const SplitParam ¶m)
: baseInfo(info), splitParam(param), splitInfo(info.bSize), costInfo(info.bSize) {}
};
struct BatchCache {
uint32_t bIdx{0U};
uint32_t s1Size{0U};
uint32_t s2Size{0U};
int64_t preTokenLeftUp{0};
int64_t nextTokenLeftUp{0};
BlockCost<int64_t> typeCost{};
};
struct S1GCache {
uint32_t bIdx{0U};
uint32_t s1GIdx{0U};
uint32_t s2Start{0U};
uint32_t s2End{0U};
int64_t s1GCost{0};
int64_t s1GLastBlockCost{0};
uint32_t s1GBlock{0U};
int64_t s1GNormalBlockCost{0};
};
struct CoreCache {
int64_t costLimit{0};
int64_t cost{0};
uint32_t block{0U};
};
struct AssignContext {
uint32_t curBIdx{0U};
uint32_t curBN2Idx{0U};
uint32_t curS1GIdx{0U};
uint32_t curS2Idx{0U};
uint32_t curCoreIdx{0U};
int64_t unassignedCost{0};
uint32_t usedCoreNum{0U};
uint32_t curKvSplitPart{1U};
int64_t bN2Cost{0};
uint32_t bN2Block{0U};
bool isFinished{false};
BatchCache batchCache{};
S1GCache s1GCache{};
CoreCache coreCache{};
};
uint32_t GetS1SeqSize(uint32_t bIdx, const BaseInfo &baseInfo);
uint32_t GetS2SeqSize(uint32_t bIdx, const BaseInfo &baseInfo);
int64_t CalcPreTokenLeftUp(uint32_t s1Size, uint32_t s2Size, const BaseInfo &baseInfo);
int64_t CalcNextTokenLeftUp(uint32_t s1Size, uint32_t s2Size, const BaseInfo &baseInfo);
Range<uint32_t> CalcS2Range(
uint32_t s1GIdx, const BaseInfo &baseInfo, const SplitParam &splitParam, const BatchCache &batchCache);
int64_t CalcCost(uint32_t basicM, uint32_t basicS2);
BlockCost<int64_t> CalcCostTable(
uint32_t s1NormalSize, uint32_t s2NormalSize, uint32_t s1GTailSize, uint32_t s2TailSize);
void CalcBatchCache(uint32_t bIdx, const SplitContext &splitContext, BatchCache &batchCache);
void CalcS1GCache(uint32_t s1GIdx, const SplitContext &splitContext, const BatchCache &batchCache, S1GCache &s1GCache);
void CopyTmpResult(SplitResult &tmpRes, SplitResult &splitRes);
void ClearTmpResult(SplitResult &tmpResult);
void CalcSplitInfo(SplitContext &splitContext);
void CalcBatchCost(uint32_t bIdx, const SplitContext &splitContext, CostInfo &costInfo);
void CalcCostInfo(SplitContext &splitContext);
void UpdateCursor(const SplitContext &splitContext, AssignContext &assignContext);
void AssignByBatch(const SplitContext &splitContext, AssignContext &assignContext);
void AssignByRow(const SplitContext &splitContext, AssignContext &assignContext);
void AssignByBlock(AssignContext &assignContext);
void ForceAssign(const SplitContext &splitContext, AssignContext &assignContext);
bool IsNeedRecordFDInfo(const AssignContext &assignContext, const SplitResult &splitRes);
void RecordFDInfo(const SplitContext &splitContext, const AssignContext &assignContext, SplitResult &result);
void SplitFD(SplitResult &result);
void CalcSplitPlan(uint32_t coreNum, int64_t costLimit, const SplitContext &splitContext, SplitResult &result);
void SplitCore(uint32_t coreNum, const BaseInfo &baseInfo, const SplitParam ¶m, SplitResult &result);
void LogSplitCoreInput(const BaseInfo &baseInfo, const SplitParam ¶m);
void LogSplitCoreResult(const SplitResult &result);
void LogAssignContext(const char *phase, const AssignContext &assignContext);
}
#endif