已合并
[flash_attn_metadata] 添加防御检查 #9078
FiguraDoge创建于 7月23日
[flash_attn_metadata] 添加防御检查 #9078
已合并
共 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 | ||
| 36 | private: | 36 | private: |
| 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 | + | ||
| 88 | inline aclnnStatus | 95 | inline aclnnStatus |
| 89 | FlashAttnMetadataCheck::CheckBaseAttr(int64_t batchSize, int64_t maxSeqlenQ, int64_t maxSeqlenKv, | 96 | FlashAttnMetadataCheck::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时,必须传入cuSeqlensKvOptional | 186 | // 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时,不可以传入cuSeqlensKvOptional | 193 | // 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 | + | ||
| 89 | void FlashAttnMetadataCpuKernel::InitDeviceInfo() | 195 | void 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 block | 219 | param.fdLeastBlock = 3; // 3: least block |
| 114 | param.fdOn = true; | 220 | param.fdOn = true; |
| 115 | } | 221 | } |
| 222 | + | ||
| 116 | void FlashAttnMetadataCpuKernel::InitBaseInfo() | 223 | void 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 | ||
| 188 | bool FlashAttnMetadataCpuKernel::BalanceSchedule(load_balance::SectionStreamKResult &splitRes) | 245 | bool 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 num | 73 | int32_t aicCoreNum_ = 36U; // 36: default aic num |
| 69 | int32_t aivCoreNum_ = 72U; // 72: default aiv num | 74 | 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 | // SplitParams | 82 | // SplitParams |
| 72 | uint32_t groupSize_ = 0; | 83 | uint32_t groupSize_ = 0; |
| 73 | uint32_t mBaseSize_ = 64; // 64: default value | 84 | uint32_t mBaseSize_ = 64; // 64: default value |