1
0
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:
nklskyoy
2026-05-25 12:39:02 +02:00
committed by GitHub
parent 07e0dd2fbd
commit c2594b41bf
12 changed files with 855 additions and 174 deletions
@@ -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 &params);
};
+9
View File
@@ -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(); }
+252
View File
@@ -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
}}
+101
View File
@@ -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
+143 -72
View File
@@ -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);
}
+19
View File
@@ -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);
+5
View File
@@ -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;
+129 -1
View File
@@ -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