mirror of
https://github.com/opencv/opencv.git
synced 2026-07-29 23:33:05 +04:00
Merge pull request #28840 from nklskyoy:key-value-cache
FP32 KV Cache #28840 OpenCV Extra: https://github.com/opencv/opencv_extra/pull/1348 ## This PR introduces basic (Paged ) KV-Cache to use on CPU ### Summary: 1. To ensure proper gemm-prepacking, 1.1. The Page Size of Key Cache is currently hardcoded as `FAST_GEMM_F32_NR`(which is 8, 12 or 16 depending on CPU architecture) 1.2. The Page Size of Values Cache is hardcoded as `FAST_GEMM_F32_PACKED_STRIDE_K` 2. there are two phases supported - prefill & generate. 2.1. prefill grows cache by `N` tokens and is allowed **only** for empty cache 2.2. generate grows cache by 1 token. 2.3. **Improtant**: it is currently not allowed to grow non-empty cache by more than one token at a time (thisbehaviour is sufficient for normal LLM querying, but should be extended if we want to implement speculative decoding) ### Pull Request Readiness Checklist See details at https://github.com/opencv/opencv/wiki/How_to_contribute#making-a-good-pull-request - [x] I agree to contribute to the project under Apache 2 License. - [x] To the best of my knowledge, the proposed patch is not based on a code under GPL or another license that is incompatible with OpenCV - [ ] The PR is proposed to the proper branch - [ ] There is a reference to the original bug report and related work - [ ] There is accuracy test, performance test and test data in opencv_extra repository, if applicable Patch to opencv_extra has the same branch name. - [ ] The feature is well documented and sample code can be built with the project CMake
This commit is contained in:
@@ -1876,6 +1876,8 @@ CV__DNN_INLINE_NS_BEGIN
|
||||
|
||||
class CV_EXPORTS AttentionOnnxAiLayer : public Layer {
|
||||
public:
|
||||
int kv_num_heads;
|
||||
|
||||
static Ptr<AttentionOnnxAiLayer> create(const LayerParams ¶ms);
|
||||
};
|
||||
|
||||
|
||||
@@ -1027,6 +1027,14 @@ CV__DNN_INLINE_NS_BEGIN
|
||||
*/
|
||||
CV_WRAP int64 getPerfProfile(CV_OUT std::vector<double>& timings);
|
||||
|
||||
/** @brief Enables KV-Cache for all AttentionOnnxI layers */
|
||||
CV_WRAP void enableKVCache();
|
||||
|
||||
/** @brief Disables KV-Cache for all AttentionOnnxI layers */
|
||||
CV_WRAP void disableKVCache();
|
||||
|
||||
/** @brief Resets KV-Cache for all AttentionOnnxI layers */
|
||||
CV_WRAP void resetKVCache();
|
||||
/** @brief Returns profiling data captured during the last forward pass.
|
||||
*
|
||||
* Entries are sorted by time in descending order. Empty vectors are returned
|
||||
@@ -1065,6 +1073,7 @@ CV__DNN_INLINE_NS_BEGIN
|
||||
bool comma=true, bool dump_details=false) const;
|
||||
std::ostream& dumpDim(std::ostream& strm, int value) const;
|
||||
|
||||
|
||||
struct Impl;
|
||||
inline Impl* getImpl() const { return impl.get(); }
|
||||
inline Impl& getImplRef() const { CV_DbgAssert(impl); return *impl.get(); }
|
||||
|
||||
@@ -0,0 +1,252 @@
|
||||
// This file is part of OpenCV project.
|
||||
// It is subject to the license terms in the LICENSE file found in the top-level directory
|
||||
// of this distribution and at http://opencv.org/license.html.
|
||||
|
||||
#include "precomp.hpp"
|
||||
#include "kv_cache_manager.hpp"
|
||||
#include "net_impl.hpp"
|
||||
|
||||
#include <memory>
|
||||
#include <opencv2/core/utils/tls.hpp>
|
||||
#include "layers/cpu_kernels/fast_gemm.hpp"
|
||||
namespace cv { namespace dnn {
|
||||
CV__DNN_INLINE_NS_BEGIN
|
||||
|
||||
|
||||
|
||||
void initKVDataRecursively(const Ptr<Graph>& graph, std::map<std::string, KCache>& kData, std::map<std::string, VCache>& vData, FastGemmOpt& opt) {
|
||||
for (const auto& layer : graph->prog()) {
|
||||
|
||||
for (const auto& subgraph : layer->subgraphs() ? *layer->subgraphs() : std::vector<Ptr<Graph>>()) {
|
||||
initKVDataRecursively(subgraph, kData, vData, opt);
|
||||
}
|
||||
|
||||
if (layer->type == "AttentionOnnxAi") {
|
||||
int kvNumHeads = layer.dynamicCast<AttentionOnnxAiLayer>()->kv_num_heads;
|
||||
|
||||
if (kvNumHeads > 0)
|
||||
{
|
||||
kData.emplace(layer->name, KCache(opt, kvNumHeads));
|
||||
vData.emplace(layer->name, VCache(opt, kvNumHeads));
|
||||
}
|
||||
else
|
||||
{
|
||||
kData.emplace(layer->name, KCache(opt));
|
||||
vData.emplace(layer->name, VCache(opt));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void setKVCacheManager(Ptr<Net::Impl> netimpl)
|
||||
{
|
||||
CV_Assert(netimpl != nullptr);
|
||||
|
||||
CV_Assert(!netimpl->layers.empty());
|
||||
|
||||
auto manager = KVCacheManager();
|
||||
manager.netimpl = netimpl;
|
||||
manager.opt.init();
|
||||
initKVDataRecursively(netimpl->mainGraph, manager.kData, manager.vData, manager.opt);
|
||||
|
||||
manager.isInitialized = true;
|
||||
netimpl->useKVCache = true;
|
||||
netimpl->kvCacheManager = std::move(manager);
|
||||
}
|
||||
|
||||
void KVCacheManager::init()
|
||||
{
|
||||
// Construction of per-layer caches happens in setKVCacheManager.
|
||||
// This method is retained as a callable hook for deferred re-init if needed.
|
||||
CV_Assert(isInitialized);
|
||||
}
|
||||
|
||||
void KVCache::grow(const Mat& newData) {
|
||||
CV_Assert(newData.dims == 4 || newData.dims == 3);
|
||||
|
||||
if (nHeads == -1) {
|
||||
CV_Assert(newData.dims == 4); // to derive shape from data, we need 4D
|
||||
headDim = newData.size[3];
|
||||
nHeads = newData.size[1];
|
||||
batchSize = newData.size[0];
|
||||
} else {
|
||||
if (headDim == -1) {
|
||||
if (newData.dims == 4)
|
||||
headDim = newData.size[3];
|
||||
else{
|
||||
CV_Assert(newData.dims == 3);
|
||||
CV_Assert(newData.size[2] % nHeads == 0);
|
||||
headDim = newData.size[2] / nHeads;
|
||||
}
|
||||
} else {
|
||||
if (newData.dims == 4)
|
||||
CV_Assert(newData.size[3] == headDim);
|
||||
else
|
||||
CV_Assert(newData.size[2] == headDim * nHeads);
|
||||
}
|
||||
if (batchSize == -1)
|
||||
batchSize = newData.size[0];
|
||||
}
|
||||
|
||||
int T = newData.dims == 4 ? newData.size[2] : newData.size[1];
|
||||
|
||||
if (T > 1 || pages.empty()) {
|
||||
// prefetch
|
||||
if(!pages.empty())
|
||||
CV_Error(
|
||||
cv::Error::StsNotImplemented,
|
||||
"storing multiple tokens to a non-empty cache is not supported yet. Either clear the cache (to reenter the prefetch phase) or provide tokens one-by-one"
|
||||
);
|
||||
|
||||
// add pages
|
||||
int totalPages = (T + pageSize - 1) / pageSize;
|
||||
for (int i = 0; i < totalPages; i++) {
|
||||
int page_size = isKCache ?
|
||||
(int)fastGemmPackBSize(pageSize, headDim, opt):
|
||||
(int)fastGemmPackBSize(headDim, pageSize, opt);
|
||||
|
||||
pages.push_back(
|
||||
Mat({batchSize, nHeads, page_size}, CV_32F, Scalar(0))
|
||||
);
|
||||
}
|
||||
growPrefill(newData, T);
|
||||
} else{
|
||||
// generate
|
||||
growGenerate(newData);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
void KVCache::growPrefill(const Mat& newData, int T){
|
||||
int totalPages = (T + pageSize - 1) / pageSize;
|
||||
int ps = isKCache ? (int)fastGemmPackBSize(pageSize, headDim, opt)
|
||||
: (int)fastGemmPackBSize(headDim, pageSize, opt);
|
||||
|
||||
bool is3Dlayout = newData.dims == 3;
|
||||
|
||||
int total = totalPages * batchSize * nHeads;
|
||||
|
||||
cv::TLSData<std::vector<float>> tls_temp_buf;
|
||||
|
||||
auto fn = [&](const Range& range) {
|
||||
for (int i = range.start; i < range.end; i++) {
|
||||
int page = i / (batchSize * nHeads);
|
||||
int b = (i - page * batchSize * nHeads) / nHeads;
|
||||
int h = i % nHeads;
|
||||
|
||||
// source
|
||||
size_t step_source = b * nHeads * T * headDim;
|
||||
if(is3Dlayout)
|
||||
step_source += page * pageSize * nHeads * headDim +
|
||||
h * headDim;
|
||||
else
|
||||
step_source += h * headDim * T + page * pageSize * headDim;
|
||||
const auto* source = newData.ptr<float>() + step_source;
|
||||
|
||||
int chunk_T = std::min(pageSize, T - page * pageSize);
|
||||
const float* actual_source = source;
|
||||
|
||||
int lds = is3Dlayout ? headDim * nHeads : headDim;
|
||||
|
||||
if (chunk_T < pageSize) {
|
||||
std::vector<float>& temp_buf = *tls_temp_buf.get();
|
||||
temp_buf.assign(pageSize * headDim, 0.0f);
|
||||
for (int i = 0; i < chunk_T; i++) {
|
||||
std::memcpy(temp_buf.data() + i * headDim, source + i * lds, headDim * sizeof(float));
|
||||
}
|
||||
actual_source = temp_buf.data();
|
||||
}
|
||||
|
||||
// dst
|
||||
size_t step_dst = b * nHeads * ps +
|
||||
h * ps;
|
||||
auto* dst = pages[page].ptr<float>() + step_dst;
|
||||
|
||||
const int N = headDim;
|
||||
const int K = pageSize;
|
||||
|
||||
fastGemmPackB(
|
||||
isKCache,
|
||||
N, K,
|
||||
actual_source, chunk_T < pageSize ? headDim : lds,
|
||||
dst,
|
||||
opt
|
||||
);
|
||||
}
|
||||
};
|
||||
parallel_for_(Range(0, total), fn);
|
||||
nTokens += T;
|
||||
}
|
||||
|
||||
void KCache::growGenerate(const Mat& newData){
|
||||
int cur_page = nTokens / pageSize;
|
||||
const int Ps = fastGemmPackBSize(pageSize, headDim, opt);
|
||||
int t0 = nTokens % pageSize;
|
||||
const int batch_size = newData.size[0];
|
||||
|
||||
if (cur_page >= (int)pages.size()) {
|
||||
pages.push_back(Mat({batchSize, nHeads, Ps}, CV_32F, Scalar(0)));
|
||||
}
|
||||
|
||||
auto* page = pages[cur_page].ptr<float>();
|
||||
const auto* data = newData.ptr<float>();
|
||||
|
||||
for (int b = 0; b < batch_size; b++){
|
||||
for (int h = 0; h < nHeads; h++){
|
||||
for(int j = 0; j < headDim; j++) {
|
||||
int step =
|
||||
b * nHeads * headDim +
|
||||
h * headDim +
|
||||
j;
|
||||
page[
|
||||
b * nHeads * Ps +
|
||||
h * Ps +
|
||||
t0 + pageSize * j
|
||||
] = *(data + step);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
nTokens += 1;
|
||||
}
|
||||
|
||||
void VCache::growGenerate(const Mat& newData){
|
||||
const int batch_size = newData.size[0];
|
||||
const int Nr = fastGemmNR(opt);
|
||||
const int Ps = fastGemmPackBSize(headDim, pageSize, opt);
|
||||
const int t0 = nTokens % pageSize;
|
||||
const int step_packed = pageSize * Nr;
|
||||
int cur_page = nTokens / pageSize;
|
||||
|
||||
if (cur_page >= (int)pages.size()) {
|
||||
pages.push_back(Mat({batchSize, nHeads, Ps}, CV_32F, Scalar(0)));
|
||||
}
|
||||
|
||||
auto* page = pages[cur_page].ptr<float>();
|
||||
const auto* data = newData.ptr<float>();
|
||||
|
||||
for (int b = 0; b < batch_size; b++){
|
||||
for (int h = 0; h < nHeads; h++){
|
||||
for (int j = 0; j <= (headDim - 1) / Nr; j++) {
|
||||
int step = b * nHeads * headDim + h * headDim + j * Nr;
|
||||
int copy_size = std::min(Nr, headDim - j * Nr);
|
||||
|
||||
auto* cur_page_ptr = page + b * nHeads * Ps + h * Ps + t0 * Nr + j * step_packed;
|
||||
const float* src_ptr = data + step;
|
||||
std::memcpy(cur_page_ptr, src_ptr, copy_size * sizeof(float));
|
||||
if (copy_size < Nr) {
|
||||
float replication_val = src_ptr[0];
|
||||
for (int k = copy_size; k < Nr; k++) {
|
||||
cur_page_ptr[k] = replication_val;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
nTokens += 1;
|
||||
}
|
||||
|
||||
|
||||
CV__DNN_INLINE_NS_END
|
||||
}}
|
||||
@@ -0,0 +1,101 @@
|
||||
// This file is part of OpenCV project.
|
||||
// It is subject to the license terms in the LICENSE file found in the top-level directory
|
||||
// of this distribution and at http://opencv.org/license.html.
|
||||
|
||||
// Copyright (C) 2020, Intel Corporation, all rights reserved.
|
||||
// Third party copyrights are property of their respective owners.
|
||||
|
||||
#ifndef __OPENCV_DNN_KV_CACHE_MANAGER_HPP__
|
||||
#define __OPENCV_DNN_KV_CACHE_MANAGER_HPP__
|
||||
|
||||
#include <opencv2/core.hpp>
|
||||
#include <opencv2/dnn/dnn.hpp>
|
||||
|
||||
#include <map>
|
||||
#include <string>
|
||||
#include <vector>
|
||||
|
||||
#include "layers/cpu_kernels/fast_gemm.hpp"
|
||||
|
||||
namespace cv { namespace dnn {
|
||||
CV__DNN_INLINE_NS_BEGIN
|
||||
|
||||
class KVCache
|
||||
{
|
||||
public:
|
||||
virtual ~KVCache() = default;
|
||||
KVCache(FastGemmOpt opt, int nHeads) : nHeads(nHeads), headDim(-1), offset(0), opt(opt) {}
|
||||
KVCache(FastGemmOpt opt) : nHeads(-1), headDim(-1),offset(0), opt(opt) {}
|
||||
void grow(const Mat& newData);
|
||||
void clear() {
|
||||
pages.clear();
|
||||
nTokens = 0;
|
||||
}
|
||||
const std::vector<Mat>& getPages() const { return pages; }
|
||||
int getPageSize() const { return pageSize; }
|
||||
int getNumTokens() const { return nTokens; }
|
||||
protected:
|
||||
void growPrefill(const Mat& newData, int T);
|
||||
|
||||
virtual void growGenerate(const Mat& newData) = 0;
|
||||
std::vector<Mat> pages;
|
||||
|
||||
int nTokens = 0;
|
||||
int pageSize = -1;
|
||||
int nHeads;
|
||||
int headDim;
|
||||
int batchSize = -1;
|
||||
int offset;
|
||||
bool isKCache = false;
|
||||
FastGemmOpt opt;
|
||||
};
|
||||
|
||||
class VCache : public KVCache
|
||||
{
|
||||
public:
|
||||
VCache(FastGemmOpt opt) : KVCache(opt) {
|
||||
isKCache = false;
|
||||
pageSize = fastGemmKC(opt);
|
||||
}
|
||||
VCache(FastGemmOpt opt, int nHeads) : KVCache(opt, nHeads) {
|
||||
isKCache = false;
|
||||
pageSize = fastGemmKC(opt);
|
||||
}
|
||||
protected:
|
||||
void growGenerate(const Mat& newData) CV_OVERRIDE;
|
||||
};
|
||||
|
||||
class KCache : public KVCache
|
||||
{
|
||||
public:
|
||||
KCache(FastGemmOpt opt) : KVCache(opt) {
|
||||
isKCache = true;
|
||||
pageSize = fastGemmNR(opt);
|
||||
}
|
||||
KCache(FastGemmOpt opt, int nHeads) : KVCache(opt, nHeads) {
|
||||
isKCache = true;
|
||||
pageSize = fastGemmNR(opt);
|
||||
}
|
||||
protected:
|
||||
void growGenerate(const Mat& newData) CV_OVERRIDE;
|
||||
};
|
||||
|
||||
|
||||
struct KVCacheManager
|
||||
{
|
||||
Net::Impl* netimpl = nullptr;
|
||||
std::map<std::string, KCache> kData;
|
||||
std::map<std::string, VCache> vData;
|
||||
FastGemmOpt opt;
|
||||
bool isInitialized = false;
|
||||
|
||||
void init();
|
||||
};
|
||||
|
||||
void setKVCacheManager(Ptr<Net::Impl> netimpl);
|
||||
|
||||
|
||||
CV__DNN_INLINE_NS_END
|
||||
}}
|
||||
|
||||
#endif
|
||||
@@ -81,15 +81,32 @@ class AttentionOnnxAiLayerImpl CV_FINAL : public AttentionOnnxAiLayer {
|
||||
|
||||
const int batch_size = inputs[0][0];
|
||||
const int seq_len_q = inputs[0][input_dims - 2];
|
||||
const int seq_len_kv = inputs[1][input_dims - 2];
|
||||
int seq_len_k = inputs[1][input_dims - 2];
|
||||
int seq_len_v = inputs[2][input_dims - 2];
|
||||
|
||||
Net::Impl* netimpl = getNetImpl(const_cast<AttentionOnnxAiLayerImpl*>(this));
|
||||
if (netimpl && netimpl->useKVCache) {
|
||||
KVCacheManager& kvCacheManager = netimpl->kvCacheManager;
|
||||
if (kvCacheManager.isInitialized) {
|
||||
auto it_k = kvCacheManager.kData.find(name);
|
||||
CV_Assert(it_k != kvCacheManager.kData.end());
|
||||
if (it_k != kvCacheManager.kData.end())
|
||||
seq_len_k += it_k->second.getNumTokens();
|
||||
|
||||
auto it_v = kvCacheManager.vData.find(name);
|
||||
CV_Assert(it_v != kvCacheManager.vData.end());
|
||||
if (it_v != kvCacheManager.vData.end())
|
||||
seq_len_v += it_v->second.getNumTokens();
|
||||
}
|
||||
}
|
||||
|
||||
const int q_hn = input_dims == 4 ?
|
||||
inputs[0][1] : q_num_heads;
|
||||
const int kv_hn =input_dims == 4 ?
|
||||
inputs[1][1] : kv_num_heads;
|
||||
|
||||
CV_CheckTrue(inputs[2][input_dims - 2] == seq_len_kv,
|
||||
"Key and query sequence lengths must be equal");
|
||||
CV_CheckTrue(seq_len_v == seq_len_k,
|
||||
"Key and value sequence lengths must be equal");
|
||||
const int nhq = input_dims == 4 ? inputs[0][1] : q_num_heads;
|
||||
|
||||
CV_CheckTrue(q_hn % kv_hn == 0,
|
||||
@@ -113,7 +130,7 @@ class AttentionOnnxAiLayerImpl CV_FINAL : public AttentionOnnxAiLayer {
|
||||
outputs.push_back(output_shape);
|
||||
}
|
||||
|
||||
MatShape attention_prob_shape{batch_size , nhq, seq_len_q, seq_len_kv};
|
||||
MatShape attention_prob_shape{batch_size , nhq, seq_len_q, seq_len_k};
|
||||
internals.push_back(attention_prob_shape);
|
||||
|
||||
return false;
|
||||
@@ -143,6 +160,13 @@ class AttentionOnnxAiLayerImpl CV_FINAL : public AttentionOnnxAiLayer {
|
||||
void forward(InputArrayOfArrays inputs_arr, OutputArrayOfArrays outputs_arr, OutputArrayOfArrays internals_arr) CV_OVERRIDE {
|
||||
opt.init();
|
||||
|
||||
Net::Impl* netimpl = getNetImpl(this);
|
||||
bool with_kv_cache = false;
|
||||
|
||||
if (netimpl && netimpl->useKVCache) {
|
||||
with_kv_cache = netimpl->kvCacheManager.isInitialized;
|
||||
}
|
||||
|
||||
if (inputs_arr.depth() == CV_16F)
|
||||
{
|
||||
forward_fallback(inputs_arr, outputs_arr, internals_arr);
|
||||
@@ -153,78 +177,109 @@ class AttentionOnnxAiLayerImpl CV_FINAL : public AttentionOnnxAiLayer {
|
||||
inputs_arr.getMatVector(inputs);
|
||||
outputs_arr.getMatVector(outputs);
|
||||
internals_arr.getMatVector(internals);
|
||||
Mat &attention_prob = internals[internals.size() - 1];
|
||||
|
||||
const int input_dims = inputs[0].dims;
|
||||
|
||||
const int batch_size = inputs[0].size[0];
|
||||
|
||||
const int seq_len_q = input_dims == 3 ?
|
||||
inputs[0].size[1]:
|
||||
inputs[0].size[2];
|
||||
|
||||
int seq_len_kv = input_dims == 3 ?
|
||||
inputs[1].size[1]:
|
||||
inputs[1].size[2];
|
||||
|
||||
const int nhq = input_dims == 3 ?
|
||||
q_num_heads :
|
||||
inputs[0].size[1];
|
||||
const int nhkv = input_dims == 3 ?
|
||||
kv_num_heads :
|
||||
inputs[1].size[1];
|
||||
q_num_heads :
|
||||
inputs[0].size[1];
|
||||
|
||||
const int qk_head_size = input_dims == 3 ?
|
||||
inputs[0].size[2] / nhq :
|
||||
inputs[0].size[3];
|
||||
|
||||
const int nhkv = input_dims == 3 ?
|
||||
kv_num_heads :
|
||||
inputs[1].size[1];
|
||||
|
||||
const int v_head_size = input_dims == 3 ?
|
||||
inputs[2].size[2] / nhkv :
|
||||
inputs[2].size[3];
|
||||
const int num_gq_groups = nhq / nhkv;
|
||||
const int seq_len_q = input_dims == 3 ?
|
||||
inputs[0].size[1]:
|
||||
inputs[0].size[2];
|
||||
const int seq_len_kv = input_dims == 3 ?
|
||||
inputs[1].size[1]:
|
||||
inputs[1].size[2];
|
||||
const auto seq_len_square = seq_len_q * seq_len_kv;
|
||||
inputs[2].size[2] / nhkv :
|
||||
inputs[2].size[3];
|
||||
|
||||
const auto* Q = inputs[0].ptr<const float>();
|
||||
const auto* K = inputs[1].ptr<const float>();
|
||||
const auto* V = inputs[2].ptr<const float>();
|
||||
std::vector<size_t> _q_offsets, _k_offsets, _v_offsets,
|
||||
_a_offsets, _o_offsets;
|
||||
|
||||
scale = is_scale_set ? scale : 1.0f / std::sqrt(static_cast<float>(qk_head_size));
|
||||
if (with_kv_cache){
|
||||
KVCacheManager& kvCacheManager = netimpl->kvCacheManager;
|
||||
|
||||
std::vector<size_t> _q_offsets(nhq * batch_size),
|
||||
_k_offsets(nhq * batch_size),
|
||||
_a_offsets(nhq * batch_size),
|
||||
_v_offsets(nhq * batch_size),
|
||||
_o_offsets(nhq * batch_size);
|
||||
auto it_k = kvCacheManager.kData.find(name);
|
||||
CV_Assert(it_k != kvCacheManager.kData.end());
|
||||
KCache&kData = it_k->second;
|
||||
|
||||
for (int b = 0; b < batch_size; b++)
|
||||
for (int n = 0; n < nhq; n++){
|
||||
_q_offsets[b * nhq + n] =
|
||||
b * seq_len_q * qk_head_size * nhq +
|
||||
(input_dims == 3 ? n * qk_head_size : n * qk_head_size * seq_len_q);
|
||||
_k_offsets[b * nhq + n] =
|
||||
b * seq_len_kv * qk_head_size * nhkv +
|
||||
(n / num_gq_groups * qk_head_size) * (input_dims == 3 ? 1 : seq_len_kv);
|
||||
_v_offsets[b * nhq + n] =
|
||||
b * seq_len_kv * v_head_size * nhkv +
|
||||
(n / num_gq_groups * v_head_size) * (input_dims == 3 ? 1 : seq_len_kv);
|
||||
_a_offsets[b * nhq + n] =
|
||||
b * seq_len_square * nhq +
|
||||
n * seq_len_square;
|
||||
_o_offsets[b * nhq + n] =
|
||||
b * seq_len_q * v_head_size * nhq +
|
||||
(input_dims == 3 ? n * v_head_size : n * v_head_size * seq_len_q);
|
||||
}
|
||||
kData.grow(inputs[1]);
|
||||
|
||||
const int ldq0 = input_dims == 3 ? qk_head_size * nhq : qk_head_size;
|
||||
const int ldk0 = input_dims == 3 ? qk_head_size * nhkv : qk_head_size;
|
||||
auto &attention_prob = internals[internals.size() - 1];
|
||||
const std::vector<Mat>& kCachePages = kData.getPages();
|
||||
seq_len_kv = kData.getNumTokens();
|
||||
|
||||
fastGemmBatch(
|
||||
batch_size * nhq,
|
||||
_q_offsets.data(), _k_offsets.data(), _a_offsets.data(),
|
||||
seq_len_q, seq_len_kv, qk_head_size , scale,
|
||||
Q, ldq0, 1,
|
||||
K, 1, ldk0,
|
||||
0.f,
|
||||
attention_prob.ptr<float>(), seq_len_kv,
|
||||
opt
|
||||
);
|
||||
scale = is_scale_set ? scale : 1.0f / std::sqrt(static_cast<float>(qk_head_size));
|
||||
|
||||
pagedAttnQKGemm(
|
||||
inputs[0], kCachePages, attention_prob,
|
||||
seq_len_q, nhq, nhkv, kData.getPageSize(),
|
||||
qk_head_size, seq_len_kv,
|
||||
scale, opt
|
||||
);
|
||||
} else {
|
||||
const auto* Q = inputs[0].ptr<const float>();
|
||||
const auto* K = inputs[1].ptr<const float>();
|
||||
|
||||
const int num_gq_groups = nhq / nhkv;
|
||||
|
||||
const auto seq_len_square = seq_len_q * seq_len_kv;
|
||||
|
||||
scale = is_scale_set ? scale : 1.0f / std::sqrt(static_cast<float>(qk_head_size));
|
||||
|
||||
_q_offsets.resize(nhq * batch_size);
|
||||
_k_offsets.resize(nhq * batch_size);
|
||||
_a_offsets.resize(nhq * batch_size);
|
||||
_v_offsets.resize(nhq * batch_size);
|
||||
_o_offsets.resize(nhq * batch_size);
|
||||
|
||||
for (int b = 0; b < batch_size; b++)
|
||||
for (int n = 0; n < nhq; n++){
|
||||
_q_offsets[b * nhq + n] =
|
||||
b * seq_len_q * qk_head_size * nhq +
|
||||
(input_dims == 3 ? n * qk_head_size : n * qk_head_size * seq_len_q);
|
||||
_k_offsets[b * nhq + n] =
|
||||
b * seq_len_kv * qk_head_size * nhkv +
|
||||
(n / num_gq_groups * qk_head_size) * (input_dims == 3 ? 1 : seq_len_kv);
|
||||
_v_offsets[b * nhq + n] =
|
||||
b * seq_len_kv * v_head_size * nhkv +
|
||||
(n / num_gq_groups * v_head_size) * (input_dims == 3 ? 1 : seq_len_kv);
|
||||
_a_offsets[b * nhq + n] =
|
||||
b * seq_len_square * nhq +
|
||||
n * seq_len_square;
|
||||
_o_offsets[b * nhq + n] =
|
||||
b * seq_len_q * v_head_size * nhq +
|
||||
(input_dims == 3 ? n * v_head_size : n * v_head_size * seq_len_q);
|
||||
}
|
||||
|
||||
const int ldq0 = input_dims == 3 ? qk_head_size * nhq : qk_head_size;
|
||||
const int ldk0 = input_dims == 3 ? qk_head_size * nhkv : qk_head_size;
|
||||
auto &attention_prob = internals[internals.size() - 1];
|
||||
|
||||
fastGemmBatch(
|
||||
batch_size * nhq,
|
||||
_q_offsets.data(), _k_offsets.data(), _a_offsets.data(),
|
||||
seq_len_q, seq_len_kv, qk_head_size , scale,
|
||||
Q, ldq0, 1,
|
||||
K, 1, ldk0,
|
||||
0.f,
|
||||
attention_prob.ptr<float>(), seq_len_kv,
|
||||
opt
|
||||
);
|
||||
}
|
||||
|
||||
fused_softmax_softcap_mask(
|
||||
attention_prob,
|
||||
@@ -236,25 +291,41 @@ class AttentionOnnxAiLayerImpl CV_FINAL : public AttentionOnnxAiLayer {
|
||||
is_causal
|
||||
);
|
||||
|
||||
if (with_kv_cache){
|
||||
KVCacheManager& kvCacheManager = netimpl->kvCacheManager;
|
||||
auto it_v = kvCacheManager.vData.find(name);
|
||||
CV_Assert(it_v != kvCacheManager.vData.end());
|
||||
VCache& vData = it_v->second;
|
||||
|
||||
const int ldv0 = input_dims == 3 ? v_head_size * nhkv : v_head_size;
|
||||
const int ldout = input_dims == 3 ? v_head_size * nhq : v_head_size;
|
||||
vData.grow(inputs[2]);
|
||||
seq_len_kv = vData.getNumTokens();
|
||||
|
||||
fastGemmBatch(
|
||||
batch_size * nhq,
|
||||
_a_offsets.data(), _v_offsets.data(), _o_offsets.data(),
|
||||
seq_len_q, v_head_size, seq_len_kv, 1.f,
|
||||
attention_prob.ptr<float>(), seq_len_kv, 1,
|
||||
V, ldv0, 1,
|
||||
0.f,
|
||||
outputs[0].ptr<float>(), ldout,
|
||||
opt
|
||||
);
|
||||
pagedAttnAVGemm(
|
||||
attention_prob, vData.getPages(), outputs[0],
|
||||
seq_len_q, nhq, nhkv, vData.getPageSize() , v_head_size, seq_len_kv,
|
||||
opt
|
||||
);
|
||||
} else {
|
||||
const auto* V = inputs[2].ptr<const float>();
|
||||
|
||||
const int ldv0 = input_dims == 3 ? v_head_size * nhkv : v_head_size;
|
||||
const int ldout = input_dims == 3 ? v_head_size * nhq : v_head_size;
|
||||
|
||||
fastGemmBatch(
|
||||
batch_size * nhq,
|
||||
_a_offsets.data(), _v_offsets.data(), _o_offsets.data(),
|
||||
seq_len_q, v_head_size, seq_len_kv, 1.f,
|
||||
attention_prob.ptr<float>(), seq_len_kv, 1,
|
||||
V, ldv0, 1,
|
||||
0.f,
|
||||
outputs[0].ptr<float>(), ldout,
|
||||
opt
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
bool is_causal;
|
||||
int kv_num_heads;
|
||||
int q_num_heads;
|
||||
int qk_matmul_output_mode;
|
||||
float scale;
|
||||
|
||||
@@ -99,6 +99,34 @@ int fastGemmKC(const FastGemmOpt &opt) {
|
||||
}
|
||||
}
|
||||
|
||||
int fastGemmNR(const FastGemmOpt &opt) {
|
||||
#if CV_TRY_NEON
|
||||
if (opt.use_neon) {
|
||||
return opt_NEON::fastGemmNR();
|
||||
} else
|
||||
#endif
|
||||
#if CV_TRY_AVX2
|
||||
if (opt.use_avx2) {
|
||||
return opt_AVX2::fastGemmNR();
|
||||
} else
|
||||
#endif
|
||||
#if CV_TRY_AVX
|
||||
if (opt.use_avx) {
|
||||
return opt_AVX::fastGemmNR();
|
||||
} else
|
||||
#endif
|
||||
#if CV_TRY_LASX
|
||||
if (opt.use_lasx) {
|
||||
return opt_LASX::fastGemmNR();
|
||||
} else
|
||||
#endif
|
||||
{
|
||||
return cpu_baseline::fastGemmNR();
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
size_t fastGemmPackBSize(size_t N, size_t K, const FastGemmOpt &opt) {
|
||||
#if CV_TRY_NEON
|
||||
if (opt.use_neon) {
|
||||
@@ -768,8 +796,8 @@ void fastGemmBatch(size_t batch,
|
||||
|
||||
void pagedAttnQKGemm(
|
||||
const Mat& Q,const std::vector<Mat> &K, Mat& A,
|
||||
int T_q, int Nq, int N_k, int T_s, int D,
|
||||
const FastGemmOpt &opt
|
||||
int T_q, int Nq, int N_k, int T_s, int D, size_t T_k,
|
||||
float sm_scale, const FastGemmOpt &opt
|
||||
) {
|
||||
size_t esz = Q.elemSize();
|
||||
|
||||
@@ -780,7 +808,7 @@ void pagedAttnQKGemm(
|
||||
CV_CheckTypeEQ(Q.type(), A.type(), "pagedAttnQKGemmKernel: Q and A should have the same type");
|
||||
|
||||
CV_CheckTrue(
|
||||
T_s % fastGemmNC(opt) == 0,
|
||||
T_s % fastGemmNR(opt) == 0,
|
||||
"pagedAttnQKGemmKernel: T_s should be divisible by the macro tile size"
|
||||
);
|
||||
|
||||
@@ -801,7 +829,7 @@ void pagedAttnQKGemm(
|
||||
);
|
||||
CV_CheckEQ(shape_k[0], B, "pagedAttnQKGemmKernel: the batch size of K should be the same as A");
|
||||
CV_CheckEQ(shape_k[1], N_k, "pagedAttnQKGemmKernel: the number of heads in K should match that of Q");
|
||||
CV_CheckEQ(shape_k[2], D * T_s, "pagedAttnQKGemmKernel: the head dimension of K should match that of Q and A");
|
||||
CV_Assert(shape_k[2] == D * T_s);
|
||||
}
|
||||
|
||||
std::vector<const char*> packed_K;
|
||||
@@ -815,8 +843,8 @@ void pagedAttnQKGemm(
|
||||
if (opt.use_neon)
|
||||
opt_NEON::pagedAttnQKGemmKernel(
|
||||
Q.ptr<const char>(), packed_K, a,
|
||||
B, T_q, Nq, N_k, T_s, D,
|
||||
esz, isQ3D
|
||||
B, T_q, Nq, N_k, T_s, D, T_k,
|
||||
sm_scale, esz, isQ3D
|
||||
);
|
||||
else
|
||||
#endif
|
||||
@@ -824,8 +852,8 @@ void pagedAttnQKGemm(
|
||||
if (opt.use_avx2)
|
||||
opt_AVX2::pagedAttnQKGemmKernel(
|
||||
Q.ptr<const char>(), packed_K, a,
|
||||
B, T_q, Nq, N_k, T_s, D,
|
||||
esz, isQ3D
|
||||
B, T_q, Nq, N_k, T_s, D, T_k,
|
||||
sm_scale, esz, isQ3D
|
||||
);
|
||||
else
|
||||
#endif
|
||||
@@ -833,8 +861,8 @@ void pagedAttnQKGemm(
|
||||
if (opt.use_avx)
|
||||
opt_AVX::pagedAttnQKGemmKernel(
|
||||
Q.ptr<const char>(), packed_K, a,
|
||||
B, T_q, Nq, N_k, T_s, D,
|
||||
esz, isQ3D
|
||||
B, T_q, Nq, N_k, T_s, D, T_k,
|
||||
sm_scale, esz, isQ3D
|
||||
);
|
||||
else
|
||||
#endif
|
||||
@@ -842,15 +870,15 @@ void pagedAttnQKGemm(
|
||||
if (opt.use_lasx)
|
||||
opt_LASX::pagedAttnQKGemmKernel(
|
||||
Q.ptr<const char>(), packed_K, a,
|
||||
B, T_q, Nq, N_k, T_s, D,
|
||||
esz, isQ3D
|
||||
B, T_q, Nq, N_k, T_s, D, T_k,
|
||||
sm_scale, esz, isQ3D
|
||||
);
|
||||
else
|
||||
#endif
|
||||
cpu_baseline::pagedAttnQKGemmKernel(
|
||||
Q.ptr<const char>(), packed_K, a,
|
||||
B, T_q, Nq, N_k, T_s, D,
|
||||
esz, isQ3D
|
||||
B, T_q, Nq, N_k, T_s, D, T_k,
|
||||
sm_scale, esz, isQ3D
|
||||
);
|
||||
|
||||
|
||||
@@ -859,7 +887,7 @@ void pagedAttnQKGemm(
|
||||
|
||||
void pagedAttnAVGemm(
|
||||
const Mat& A,const std::vector<Mat> &V, Mat& Out,
|
||||
int T_q, int Nq, int N_k, int T_s, int D,
|
||||
int T_q, int Nq, int N_k, int T_s, int D, int T_v,
|
||||
const FastGemmOpt &opt
|
||||
) {
|
||||
size_t esz = A.elemSize();
|
||||
@@ -904,7 +932,7 @@ void pagedAttnAVGemm(
|
||||
if (opt.use_neon)
|
||||
opt_NEON::pagedAttnAVGemmKernel(
|
||||
A.ptr<const char>(), packed_V, Out.ptr<char>(),
|
||||
B, T_q, Nq, N_k, T_s, D,
|
||||
B, T_q, Nq, N_k, T_s, D, T_v,
|
||||
esz, canonical_layout, fastGemmPackBSize(D, T_s, opt)
|
||||
);
|
||||
else
|
||||
@@ -913,7 +941,7 @@ void pagedAttnAVGemm(
|
||||
if (opt.use_avx2) {
|
||||
opt_AVX2::pagedAttnAVGemmKernel(
|
||||
A.ptr<const char>(), packed_V, Out.ptr<char>(),
|
||||
B, T_q, Nq, N_k, T_s, D,
|
||||
B, T_q, Nq, N_k, T_s, D, T_v,
|
||||
esz, canonical_layout, fastGemmPackBSize(D, T_s, opt)
|
||||
);
|
||||
}
|
||||
@@ -923,7 +951,7 @@ void pagedAttnAVGemm(
|
||||
if (opt.use_avx){
|
||||
opt_AVX::pagedAttnAVGemmKernel(
|
||||
A.ptr<const char>(), packed_V, Out.ptr<char>(),
|
||||
B, T_q, Nq, N_k, T_s, D,
|
||||
B, T_q, Nq, N_k, T_s, D, T_v,
|
||||
esz, canonical_layout, fastGemmPackBSize(D, T_s, opt)
|
||||
);
|
||||
}
|
||||
@@ -933,7 +961,7 @@ void pagedAttnAVGemm(
|
||||
if (opt.use_lasx){
|
||||
opt_LASX::pagedAttnAVGemmKernel(
|
||||
A.ptr<const char>(), packed_V, Out.ptr<char>(),
|
||||
B, T_q, Nq, N_k, T_s, D,
|
||||
B, T_q, Nq, N_k, T_s, D, T_v,
|
||||
esz, canonical_layout, fastGemmPackBSize(D, T_s, opt)
|
||||
);
|
||||
}
|
||||
@@ -942,7 +970,7 @@ void pagedAttnAVGemm(
|
||||
{
|
||||
cpu_baseline::pagedAttnAVGemmKernel(
|
||||
A.ptr<const char>(), packed_V, Out.ptr<char>(),
|
||||
B, T_q, Nq, N_k, T_s, D,
|
||||
B, T_q, Nq, N_k, T_s, D, T_v,
|
||||
esz, canonical_layout, fastGemmPackBSize(D, T_s, opt)
|
||||
);
|
||||
|
||||
|
||||
@@ -159,6 +159,7 @@ void fastGemmPackB(const Mat &m, std::vector<float> &packed_B, bool trans, FastG
|
||||
int fastGemmMC(const FastGemmOpt &opt);
|
||||
int fastGemmNC(const FastGemmOpt &opt);
|
||||
int fastGemmKC(const FastGemmOpt &opt);
|
||||
int fastGemmNR(const FastGemmOpt &opt);
|
||||
|
||||
void fastGemmPackB(bool trans, size_t N, size_t K, const float *B, size_t ldb, float *packed_B, const FastGemmOpt &opt);
|
||||
|
||||
@@ -196,12 +197,12 @@ void fastGemmThin(int M, int N, int K, float alpha,
|
||||
|
||||
void pagedAttnQKGemm(
|
||||
const Mat& Q, const std::vector<Mat> &K, Mat& A,
|
||||
int T_q, int Nq, int N_k, int T_s, int D,
|
||||
const FastGemmOpt &opts
|
||||
int T_q, int Nq, int N_k, int T_s, int D, size_t T_k,
|
||||
float sm_scale, const FastGemmOpt &opts
|
||||
);
|
||||
void pagedAttnAVGemm(
|
||||
const Mat& A,const std::vector<Mat> &V, Mat& Out,
|
||||
int T_q, int Nq, int N_k, int T_s, int D,
|
||||
int T_q, int Nq, int N_k, int T_s, int D, int T_v,
|
||||
const FastGemmOpt &opt
|
||||
);
|
||||
|
||||
|
||||
@@ -100,13 +100,13 @@ void fastGemmBatchKernel(size_t batch, const size_t *A_offsets, const size_t *B_
|
||||
const char *packed_B, float beta, char *C, size_t ldc, size_t esz);
|
||||
void pagedAttnAVGemmKernel(
|
||||
const char* A, const std::vector<const char *> &V, char*Out,
|
||||
size_t B, size_t T_a, size_t Nq, size_t N_k, size_t T_s, size_t D,
|
||||
size_t B, size_t T_a, size_t Nq, size_t N_k, size_t T_s, size_t D, size_t T_v,
|
||||
size_t esz, bool canonical_layout, size_t packed_stride
|
||||
);
|
||||
void pagedAttnQKGemmKernel(
|
||||
const char *Q, const std::vector<const char *> &K, char *A,
|
||||
size_t B, size_t T_q, size_t Nq, size_t N_k, size_t T_s, size_t D,
|
||||
size_t esz, bool isQ3d
|
||||
size_t B, size_t T_q, size_t Nq, size_t N_k, size_t T_s, size_t D, size_t T_k,
|
||||
float sm_scale, size_t esz, bool isQ3d
|
||||
);
|
||||
FAST_GEMM_IMPLEMENT_PACK(8, _f32, float, float)
|
||||
FAST_GEMM_IMPLEMENT_PACK(12, _f32, float, float)
|
||||
@@ -130,6 +130,10 @@ int fastGemmKC() {
|
||||
return FAST_GEMM_F32_PACKED_STRIDE_K;
|
||||
}
|
||||
|
||||
int fastGemmNR() {
|
||||
return FAST_GEMM_F32_NR;
|
||||
}
|
||||
|
||||
void fastGemmPackBKernel(const char *B, char *packed_B, size_t N, size_t K, size_t ldb0, size_t ldb1, size_t esz) {
|
||||
size_t GEMM_NC = FAST_GEMM_F32_NC, GEMM_NR = FAST_GEMM_F32_NR;
|
||||
size_t NC = (((GEMM_NC < N ? GEMM_NC : N) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR;
|
||||
@@ -495,8 +499,8 @@ void fastGemmBatchKernel(size_t batch, const size_t *A_offsets, const size_t *B_
|
||||
// tot. seq len T = T_s * S
|
||||
void pagedAttnQKGemmKernel(
|
||||
const char *Q, const std::vector<const char *> &K, char *A,
|
||||
size_t B, size_t T_q, size_t Nq, size_t N_k, size_t T_s, size_t D,
|
||||
size_t esz, bool isQ3d
|
||||
size_t B, size_t T_q, size_t Nq, size_t N_k, size_t T_s, size_t D, size_t T_k,
|
||||
float sm_scale, size_t esz, bool isQ3d
|
||||
) {
|
||||
size_t GEMM_MC = static_cast<size_t>(FAST_GEMM_F32_MC),
|
||||
GEMM_NC = static_cast<size_t>(FAST_GEMM_F32_NC),
|
||||
@@ -514,11 +518,14 @@ void pagedAttnQKGemmKernel(
|
||||
const size_t tiles_per_mat = m_tiles * n_tiles;
|
||||
const size_t S = K.size();
|
||||
const size_t T = S * T_s;
|
||||
const size_t T_r = T_s - (T - T_k);
|
||||
const int n_kq_groups = Nq / N_k;
|
||||
size_t batch = Nq * B;
|
||||
size_t batch = Nq * B;
|
||||
// batch = B x N_q
|
||||
int total = S * batch * tiles_per_mat;
|
||||
|
||||
size_t packed_stride = fastGemmPackBSize(T_s, D);
|
||||
|
||||
auto fn = [&](const Range &r) {
|
||||
cv::AutoBuffer<char, FAST_GEMM_MAX_STACKBUF> packed_q_buff;
|
||||
packed_q_buff.allocate(buff_size);
|
||||
@@ -543,38 +550,48 @@ void pagedAttnQKGemmKernel(
|
||||
size_t mc = T_q - i0 < MC ? T_q - i0 : MC;
|
||||
size_t nc = T_s - j0 < NC ? T_s - j0 : NC;
|
||||
|
||||
size_t rc = nc;
|
||||
if (s == S - 1) {
|
||||
if (j0 >= T_r) continue; // Completely out of bounds, skip this tile
|
||||
rc = std::min(nc, T_r - j0);
|
||||
}
|
||||
|
||||
size_t q_offset = b * Nq * T_q * D +
|
||||
(isQ3d ? nq * D : nq * T_q * D);
|
||||
const char *q_block = Q + q_offset * esz;
|
||||
|
||||
size_t k_offset = b * N_k * T_s * D +
|
||||
n_k * T_s * D +
|
||||
size_t k_offset = b * N_k * packed_stride +
|
||||
n_k * packed_stride +
|
||||
j0 * D;
|
||||
const char *k_block = (const char *)K[s] + k_offset * esz;
|
||||
|
||||
// save result to A[b, n_q, : , T_s * s : T_s * (s + 1)]
|
||||
const int a_offset = b * Nq * T_q * T +
|
||||
nq * T_q * T +
|
||||
T_s * s +
|
||||
i0 * T + j0;
|
||||
const size_t a_offset = b * Nq * T_q * T_k +
|
||||
nq * T_q * T_k +
|
||||
T_s * s +
|
||||
i0 * T_k + j0;
|
||||
char* a_block = A + a_offset * esz;
|
||||
|
||||
int _nc = static_cast<int>((nc + GEMM_NR - 1) / GEMM_NR) * GEMM_NR * esz;
|
||||
for(size_t i = 0; i < mc; i++) {
|
||||
memset(a_block + i * T_k * esz, 0, rc * esz);
|
||||
}
|
||||
|
||||
for(int k0 = 0; k0 < D; k0 += KC)
|
||||
int _nc = static_cast<int>((nc + GEMM_NR - 1) / GEMM_NR) * GEMM_NR * esz;
|
||||
for(int k0 = 0; k0 < (int)D; k0 += (int)KC)
|
||||
{
|
||||
int kc = D - k0 < KC ? D - k0 : KC;
|
||||
int kc = D - k0 < (int)KC ? (int)D - k0 : (int)KC;
|
||||
// pack q
|
||||
size_t step_q = (i0 * ldq0 + k0) * esz;
|
||||
fast_gemm_pack8_f32(mc, kc, q_block + step_q, ldq0, 1, packed_q);
|
||||
// run kernel
|
||||
fast_gemm_macro_kernel(mc, nc, kc, packed_q, k_block, 1.f, a_block, T, esz);
|
||||
fast_gemm_macro_kernel(mc, rc, kc, packed_q, k_block, sm_scale, a_block, T_k, esz);
|
||||
k_block += _nc * kc;
|
||||
}
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
|
||||
int cost_per_thread = static_cast<int>((D / KC) * (MC / GEMM_MR) * (NC / GEMM_NR));
|
||||
double nstripes = (size_t)total * cost_per_thread * (1 / 1024.0);
|
||||
parallel_for_(Range(0, total), fn, nstripes);
|
||||
@@ -589,8 +606,8 @@ void pagedAttnQKGemmKernel(
|
||||
// tot. seq len T = T_s * S
|
||||
void pagedAttnAVGemmKernel(
|
||||
const char* A, const std::vector<const char *> &V, char*Out,
|
||||
size_t B, size_t T_a, size_t Nq, size_t N_k, size_t T_s, size_t D,
|
||||
size_t esz, bool canonical_layout, size_t packed_stride
|
||||
size_t B, size_t T_a, size_t Nq, size_t N_k, size_t T_s, size_t D, size_t T_v,
|
||||
size_t esz, bool canonical_layout, size_t packed_stride
|
||||
) {
|
||||
|
||||
size_t GEMM_MC = static_cast<size_t>(FAST_GEMM_F32_MC),
|
||||
@@ -598,12 +615,9 @@ void pagedAttnAVGemmKernel(
|
||||
GEMM_MR = static_cast<size_t>(FAST_GEMM_F32_MR),
|
||||
GEMM_NR = static_cast<size_t>(FAST_GEMM_F32_NR);
|
||||
|
||||
const size_t S = V.size();
|
||||
const size_t T = S * T_s;
|
||||
|
||||
size_t MC = (((GEMM_MC < T_a ? GEMM_MC : T_a) + GEMM_MR - 1) / GEMM_MR) * GEMM_MR;
|
||||
size_t NC = (((GEMM_NC < D ? GEMM_NC : D) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR;
|
||||
size_t KC = std::min( static_cast<size_t>(FAST_GEMM_F32_PACKED_STRIDE_K), T);
|
||||
size_t KC = std::min( static_cast<size_t>(FAST_GEMM_F32_PACKED_STRIDE_K), T_v);
|
||||
|
||||
size_t buff_size = KC * MC * esz;
|
||||
|
||||
@@ -619,9 +633,13 @@ void pagedAttnAVGemmKernel(
|
||||
packed_a_buff.allocate(buff_size);
|
||||
char* packed_a = packed_a_buff.data();
|
||||
|
||||
cv::AutoBuffer<char, FAST_GEMM_MAX_STACKBUF> repacked_v_buff;
|
||||
repacked_v_buff.allocate(NC * T_s * esz);
|
||||
char* repacked_v = repacked_v_buff.data();
|
||||
|
||||
size_t start = r.start;
|
||||
size_t end = r.end;
|
||||
size_t ldc0 = canonical_layout ? D : Nq * D;
|
||||
size_t ldc0 = canonical_layout ? Nq * D : D;
|
||||
|
||||
for (size_t tile_idx = start; tile_idx < end; tile_idx++) {
|
||||
size_t idx = tile_idx / tiles_per_mat;
|
||||
@@ -638,9 +656,9 @@ void pagedAttnAVGemmKernel(
|
||||
size_t mc = T_a - i0 < MC ? T_a - i0 : MC;
|
||||
size_t nc = D - j0 < NC ? D - j0 : NC;
|
||||
|
||||
const size_t a_offset = b * Nq * T_a * T +
|
||||
nq * T_a * T +
|
||||
i0 * T;
|
||||
const size_t a_offset = b * Nq * T_a * T_v +
|
||||
nq * T_a * T_v +
|
||||
i0 * T_v;
|
||||
const char*a_block = A + a_offset * esz;
|
||||
|
||||
// start at the 0th row
|
||||
@@ -651,35 +669,49 @@ void pagedAttnAVGemmKernel(
|
||||
|
||||
size_t o_offset = b * Nq * T_a * D;
|
||||
if (canonical_layout)
|
||||
o_offset += nq * T_a * D +
|
||||
i0 * D + j0;
|
||||
o_offset += i0 * Nq * D + nq * D + j0;
|
||||
else
|
||||
o_offset += i0 * Nq * D +
|
||||
nq * D + j0;
|
||||
o_offset += nq * T_a * D + i0 * D + j0;
|
||||
char* out_block = Out + o_offset * esz;
|
||||
|
||||
int _nc = static_cast<int>((nc + GEMM_NR - 1) / GEMM_NR) * GEMM_NR * esz;
|
||||
for(size_t i = 0; i < mc; i++) {
|
||||
memset(out_block + i * ldc0 * esz, 0, nc * esz);
|
||||
}
|
||||
|
||||
for(int k0 = 0; k0 < T; k0 += KC)
|
||||
// int _nc = static_cast<int>((nc + GEMM_NR - 1) / GEMM_NR) * GEMM_NR * esz;
|
||||
|
||||
for(int k0 = 0; k0 < (int)T_v; )
|
||||
{
|
||||
// T_s should be divisible by KC
|
||||
// pack a
|
||||
fast_gemm_pack8_f32(mc, KC, a_block + k0 * esz, T, 1, packed_a);
|
||||
|
||||
// get the k-block
|
||||
size_t sk = k0 / T_s; // which page
|
||||
size_t k = k0 % T_s;
|
||||
size_t v_offset = v_offset_base * esz + k * _nc;
|
||||
int kc = std::min((int)KC, (int)T_v - k0);
|
||||
kc = std::min(kc, (int)T_s - (int)k);
|
||||
|
||||
size_t v_offset = v_offset_base * esz + k * GEMM_NR * esz;
|
||||
const char *v_block = V[sk] + v_offset;
|
||||
|
||||
// pack a
|
||||
fast_gemm_pack8_f32(mc, kc, a_block + k0 * esz, T_v, 1, packed_a);
|
||||
|
||||
const char* v_block_to_use = v_block;
|
||||
if (kc < (int)T_s) {
|
||||
for (size_t j = 0; j < nc; j += GEMM_NR) {
|
||||
const char *src_panel = v_block + j * T_s * esz;
|
||||
char *dst_panel = repacked_v + j * kc * esz;
|
||||
memcpy(dst_panel, src_panel, kc * GEMM_NR * esz);
|
||||
}
|
||||
v_block_to_use = repacked_v;
|
||||
}
|
||||
|
||||
// run kernel
|
||||
fast_gemm_macro_kernel(mc, nc, KC, packed_a, v_block, 1.f, out_block, ldc0, esz);
|
||||
fast_gemm_macro_kernel(mc, nc, kc, packed_a, v_block_to_use, 1.f, out_block, ldc0, esz);
|
||||
k0 += kc;
|
||||
}
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
int cost_per_thread = static_cast<int>((T / KC) * (MC / GEMM_MR) * (NC / GEMM_NR));
|
||||
int cost_per_thread = static_cast<int>((T_v / KC) * (MC / GEMM_MR) * (NC / GEMM_NR));
|
||||
double nstripes = (size_t)total * cost_per_thread * (1 / 1024.0);
|
||||
parallel_for_(Range(0, total), fn, nstripes);
|
||||
}
|
||||
|
||||
@@ -120,6 +120,7 @@ size_t fastGemmPackBSize(int N, int K);
|
||||
int fastGemmMC();
|
||||
int fastGemmNC();
|
||||
int fastGemmKC();
|
||||
int fastGemmNR();
|
||||
|
||||
void fastGemmPackBKernel(const char *B, char *packed_B, size_t N, size_t K, size_t ldb0, size_t ldb1, size_t esz);
|
||||
|
||||
@@ -140,13 +141,13 @@ void fastGemmBatchKernel(size_t batch, const size_t *A_offsets, const size_t *B_
|
||||
|
||||
void pagedAttnQKGemmKernel(
|
||||
const char *Q, const std::vector<const char *> &K, char *A,
|
||||
size_t B, size_t T_q, size_t Nq, size_t N_k, size_t T_s, size_t D,
|
||||
size_t esz, bool isQ3d
|
||||
size_t B, size_t T_q, size_t Nq, size_t N_k, size_t T_s, size_t D, size_t T_k,
|
||||
float sm_scale, size_t esz, bool isQ3d
|
||||
);
|
||||
|
||||
void pagedAttnAVGemmKernel(
|
||||
const char* A, const std::vector<const char *> &K, char*Out,
|
||||
size_t B, size_t T_a, size_t Nq, size_t N_k, size_t T_s, size_t D,
|
||||
const char* A, const std::vector<const char *> &V, char*Out,
|
||||
size_t B, size_t T_a, size_t Nq, size_t N_k, size_t T_s, size_t D, size_t T_v,
|
||||
size_t esz, bool canonical_layout, size_t packed_stride
|
||||
);
|
||||
|
||||
@@ -529,6 +530,7 @@ static inline void fast_gemm_macro_kernel(int m, int n, int k,
|
||||
for(int p = 0; p < mr; p++)
|
||||
memcpy(cptr + p * (ldc * esz), cptr0 + p * ldc0_esz, nr_esz);
|
||||
}
|
||||
|
||||
#if CV_NEON && CV_NEON_AARCH64
|
||||
fast_gemm8x12_f32(k, packed_A + i * k * esz, packed_B + j * k * esz, cptr, ldc, alpha);
|
||||
#elif CV_AVX
|
||||
@@ -557,6 +559,7 @@ size_t fastGemmPackBSize(int N, int K) {
|
||||
int fastGemmMC() {return FAST_GEMM_F32_MC;}
|
||||
int fastGemmNC() {return FAST_GEMM_F32_NC;}
|
||||
int fastGemmKC() {return FAST_GEMM_F32_PACKED_STRIDE_K;}
|
||||
int fastGemmNR() {return FAST_GEMM_F32_NR;}
|
||||
|
||||
void fastGemmPackBKernel(const char *B, char *packed_B, size_t N, size_t K, size_t ldb0, size_t ldb1, size_t esz) {
|
||||
size_t GEMM_NC = FAST_GEMM_F32_NC, GEMM_NR = FAST_GEMM_F32_NR;
|
||||
@@ -940,9 +943,9 @@ void fastGemmBatchKernel(size_t batch, const size_t *A_offsets, const size_t *B_
|
||||
// tot. seq len T = T_s * S
|
||||
void pagedAttnQKGemmKernel(
|
||||
const char *Q, const std::vector<const char *> &K, char *A,
|
||||
size_t B, size_t T_q, size_t Nq, size_t N_k, size_t T_s, size_t D,
|
||||
size_t esz, bool isQ3d
|
||||
) {
|
||||
size_t B, size_t T_q, size_t Nq, size_t N_k, size_t T_s, size_t D, size_t T_k,
|
||||
float sm_scale, size_t esz, bool isQ3d
|
||||
){
|
||||
size_t GEMM_MC = static_cast<size_t>(FAST_GEMM_F32_MC),
|
||||
GEMM_NC = static_cast<size_t>(FAST_GEMM_F32_NC),
|
||||
GEMM_MR = static_cast<size_t>(FAST_GEMM_F32_MR),
|
||||
@@ -959,11 +962,14 @@ void pagedAttnQKGemmKernel(
|
||||
const size_t tiles_per_mat = m_tiles * n_tiles;
|
||||
const size_t S = K.size();
|
||||
const size_t T = S * T_s;
|
||||
const size_t T_r = T_s - (T - T_k);
|
||||
const int n_kq_groups = Nq / N_k;
|
||||
size_t batch = Nq * B;
|
||||
// batch = B x N_q
|
||||
int total = S * batch * tiles_per_mat;
|
||||
|
||||
size_t packed_stride = fastGemmPackBSize(T_s, D);
|
||||
|
||||
auto fn = [&](const Range &r) {
|
||||
cv::AutoBuffer<char, FAST_GEMM_MAX_STACKBUF> packed_q_buff;
|
||||
packed_q_buff.allocate(buff_size);
|
||||
@@ -988,22 +994,32 @@ void pagedAttnQKGemmKernel(
|
||||
size_t mc = T_q - i0 < MC ? T_q - i0 : MC;
|
||||
size_t nc = T_s - j0 < NC ? T_s - j0 : NC;
|
||||
|
||||
size_t rc = nc;
|
||||
if (s == S - 1) {
|
||||
if (j0 >= T_r) continue; // Completely out of bounds, skip this tile
|
||||
rc = std::min(nc, T_r - j0);
|
||||
}
|
||||
|
||||
size_t q_offset = b * Nq * T_q * D +
|
||||
(isQ3d ? nq * D : nq * T_q * D);
|
||||
const char *q_block = Q + q_offset * esz;
|
||||
|
||||
size_t k_offset = b * N_k * T_s * D +
|
||||
n_k * T_s * D +
|
||||
size_t k_offset = b * N_k * packed_stride +
|
||||
n_k * packed_stride +
|
||||
j0 * D;
|
||||
const char *k_block = (const char *)K[s] + k_offset * esz;
|
||||
|
||||
// save result to A[b, n_q, : , T_s * s : T_s * (s + 1)]
|
||||
const int a_offset = b * Nq * T_q * T +
|
||||
nq * T_q * T +
|
||||
T_s * s +
|
||||
i0 * T + j0;
|
||||
const size_t a_offset = b * Nq * T_q * T_k +
|
||||
nq * T_q * T_k +
|
||||
T_s * s +
|
||||
i0 * T_k + j0;
|
||||
char* a_block = A + a_offset * esz;
|
||||
|
||||
for(size_t i = 0; i < mc; i++) {
|
||||
memset(a_block + i * T_k * esz, 0, rc * esz);
|
||||
}
|
||||
|
||||
int _nc = static_cast<int>((nc + GEMM_NR - 1) / GEMM_NR) * GEMM_NR * esz;
|
||||
for(int k0 = 0; k0 < D; k0 += KC)
|
||||
{
|
||||
@@ -1020,7 +1036,7 @@ void pagedAttnQKGemmKernel(
|
||||
fast_gemm_pack8_f32(mc, kc, q_block + step_q, ldq0, 1, packed_q);
|
||||
#endif
|
||||
// run kernel
|
||||
fast_gemm_macro_kernel(mc, nc, kc, packed_q, k_block, 1.f, a_block, T, esz);
|
||||
fast_gemm_macro_kernel(mc, rc, kc, packed_q, k_block, sm_scale, a_block, T_k, esz);
|
||||
k_block += _nc * kc;
|
||||
}
|
||||
}
|
||||
@@ -1041,7 +1057,7 @@ void pagedAttnQKGemmKernel(
|
||||
// tot. seq len T = T_s * S
|
||||
void pagedAttnAVGemmKernel(
|
||||
const char* A, const std::vector<const char *> &V, char*Out,
|
||||
size_t B, size_t T_a, size_t Nq, size_t N_k, size_t T_s, size_t D,
|
||||
size_t B, size_t T_a, size_t Nq, size_t N_k, size_t T_s, size_t D, size_t T_v,
|
||||
size_t esz, bool canonical_layout, size_t packed_stride
|
||||
) {
|
||||
size_t GEMM_MC = static_cast<size_t>(FAST_GEMM_F32_MC),
|
||||
@@ -1049,12 +1065,9 @@ void pagedAttnAVGemmKernel(
|
||||
GEMM_MR = static_cast<size_t>(FAST_GEMM_F32_MR),
|
||||
GEMM_NR = static_cast<size_t>(FAST_GEMM_F32_NR);
|
||||
|
||||
const size_t S = V.size();
|
||||
const size_t T = S * T_s;
|
||||
|
||||
size_t MC = (((GEMM_MC < T_a ? GEMM_MC : T_a) + GEMM_MR - 1) / GEMM_MR) * GEMM_MR;
|
||||
size_t NC = (((GEMM_NC < D ? GEMM_NC : D) + GEMM_NR - 1) / GEMM_NR) * GEMM_NR;
|
||||
size_t KC = std::min( static_cast<size_t>(FAST_GEMM_F32_PACKED_STRIDE_K), T);
|
||||
size_t KC = std::min( static_cast<size_t>(FAST_GEMM_F32_PACKED_STRIDE_K), T_v);
|
||||
|
||||
size_t buff_size = KC * MC * esz;
|
||||
|
||||
@@ -1070,9 +1083,13 @@ void pagedAttnAVGemmKernel(
|
||||
packed_a_buff.allocate(buff_size);
|
||||
char* packed_a = packed_a_buff.data();
|
||||
|
||||
cv::AutoBuffer<char, FAST_GEMM_MAX_STACKBUF> repacked_v_buff;
|
||||
repacked_v_buff.allocate(NC * T_s * esz);
|
||||
char* repacked_v = repacked_v_buff.data();
|
||||
|
||||
size_t start = r.start;
|
||||
size_t end = r.end;
|
||||
size_t ldc0 = canonical_layout ? D : Nq * D;
|
||||
size_t ldc0 = canonical_layout ? Nq * D : D;
|
||||
|
||||
for (size_t tile_idx = start; tile_idx < end; tile_idx++) {
|
||||
size_t idx = tile_idx / tiles_per_mat;
|
||||
@@ -1089,9 +1106,9 @@ void pagedAttnAVGemmKernel(
|
||||
size_t mc = T_a - i0 < MC ? T_a - i0 : MC;
|
||||
size_t nc = D - j0 < NC ? D - j0 : NC;
|
||||
|
||||
const size_t a_offset = b * Nq * T_a * T +
|
||||
nq * T_a * T +
|
||||
i0 * T;
|
||||
const size_t a_offset = b * Nq * T_a * T_v +
|
||||
nq * T_a * T_v +
|
||||
i0 * T_v;
|
||||
const char*a_block = A + a_offset * esz;
|
||||
|
||||
// start at the 0th row
|
||||
@@ -1102,42 +1119,58 @@ void pagedAttnAVGemmKernel(
|
||||
|
||||
size_t o_offset = b * Nq * T_a * D;
|
||||
if (canonical_layout)
|
||||
o_offset += nq * T_a * D +
|
||||
i0 * D + j0;
|
||||
o_offset += i0 * Nq * D + nq * D + j0;
|
||||
else
|
||||
o_offset += i0 * Nq * D +
|
||||
nq * D + j0;
|
||||
o_offset += nq * T_a * D + i0 * D + j0;
|
||||
// nq * T_a * D + i0 * D + j0;
|
||||
char* out_block = Out + o_offset * esz;
|
||||
|
||||
int _nc = static_cast<int>((nc + GEMM_NR - 1) / GEMM_NR) * GEMM_NR * esz;
|
||||
for(size_t i = 0; i < mc; i++) {
|
||||
memset(out_block + i * ldc0 * esz, 0, nc * esz);
|
||||
}
|
||||
|
||||
for(int k0 = 0; k0 < T; k0 += KC)
|
||||
// int _nc = static_cast<int>((nc + GEMM_NR - 1) / GEMM_NR) * GEMM_NR * esz;
|
||||
|
||||
for(int k0 = 0; k0 < (int)T_v; )
|
||||
{
|
||||
// get the k-block
|
||||
size_t sk = k0 / T_s; // which page
|
||||
size_t k = k0 % T_s;
|
||||
size_t v_offset = v_offset_base * esz + k * _nc;
|
||||
int kc = std::min((int)KC, (int)T_v - k0);
|
||||
kc = std::min(kc, (int)T_s - (int)k);
|
||||
|
||||
size_t v_offset = v_offset_base * esz + k * GEMM_NR * esz;
|
||||
const char *v_block = V[sk] + v_offset;
|
||||
|
||||
// T_s should be divisible by KC
|
||||
// pack
|
||||
#if CV_NEON && CV_NEON_AARCH64
|
||||
fast_gemm_pack8_f32(mc, KC, a_block + k0 * esz, T, 1, packed_a);
|
||||
fast_gemm_pack8_f32(mc, kc, a_block + k0 * esz, T_v, 1, packed_a);
|
||||
#elif CV_AVX
|
||||
fast_gemm_pack12_f32(mc, KC, a_block + k0 * esz, T, 1, packed_a);
|
||||
fast_gemm_pack12_f32(mc, kc, a_block + k0 * esz, T_v, 1, packed_a);
|
||||
#elif CV_LASX
|
||||
fast_gemm_pack12_f32(mc, KC, a_block + k0 * esz, T, 1, packed_a);
|
||||
fast_gemm_pack12_f32(mc, kc, a_block + k0 * esz, T_v, 1, packed_a);
|
||||
#elif CV_SIMD128
|
||||
fast_gemm_pack8_f32(mc, KC, a_block + k0 * esz, T, 1, packed_a);
|
||||
fast_gemm_pack8_f32(mc, kc, a_block + k0 * esz, T_v, 1, packed_a);
|
||||
#endif
|
||||
|
||||
const char* v_block_to_use = v_block;
|
||||
if (kc < (int)T_s) {
|
||||
for (size_t j = 0; j < nc; j += GEMM_NR) {
|
||||
const char *src_panel = v_block + j * T_s * esz;
|
||||
char *dst_panel = repacked_v + j * kc * esz;
|
||||
memcpy(dst_panel, src_panel, kc * GEMM_NR * esz);
|
||||
}
|
||||
v_block_to_use = repacked_v;
|
||||
}
|
||||
|
||||
// run kernel
|
||||
fast_gemm_macro_kernel(mc, nc, KC, packed_a, v_block, 1.f, out_block, ldc0, esz);
|
||||
fast_gemm_macro_kernel(mc, nc, kc, packed_a, v_block_to_use, 1.f, out_block, ldc0, esz);
|
||||
k0 += kc;
|
||||
}
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
int cost_per_thread = static_cast<int>((T / KC) * (MC / GEMM_MR) * (NC / GEMM_NR));
|
||||
int cost_per_thread = static_cast<int>((T_v / KC) * (MC / GEMM_MR) * (NC / GEMM_NR));
|
||||
double nstripes = (size_t)total * cost_per_thread * (1 / 1024.0);
|
||||
parallel_for_(Range(0, total), fn, nstripes);
|
||||
}
|
||||
|
||||
@@ -493,6 +493,25 @@ bool Net::haveArg(const std::string& name) const
|
||||
return impl->haveArg(name);
|
||||
}
|
||||
|
||||
void Net::enableKVCache()
|
||||
{
|
||||
CV_Assert(impl);
|
||||
setKVCacheManager(impl);
|
||||
}
|
||||
|
||||
void Net::disableKVCache()
|
||||
{
|
||||
CV_Assert(impl);
|
||||
impl->kvCacheManager = KVCacheManager();
|
||||
}
|
||||
|
||||
void Net::resetKVCache()
|
||||
{
|
||||
CV_Assert(impl);
|
||||
setKVCacheManager(impl);
|
||||
}
|
||||
|
||||
|
||||
Ptr<Graph> Net::getMainGraph() const
|
||||
{
|
||||
CV_Assert(impl);
|
||||
|
||||
@@ -25,6 +25,8 @@
|
||||
|
||||
#include "legacy_backend.hpp" // wrapMat BlobManager OpenCLBackendWrapper
|
||||
|
||||
#include "kv_cache_manager.hpp"
|
||||
|
||||
#include <unordered_map>
|
||||
|
||||
#ifdef HAVE_ONNXRUNTIME
|
||||
@@ -107,6 +109,8 @@ struct Net::Impl : public detail::NetImplBase
|
||||
bool fusion;
|
||||
bool isAsync; // FIXIT: drop
|
||||
bool useWinograd;
|
||||
bool useKVCache = false;
|
||||
|
||||
std::vector<int64> layersTimings;
|
||||
|
||||
std::string modelFileName;
|
||||
@@ -125,6 +129,7 @@ struct Net::Impl : public detail::NetImplBase
|
||||
std::vector<Mat> buffers;
|
||||
std::vector<Mat> scratchBufs;
|
||||
std::vector<Ptr<Graph> > allgraphs;
|
||||
KVCacheManager kvCacheManager;
|
||||
|
||||
Ptr<Graph> mainGraph;
|
||||
int globGraphIdx;
|
||||
|
||||
@@ -3005,6 +3005,134 @@ TEST(ConvolutionWinograd, Accuracy)
|
||||
normAssert(outLarge, refLarge, "Large input after small", 0.0, 0.0);
|
||||
}
|
||||
|
||||
class TESTKVCache : public testing::TestWithParam<std::string>
|
||||
{
|
||||
public:
|
||||
void testKVCache(const std::string& layout)
|
||||
{
|
||||
auto engine_forced = static_cast<cv::dnn::EngineType>(
|
||||
cv::utils::getConfigurationParameterSizeT("OPENCV_FORCE_DNN_ENGINE", cv::dnn::ENGINE_AUTO));
|
||||
if (engine_forced == cv::dnn::ENGINE_CLASSIC)
|
||||
{
|
||||
// Mark the test as skipped and exit early.
|
||||
applyTestTag(CV_TEST_TAG_DNN_SKIP_PARSER);
|
||||
return;
|
||||
}
|
||||
|
||||
std::string model_path = "dnn/onnx/models/test_attention_kv_cache_" + layout + ".onnx";
|
||||
|
||||
Net netWithKVCache = readNetFromONNX(findDataFile(model_path, true), cv::dnn::ENGINE_NEW);
|
||||
netWithKVCache.enableKVCache();
|
||||
Net netWithoutKVCache = readNetFromONNX(findDataFile(model_path, true), cv::dnn::ENGINE_NEW);
|
||||
|
||||
int T = 523, Nq = 8, Nkv = 4, D = 256;
|
||||
int T_pref = T;
|
||||
|
||||
std::vector<int> q_sz, k_sz, v_sz;
|
||||
if (layout == "3d") {
|
||||
q_sz = {1, T, Nq * D};
|
||||
k_sz = {1, T, Nkv * D};
|
||||
v_sz = {1, T, Nkv * D};
|
||||
} else {
|
||||
q_sz = {1, Nq, T, D};
|
||||
k_sz = {1, Nkv, T, D};
|
||||
v_sz = {1, Nkv, T, D};
|
||||
}
|
||||
|
||||
Mat Q_all(q_sz, CV_32F);
|
||||
Mat K_all(k_sz, CV_32F);
|
||||
Mat V_all(v_sz, CV_32F);
|
||||
|
||||
cv::randn(Q_all, 0.0, 1.0);
|
||||
cv::randn(K_all, 0.0, 1.0);
|
||||
cv::randn(V_all, 0.0, 1.0);
|
||||
|
||||
std::vector<int> mask_sz = {1, Nq, T, T};
|
||||
Mat mask(mask_sz, CV_32S, cv::Scalar(0));
|
||||
|
||||
int* mask_ptr = (int*)mask.data;
|
||||
for (int n = 0; n < Nq; n++) {
|
||||
for (int i = 0; i < T; i++) {
|
||||
for (int j = 0; j < T; j++) {
|
||||
int idx = n * T * T +
|
||||
i * T + j;
|
||||
if (i < T_pref) {
|
||||
if (j < T_pref) mask_ptr[idx] = 1;
|
||||
} else {
|
||||
if (j <= i) mask_ptr[idx] = 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Mat Y;
|
||||
if (layout == "3d") {
|
||||
std::vector<int> sz = {1, T, Nq * D};
|
||||
Y = Mat(sz, CV_32F);
|
||||
} else {
|
||||
std::vector<int> sz = {1, Nq, T, D};
|
||||
Y = Mat(sz, CV_32F);
|
||||
}
|
||||
Y.setTo(0);
|
||||
|
||||
std::vector<Range> ranges_pref;
|
||||
if (layout == "3d") {
|
||||
ranges_pref = {Range::all(), Range(0, T_pref), Range::all()};
|
||||
} else {
|
||||
ranges_pref = {Range::all(), Range::all(), Range(0, T_pref), Range::all()};
|
||||
}
|
||||
|
||||
Mat Q_pref = Q_all(ranges_pref);
|
||||
Mat K_pref = K_all(ranges_pref);
|
||||
Mat V_pref = V_all(ranges_pref);
|
||||
|
||||
// 1. Prefill
|
||||
netWithKVCache.setInput(Q_pref, "Q");
|
||||
netWithKVCache.setInput(K_pref, "K");
|
||||
netWithKVCache.setInput(V_pref, "V");
|
||||
Mat prefillResult = netWithKVCache.forward(); // prefill
|
||||
prefillResult.copyTo(Y(ranges_pref));
|
||||
// 2. Generate
|
||||
for(int t = T_pref; t < T; t++)
|
||||
{
|
||||
std::vector<Range> ranges_gen;
|
||||
if (layout == "3d") {
|
||||
ranges_gen = {Range::all(), Range(t, t + 1), Range::all()};
|
||||
} else {
|
||||
ranges_gen = {Range::all(), Range::all(), Range(t, t + 1), Range::all()};
|
||||
}
|
||||
|
||||
netWithKVCache.setInput(Q_all(ranges_gen), "Q");
|
||||
netWithKVCache.setInput(K_all(ranges_gen), "K");
|
||||
netWithKVCache.setInput(V_all(ranges_gen), "V");
|
||||
|
||||
Mat nextToken = netWithKVCache.forward();
|
||||
nextToken.copyTo(Y(ranges_gen));
|
||||
}
|
||||
|
||||
// 3. Standard path
|
||||
netWithoutKVCache.setInput(Q_all, "Q");
|
||||
netWithoutKVCache.setInput(K_all, "K");
|
||||
netWithoutKVCache.setInput(V_all, "V");
|
||||
netWithoutKVCache.setInput(mask, "Mask");
|
||||
|
||||
Mat Yref = netWithoutKVCache.forward();
|
||||
|
||||
std::string msg = "Attention generate " + layout + ": KV vs standard";
|
||||
normAssert(Y, Yref, msg.c_str(), 1e-3, 1e-3);
|
||||
}
|
||||
};
|
||||
|
||||
TEST_P(TESTKVCache, layouts)
|
||||
{
|
||||
testKVCache(GetParam());
|
||||
}
|
||||
|
||||
INSTANTIATE_TEST_CASE_P(KV_Cache, TESTKVCache, testing::Values("3d", "4d"));
|
||||
|
||||
|
||||
|
||||
TEST(Layer_Test_GeluApprox, NoNaN_LargeInput)
|
||||
{
|
||||
LayerParams lp;
|
||||
@@ -3060,4 +3188,4 @@ TEST(Layer_Test_Softmax, NoNaN_AllNegInf)
|
||||
}
|
||||
}
|
||||
|
||||
}} // namespace
|
||||
}} // namespace
|
||||
Reference in New Issue
Block a user