已合并
[flash_attn_metadata] 添加防御检查 #9078
[flash_attn_metadata] 添加防御检查 #9078
已合并
FiguraDoge创建于 7月23日
4 个文件变更+160-67
@@ -237,7 +237,7 @@ SectionStreamKImpl::Compute(const DeviceInfo &deviceInfo, const IBaseInfo &baseI
237 if (computeContext.gridInfo.isEmpty) {237 if (computeContext.gridInfo.isEmpty) {
238 result.emplace_back(deviceInfo.aicCoreMaxNum, deviceInfo.aivCoreMaxNum);238 result.emplace_back(deviceInfo.aicCoreMaxNum, deviceInfo.aivCoreMaxNum);
239 result[0].usedCoreNum = 1U;239 result[0].usedCoreNum = 1U;
240- result[0].bN2End[0] = NumToIndex(baseInfo.GetBatchSize() * baseInfo.GetQueryHeadNum());240+ result[0].bN2End[0] = baseInfo.GetBatchSize() * baseInfo.GetQueryHeadNum();
241 result[0].gS1End[0] = 0U;241 result[0].gS1End[0] = 0U;
242 result[0].s2End[0] = 0U;242 result[0].s2End[0] = 0U;
243 return result;243 return result;
@@ -34,8 +34,10 @@ public:
34 const char *layoutOut, const aclTensor *metadata);34 const char *layoutOut, const aclTensor *metadata);
35 35 
36private:36private:
37- static inline bool IsTensorExist(const aclTensor *tensor);37+ static constexpr int64_t NONE_VALUE = -1;
38 38 
39+ static inline bool IsTensorExist(const aclTensor *tensor);
40+ static inline bool IsPA(const char *layout);
39 static inline aclnnStatus CheckSeqLens(bool isCu, int64_t batchSize, const aclTensor *seqLens);41 static inline aclnnStatus CheckSeqLens(bool isCu, int64_t batchSize, const aclTensor *seqLens);
40 42 
41 static inline aclnnStatus CheckBaseAttr(int64_t batchSize, int64_t maxSeqlenQ, int64_t maxSeqlenKv,43 static inline aclnnStatus CheckBaseAttr(int64_t batchSize, int64_t maxSeqlenQ, int64_t maxSeqlenKv,
@@ -85,20 +87,27 @@ inline bool FlashAttnMetadataCheck::IsTensorExist(const aclTensor *tensor)
85 (tensor->GetData() != nullptr);87 (tensor->GetData() != nullptr);
86}88}
87 89 
90+inline bool FlashAttnMetadataCheck::IsPA(const char *layout)
91+{
92+ return (strcmp(layout, "PA_BNBD") == 0 || strcmp(layout, "PA_BBND") == 0 || strcmp(layout, "PA_NZ") == 0);
93+}
94+ 
88inline aclnnStatus95inline aclnnStatus
89FlashAttnMetadataCheck::CheckBaseAttr(int64_t batchSize, int64_t maxSeqlenQ, int64_t maxSeqlenKv,96FlashAttnMetadataCheck::CheckBaseAttr(int64_t batchSize, int64_t maxSeqlenQ, int64_t maxSeqlenKv,
90 int64_t numHeadsQ, int64_t numHeadsKv, int64_t headDim,97 int64_t numHeadsQ, int64_t numHeadsKv, int64_t headDim,
91 const char *layoutQ, const char *layoutKv, const char *layoutOut)98 const char *layoutQ, const char *layoutKv, const char *layoutOut)
92{99{
93- CHECK_COND((batchSize == -1 || batchSize > 0), ACLNN_ERR_RUNTIME_ERROR,100+ int64_t MIN_BATCH = 0;
94- "batchSize must be -1 or greater than 0, but got %ld", batchSize);101+ int64_t MAX_BATCH = 65536;
95- CHECK_COND((maxSeqlenQ == -1 || maxSeqlenQ > 0), ACLNN_ERR_RUNTIME_ERROR,102+ CHECK_COND((batchSize > MIN_BATCH && batchSize < MAX_BATCH), ACLNN_ERR_RUNTIME_ERROR,
96- "maxSeqlenQ must be -1 or greater than 0, but got %ld", maxSeqlenQ);103+ "batchSize must be %ld or between (%ld, %ld), but got %ld", NONE_VALUE, MIN_BATCH, MAX_BATCH, batchSize);
97- CHECK_COND((maxSeqlenKv == -1 || maxSeqlenKv > 0), ACLNN_ERR_RUNTIME_ERROR,104+ CHECK_COND((maxSeqlenQ == NONE_VALUE || maxSeqlenQ >= 0), ACLNN_ERR_RUNTIME_ERROR,
98- "maxSeqlenKv must be -1 or greater than 0, but got %ld", maxSeqlenKv);105+ "maxSeqlenQ must be %ld or greater than or equal to 0, but got %ld", NONE_VALUE, maxSeqlenQ);
106+ CHECK_COND((maxSeqlenKv == NONE_VALUE || maxSeqlenKv >= 0), ACLNN_ERR_RUNTIME_ERROR,
107+ "maxSeqlenKv must be %ld or greater than or equal to 0, but got %ld", NONE_VALUE, maxSeqlenKv);
99 108 
100 CHECK_COND(numHeadsQ > 0, ACLNN_ERR_RUNTIME_ERROR, "numHeadsQ must be greater than 0, but got %ld", numHeadsQ);109 CHECK_COND(numHeadsQ > 0, ACLNN_ERR_RUNTIME_ERROR, "numHeadsQ must be greater than 0, but got %ld", numHeadsQ);
101- CHECK_COND(numHeadsKv > 0, ACLNN_ERR_RUNTIME_ERROR, "numHeadsKv must be greater than 0, but got %ld", maxSeqlenKv);110+ CHECK_COND(numHeadsKv > 0, ACLNN_ERR_RUNTIME_ERROR, "numHeadsKv must be greater than 0, but got %ld", numHeadsKv);
102 111 
103 constexpr int64_t HEAD_DIM_64 = 64;112 constexpr int64_t HEAD_DIM_64 = 64;
104 constexpr int64_t HEAD_DIM_128 = 128;113 constexpr int64_t HEAD_DIM_128 = 128;
@@ -132,8 +141,20 @@ FlashAttnMetadataCheck::CheckMask(int64_t maskMode, int64_t winLeft, int64_t win
132 static const std::unordered_set<int64_t> maskSet = { NO_MASK, CAUSAL_MASK, WINDOW_MASK };141 static const std::unordered_set<int64_t> maskSet = { NO_MASK, CAUSAL_MASK, WINDOW_MASK };
133 CHECK_COND(maskSet.count(maskMode) > 0, ACLNN_ERR_RUNTIME_ERROR,142 CHECK_COND(maskSet.count(maskMode) > 0, ACLNN_ERR_RUNTIME_ERROR,
134 "maskMode only supports %ld, %ld, %ld, but got %ld", NO_MASK, CAUSAL_MASK, WINDOW_MASK, maskMode);143 "maskMode only supports %ld, %ld, %ld, but got %ld", NO_MASK, CAUSAL_MASK, WINDOW_MASK, maskMode);
135- CHECK_COND(winLeft >= -1, ACLNN_ERR_RUNTIME_ERROR, "winLeft must be -1 or at least 0, but got %ld", winLeft);144+ 
136- CHECK_COND(winRight >= -1, ACLNN_ERR_RUNTIME_ERROR, "winRight must be -1 or at least 0, but got %ld", winRight);145+ if (maskMode == NO_MASK || maskMode == CAUSAL_MASK) {
146+ CHECK_COND(winLeft == NONE_VALUE, ACLNN_ERR_RUNTIME_ERROR,
147+ "When maskMode is %ld, winLeft must be %ld, but got %ld", maskMode, NONE_VALUE, winLeft);
148+ CHECK_COND(winRight == NONE_VALUE, ACLNN_ERR_RUNTIME_ERROR,
149+ "When maskMode is %ld, winRight must be %ld, but got %ld", maskMode, NONE_VALUE, winRight);
150+ } else if (maskMode == WINDOW_MASK) {
151+ CHECK_COND(winLeft >= NONE_VALUE, ACLNN_ERR_RUNTIME_ERROR,
152+ "When maskMode is %ld, winLeft must be %ld or greater than or equal to 0, but got %ld",
153+ maskMode, NONE_VALUE, winLeft);
154+ CHECK_COND(winRight >= NONE_VALUE, ACLNN_ERR_RUNTIME_ERROR,
155+ "When maskMode is %ld, winRight must be %ld or greater than or equal to 0, but got %ld",
156+ maskMode, NONE_VALUE, winRight);
157+ }
137 158 
138 return ACLNN_SUCCESS;159 return ACLNN_SUCCESS;
139}160}
@@ -157,7 +178,7 @@ FlashAttnMetadataCheck::CheckExistency(int64_t maxSeqlenQ, int64_t maxSeqlenKv,
157 "When layoutQ is not TND, cuSeqlensQOptional should not be provided, but got non-null");178 "When layoutQ is not TND, cuSeqlensQOptional should not be provided, but got non-null");
158 179 
159 // maxSeqlenQ和sequsedQOptional必须有一个(-1表示不传)180 // maxSeqlenQ和sequsedQOptional必须有一个(-1表示不传)
160- CHECK_COND(((maxSeqlenQ >= 0) || IsTensorExist(sequsedQOptional)), ACLNN_ERR_RUNTIME_ERROR,181+ CHECK_COND(((maxSeqlenQ > 0) || IsTensorExist(sequsedQOptional)), ACLNN_ERR_RUNTIME_ERROR,
161 "When layoutQ is not TND, at least one of maxSeqlenQ or sequsedQOptional must be provided");182 "When layoutQ is not TND, at least one of maxSeqlenQ or sequsedQOptional must be provided");
162 }183 }
163 184 
@@ -165,14 +186,18 @@ FlashAttnMetadataCheck::CheckExistency(int64_t maxSeqlenQ, int64_t maxSeqlenKv,
165 // layoutQ为TND时,必须传入cuSeqlensKvOptional186 // layoutQ为TND时,必须传入cuSeqlensKvOptional
166 CHECK_COND(IsTensorExist(cuSeqlensKvOptional), ACLNN_ERR_RUNTIME_ERROR,187 CHECK_COND(IsTensorExist(cuSeqlensKvOptional), ACLNN_ERR_RUNTIME_ERROR,
167 "When layoutKv is TND, cuSeqlensKvOptional should be provided, but got null");188 "When layoutKv is TND, cuSeqlensKvOptional should be provided, but got null");
189+ } else if (IsPA(layoutKv)) {
190+ CHECK_COND(IsTensorExist(sequsedKvOptional), ACLNN_ERR_RUNTIME_ERROR,
191+ "When layoutKv is PA, sequsedKvOptional must be provided");
168 } else {192 } else {
169 // layoutKv不为TND时,不可以传入cuSeqlensKvOptional193 // layoutKv不为TND时,不可以传入cuSeqlensKvOptional
170 CHECK_COND(!IsTensorExist(cuSeqlensKvOptional), ACLNN_ERR_RUNTIME_ERROR,194 CHECK_COND(!IsTensorExist(cuSeqlensKvOptional), ACLNN_ERR_RUNTIME_ERROR,
171 "When layoutKv is not TND, cuSeqlensKvOptional should not be provided, but got non-null");195 "When layoutKv is not TND, cuSeqlensKvOptional should not be provided, but got non-null");
172 // maxSeqlenKv和sequsedKvOptional必须有一个(-1表示不传)196 // maxSeqlenKv和sequsedKvOptional必须有一个(-1表示不传)
173- CHECK_COND(((maxSeqlenKv >= 0) || IsTensorExist(sequsedKvOptional)), ACLNN_ERR_RUNTIME_ERROR,197+ CHECK_COND(((maxSeqlenKv > 0) || IsTensorExist(sequsedKvOptional)), ACLNN_ERR_RUNTIME_ERROR,
174 "When layoutKv is not TND, at least one of maxSeqlenKv or sequsedKvOptional must be provided");198 "When layoutKv is not TND, at least one of maxSeqlenKv or sequsedKvOptional must be provided");
175 }199 }
200+ 
176 return ACLNN_SUCCESS;201 return ACLNN_SUCCESS;
177}202}
178 203 
@@ -75,6 +75,8 @@ bool FlashAttnMetadataCpuKernel::Prepare(CpuKernelContext &ctx)
75 GetAttrValueOpt(ctx, "layout_q", layoutQ_);75 GetAttrValueOpt(ctx, "layout_q", layoutQ_);
76 GetAttrValueOpt(ctx, "layout_kv", layoutKv_);76 GetAttrValueOpt(ctx, "layout_kv", layoutKv_);
77 GetAttrValueOpt(ctx, "layout_out", layoutOut_);77 GetAttrValueOpt(ctx, "layout_out", layoutOut_);
78+ 
79+ KERNEL_CHECK_FALSE(ParamsCheck(), false, "Params check failed");
78 return ParamsInit();80 return ParamsInit();
79}81}
80 82 
@@ -86,6 +88,110 @@ bool FlashAttnMetadataCpuKernel::ParamsInit()
86 return true;88 return true;
87}89}
88 90 
91+bool FlashAttnMetadataCpuKernel::ParamsCheck()
92+{
93+ KERNEL_CHECK_FALSE(CheckActualQuerySeq(), false, "Check query sequence failed");
94+ KERNEL_CHECK_FALSE(CheckActualKvSeq(), false, "Check kv sequence failed");
95+ return true;
96+}
97+ 
98+bool FlashAttnMetadataCpuKernel::CheckActualQuerySeq()
99+{
100+ isActualSeqlenQAccum_ = false;
101+ actualSeqlenQ_.clear();
102+ std::vector<int64_t> cuSeqlensQ {};
103+ std::vector<int64_t> sequsedQ {};
104+ 
105+ if (cuSeqlensQ_ != nullptr && cuSeqlensQ_->GetData() != nullptr) {
106+ cuSeqlensQ = GetTensorDataAsInt64(cuSeqlensQ_, cuSeqlensQ_->GetTensorShape()->GetDimSize(0));
107+ }
108+ if (sequsedQ_ != nullptr && sequsedQ_->GetData() != nullptr) {
109+ sequsedQ = GetTensorDataAsInt64(sequsedQ_, sequsedQ_->GetTensorShape()->GetDimSize(0));
110+ }
111+ 
112+ for (size_t i = 0; i < sequsedQ.size(); ++i) {
113+ if (sequsedQ[i] < 0) {
114+ KERNEL_LOG_ERROR("The elements of sequsedQ must be non-negative, but %zuth element is %ld", i, sequsedQ[i]);
115+ return false;
116+ }
117+ }
118+ 
119+ if (!cuSeqlensQ.empty()) {
120+ if (cuSeqlensQ[0] != 0) {
121+ KERNEL_LOG_ERROR("The first element of cuSeqlensQ must be 0, but got %ld", cuSeqlensQ[0]);
122+ return false;
123+ }
124+ }
125+ 
126+ for (size_t i = 1; i < cuSeqlensQ.size(); ++i) {
127+ if (cuSeqlensQ[i] < cuSeqlensQ[i - 1]) {
128+ KERNEL_LOG_ERROR(
129+ "The %zuth element of cuSeqlensQ must be greather than the %zuth element, but got %ld and %ld",
130+ cuSeqlensQ[i], cuSeqlensQ[i-1]);
131+ return false;
132+ }
133+ }
134+ 
135+ if (!sequsedQ.empty()) {
136+ isActualSeqlenQAccum_ = false;
137+ actualSeqlenQ_ = sequsedQ;
138+ } else if (!cuSeqlensQ.empty()) {
139+ isActualSeqlenQAccum_ = true;
140+ actualSeqlenQ_ = cuSeqlensQ;
141+ }
142+ 
143+ return true;
144+}
145+ 
146+bool FlashAttnMetadataCpuKernel::CheckActualKvSeq()
147+{
148+ isActualSeqlenKvAccum_ = false;
149+ actualSeqlenKv_.clear();
150+ std::vector<int64_t> cuSeqlensKv {};
151+ std::vector<int64_t> sequsedKv {};
152+ 
153+ if (cuSeqlensKv_ != nullptr && cuSeqlensKv_->GetData() != nullptr) {
154+ cuSeqlensKv = GetTensorDataAsInt64(cuSeqlensKv_, cuSeqlensKv_->GetTensorShape()->GetDimSize(0));
155+ }
156+ if (sequsedKv_ != nullptr && sequsedKv_->GetData() != nullptr) {
157+ sequsedKv = GetTensorDataAsInt64(sequsedKv_, sequsedKv_->GetTensorShape()->GetDimSize(0));
158+ }
159+ 
160+ for (size_t i = 0; i < sequsedKv.size(); ++i) {
161+ if (sequsedKv[i] < 0) {
162+ KERNEL_LOG_ERROR("The elements of sequsedKv must be non-negative, but %zuth element is %ld",
163+ i, sequsedKv[i]);
164+ return false;
165+ }
166+ }
167+ 
168+ if (!cuSeqlensKv.empty()) {
169+ if (cuSeqlensKv[0] != 0) {
170+ KERNEL_LOG_ERROR("The first element of cuSeqlensKv must be 0, but got %ld", cuSeqlensKv[0]);
171+ return false;
172+ }
173+ }
174+ 
175+ for (size_t i = 1; i < cuSeqlensKv.size(); ++i) {
176+ if (cuSeqlensKv[i] < cuSeqlensKv[i - 1]) {
177+ KERNEL_LOG_ERROR(
178+ "The %zuth element of cuSeqlensKv must be greather than the %zuth element, but got %ld and %ld",
179+ cuSeqlensKv[i], cuSeqlensKv[i-1]);
180+ return false;
181+ }
182+ }
183+ 
184+ if (!sequsedKv.empty()) {
185+ isActualSeqlenKvAccum_ = false;
186+ actualSeqlenKv_ = sequsedKv;
187+ } else if (!cuSeqlensKv.empty()) {
188+ isActualSeqlenKvAccum_ = true;
189+ actualSeqlenKv_ = cuSeqlensKv;
190+ }
191+ 
192+ return true;
193+}
194+ 
89void FlashAttnMetadataCpuKernel::InitDeviceInfo()195void FlashAttnMetadataCpuKernel::InitDeviceInfo()
90{196{
91 deviceInfo.aicCoreMaxNum = aicCoreNum_;197 deviceInfo.aicCoreMaxNum = aicCoreNum_;
@@ -113,6 +219,7 @@ void FlashAttnMetadataCpuKernel::InitLoadBalanceParams()
113 param.fdLeastBlock = 3; // 3: least block219 param.fdLeastBlock = 3; // 3: least block
114 param.fdOn = true;220 param.fdOn = true;
115}221}
222+ 
116void FlashAttnMetadataCpuKernel::InitBaseInfo()223void FlashAttnMetadataCpuKernel::InitBaseInfo()
117{224{
118 baseInfo.batchSize = batchSize_;225 baseInfo.batchSize = batchSize_;
@@ -129,60 +236,10 @@ void FlashAttnMetadataCpuKernel::InitBaseInfo()
129 baseInfo.layoutKv = load_balance::ConvertToLayout(layoutKv_);236 baseInfo.layoutKv = load_balance::ConvertToLayout(layoutKv_);
130 baseInfo.queryType = load_balance::DataType::FP16;237 baseInfo.queryType = load_balance::DataType::FP16;
131 baseInfo.kvType = load_balance::DataType::FP16;238 baseInfo.kvType = load_balance::DataType::FP16;
132- LoadActualQuerySeq();239+ baseInfo.isCumulativeKvSeq = isActualSeqlenKvAccum_;
133- LoadActualKvSeq();240+ baseInfo.actualKvSeqSize = actualSeqlenKv_;
134-}241+ baseInfo.isCumulativeQuerySeq = isActualSeqlenQAccum_;
135- 242+ baseInfo.actualQuerySeqSize = actualSeqlenQ_;
136-void FlashAttnMetadataCpuKernel::LoadActualQuerySeq()
137-{
138- baseInfo.actualQuerySeqSize.clear();
139- baseInfo.isCumulativeQuerySeq = (layoutQ_ == "TND" || layoutQ_ == "NTD");
140- 
141- if (sequsedQ_ != nullptr && sequsedQ_->GetData() != nullptr) {
142- batchSize_ = sequsedQ_->GetTensorShape()->GetDimSize(0);
143- baseInfo.batchSize = batchSize_;
144- auto tmpSeq = GetTensorDataAsInt64(sequsedQ_, sequsedQ_->GetTensorShape()->GetDimSize(0));
145- baseInfo.querySeqSize = static_cast<uint32_t>(*std::max_element(tmpSeq.begin(), tmpSeq.end()));
146- baseInfo.actualQuerySeqSize.assign(tmpSeq.begin(), tmpSeq.end());
147- if (baseInfo.isCumulativeQuerySeq) {
148- std::partial_sum(tmpSeq.begin(), tmpSeq.end(), baseInfo.actualQuerySeqSize.begin());
149- }
150- } else if (cuSeqlensQ_ != nullptr && cuSeqlensQ_->GetData() != nullptr) {
151- batchSize_ = cuSeqlensQ_->GetTensorShape()->GetDimSize(0) - 1U;
152- baseInfo.batchSize = batchSize_;
153- auto tmpSeq = GetTensorDataAsInt64(cuSeqlensQ_, cuSeqlensQ_->GetTensorShape()->GetDimSize(0));
154- baseInfo.actualQuerySeqSize.assign(tmpSeq.begin() + 1, tmpSeq.end());
155- baseInfo.querySeqSize = 0U;
156- for (size_t i = 1; i < tmpSeq.size(); ++i) {
157- auto seq = (baseInfo.isCumulativeQuerySeq) ? tmpSeq[i] - tmpSeq[i - 1] : tmpSeq[i];
158- baseInfo.querySeqSize = std::max(baseInfo.querySeqSize, static_cast<uint32_t>(seq));
159- }
160- }
161- return;
162-}
163- 
164-void FlashAttnMetadataCpuKernel::LoadActualKvSeq()
165-{
166- baseInfo.actualKvSeqSize.clear();
167- baseInfo.isCumulativeKvSeq = (layoutKv_ == "TND" || layoutKv_ == "NTD");
168- 
169- if (sequsedKv_ != nullptr && sequsedKv_->GetData() != nullptr) {
170- auto tmpSeq = GetTensorDataAsInt64(sequsedKv_, sequsedKv_->GetTensorShape()->GetDimSize(0));
171- baseInfo.kvSeqSize = static_cast<uint32_t>(*std::max_element(tmpSeq.begin(), tmpSeq.end()));
172- baseInfo.actualKvSeqSize.assign(tmpSeq.begin(), tmpSeq.end());
173- if (baseInfo.isCumulativeKvSeq) {
174- std::partial_sum(tmpSeq.begin(), tmpSeq.end(), baseInfo.actualKvSeqSize.begin());
175- }
176- } else if (cuSeqlensKv_ != nullptr && cuSeqlensKv_->GetData() != nullptr) {
177- auto tmpSeq = GetTensorDataAsInt64(cuSeqlensKv_, cuSeqlensKv_->GetTensorShape()->GetDimSize(0));
178- baseInfo.actualKvSeqSize.assign(tmpSeq.begin() + 1, tmpSeq.end());
179- baseInfo.kvSeqSize = 0U;
180- for (size_t i = 1; i < tmpSeq.size(); ++i) {
181- auto seq = (baseInfo.isCumulativeKvSeq) ? tmpSeq[i] - tmpSeq[i - 1] : tmpSeq[i];
182- baseInfo.kvSeqSize = std::max(baseInfo.kvSeqSize, static_cast<uint32_t>(seq));
183- }
184- }
185- return;
186}243}
187 244 
188bool FlashAttnMetadataCpuKernel::BalanceSchedule(load_balance::SectionStreamKResult &splitRes)245bool FlashAttnMetadataCpuKernel::BalanceSchedule(load_balance::SectionStreamKResult &splitRes)
@@ -34,6 +34,11 @@ private:
34 bool Prepare(CpuKernelContext &ctx);34 bool Prepare(CpuKernelContext &ctx);
35 bool BalanceSchedule(load_balance::SectionStreamKResult &splitRes);35 bool BalanceSchedule(load_balance::SectionStreamKResult &splitRes);
36 bool GenMetadata(load_balance::SectionStreamKResult &splitRes);36 bool GenMetadata(load_balance::SectionStreamKResult &splitRes);
37+ 
38+ bool ParamsCheck();
39+ bool CheckActualQuerySeq();
40+ bool CheckActualKvSeq();
41+ 
37 bool ParamsInit();42 bool ParamsInit();
38 void InitDeviceInfo();43 void InitDeviceInfo();
39 void InitBaseInfo();44 void InitBaseInfo();
@@ -68,6 +73,12 @@ private:
68 int32_t aicCoreNum_ = 36U; // 36: default aic num73 int32_t aicCoreNum_ = 36U; // 36: default aic num
69 int32_t aivCoreNum_ = 72U; // 72: default aiv num74 int32_t aivCoreNum_ = 72U; // 72: default aiv num
70 75 
76+ // BaseInfo
77+ bool isActualSeqlenQAccum_ = false;
78+ bool isActualSeqlenKvAccum_ = false;
79+ std::vector<int64_t> actualSeqlenQ_ {};
80+ std::vector<int64_t> actualSeqlenKv_ {};
81+ 
71 // SplitParams82 // SplitParams
72 uint32_t groupSize_ = 0;83 uint32_t groupSize_ = 0;
73 uint32_t mBaseSize_ = 64; // 64: default value84 uint32_t mBaseSize_ = 64; // 64: default value