From c2594b41bfed77d8f6cf216eac187c44f2c4db2e Mon Sep 17 00:00:00 2001 From: nklskyoy <31723634+nklskyoy@users.noreply.github.com> Date: Mon, 25 May 2026 12:39:02 +0200 Subject: [PATCH] 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 --- .../dnn/include/opencv2/dnn/all_layers.hpp | 2 + modules/dnn/include/opencv2/dnn/dnn.hpp | 9 + modules/dnn/src/kv_cache_manager.cpp | 252 ++++++++++++++++++ modules/dnn/src/kv_cache_manager.hpp | 101 +++++++ .../dnn/src/layers/attention_onnxai_layer.cpp | 215 ++++++++++----- .../dnn/src/layers/cpu_kernels/fast_gemm.cpp | 68 +++-- .../dnn/src/layers/cpu_kernels/fast_gemm.hpp | 7 +- .../cpu_kernels/fast_gemm_kernels.default.hpp | 112 +++++--- .../cpu_kernels/fast_gemm_kernels.simd.hpp | 109 +++++--- modules/dnn/src/net.cpp | 19 ++ modules/dnn/src/net_impl.hpp | 5 + modules/dnn/test/test_layers.cpp | 130 ++++++++- 12 files changed, 855 insertions(+), 174 deletions(-) create mode 100644 modules/dnn/src/kv_cache_manager.cpp create mode 100644 modules/dnn/src/kv_cache_manager.hpp diff --git a/modules/dnn/include/opencv2/dnn/all_layers.hpp b/modules/dnn/include/opencv2/dnn/all_layers.hpp index 905d36f047..34f6c10e01 100644 --- a/modules/dnn/include/opencv2/dnn/all_layers.hpp +++ b/modules/dnn/include/opencv2/dnn/all_layers.hpp @@ -1876,6 +1876,8 @@ CV__DNN_INLINE_NS_BEGIN class CV_EXPORTS AttentionOnnxAiLayer : public Layer { public: + int kv_num_heads; + static Ptr create(const LayerParams ¶ms); }; diff --git a/modules/dnn/include/opencv2/dnn/dnn.hpp b/modules/dnn/include/opencv2/dnn/dnn.hpp index 83b1a7ae0a..37fe53f7e2 100644 --- a/modules/dnn/include/opencv2/dnn/dnn.hpp +++ b/modules/dnn/include/opencv2/dnn/dnn.hpp @@ -1027,6 +1027,14 @@ CV__DNN_INLINE_NS_BEGIN */ CV_WRAP int64 getPerfProfile(CV_OUT std::vector& 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(); } diff --git a/modules/dnn/src/kv_cache_manager.cpp b/modules/dnn/src/kv_cache_manager.cpp new file mode 100644 index 0000000000..e89e10726f --- /dev/null +++ b/modules/dnn/src/kv_cache_manager.cpp @@ -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 +#include +#include "layers/cpu_kernels/fast_gemm.hpp" +namespace cv { namespace dnn { +CV__DNN_INLINE_NS_BEGIN + + + +void initKVDataRecursively(const Ptr& graph, std::map& kData, std::map& vData, FastGemmOpt& opt) { + for (const auto& layer : graph->prog()) { + + for (const auto& subgraph : layer->subgraphs() ? *layer->subgraphs() : std::vector>()) { + initKVDataRecursively(subgraph, kData, vData, opt); + } + + if (layer->type == "AttentionOnnxAi") { + int kvNumHeads = layer.dynamicCast()->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 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> 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() + 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& 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() + 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(); + const auto* data = newData.ptr(); + + 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(); + const auto* data = newData.ptr(); + + 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 +}} diff --git a/modules/dnn/src/kv_cache_manager.hpp b/modules/dnn/src/kv_cache_manager.hpp new file mode 100644 index 0000000000..89b2961499 --- /dev/null +++ b/modules/dnn/src/kv_cache_manager.hpp @@ -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 +#include + +#include +#include +#include + +#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& 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 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 kData; + std::map vData; + FastGemmOpt opt; + bool isInitialized = false; + + void init(); +}; + +void setKVCacheManager(Ptr netimpl); + + +CV__DNN_INLINE_NS_END +}} + +#endif diff --git a/modules/dnn/src/layers/attention_onnxai_layer.cpp b/modules/dnn/src/layers/attention_onnxai_layer.cpp index 209e70242e..aac78c5f97 100644 --- a/modules/dnn/src/layers/attention_onnxai_layer.cpp +++ b/modules/dnn/src/layers/attention_onnxai_layer.cpp @@ -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(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 auto* K = inputs[1].ptr(); - const auto* V = inputs[2].ptr(); + std::vector _q_offsets, _k_offsets, _v_offsets, + _a_offsets, _o_offsets; - scale = is_scale_set ? scale : 1.0f / std::sqrt(static_cast(qk_head_size)); + if (with_kv_cache){ + KVCacheManager& kvCacheManager = netimpl->kvCacheManager; - std::vector _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& 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(), seq_len_kv, - opt - ); + scale = is_scale_set ? scale : 1.0f / std::sqrt(static_cast(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 auto* K = inputs[1].ptr(); + + 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(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(), 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(), seq_len_kv, 1, - V, ldv0, 1, - 0.f, - outputs[0].ptr(), 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 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(), seq_len_kv, 1, + V, ldv0, 1, + 0.f, + outputs[0].ptr(), ldout, + opt + ); + } } private: bool is_causal; - int kv_num_heads; int q_num_heads; int qk_matmul_output_mode; float scale; diff --git a/modules/dnn/src/layers/cpu_kernels/fast_gemm.cpp b/modules/dnn/src/layers/cpu_kernels/fast_gemm.cpp index b7b3232052..6516a69aac 100644 --- a/modules/dnn/src/layers/cpu_kernels/fast_gemm.cpp +++ b/modules/dnn/src/layers/cpu_kernels/fast_gemm.cpp @@ -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 &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 packed_K; @@ -815,8 +843,8 @@ void pagedAttnQKGemm( if (opt.use_neon) opt_NEON::pagedAttnQKGemmKernel( Q.ptr(), 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(), 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(), 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(), 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(), 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 &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(), packed_V, Out.ptr(), - 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(), packed_V, Out.ptr(), - 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(), packed_V, Out.ptr(), - 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(), packed_V, Out.ptr(), - 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(), packed_V, Out.ptr(), - 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) ); diff --git a/modules/dnn/src/layers/cpu_kernels/fast_gemm.hpp b/modules/dnn/src/layers/cpu_kernels/fast_gemm.hpp index b0f9740074..cbe8b819f0 100644 --- a/modules/dnn/src/layers/cpu_kernels/fast_gemm.hpp +++ b/modules/dnn/src/layers/cpu_kernels/fast_gemm.hpp @@ -159,6 +159,7 @@ void fastGemmPackB(const Mat &m, std::vector &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 &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 &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 ); diff --git a/modules/dnn/src/layers/cpu_kernels/fast_gemm_kernels.default.hpp b/modules/dnn/src/layers/cpu_kernels/fast_gemm_kernels.default.hpp index 5ef81f1d50..38c2f66594 100644 --- a/modules/dnn/src/layers/cpu_kernels/fast_gemm_kernels.default.hpp +++ b/modules/dnn/src/layers/cpu_kernels/fast_gemm_kernels.default.hpp @@ -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 &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 &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 &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(FAST_GEMM_F32_MC), GEMM_NC = static_cast(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 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((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((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((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 &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(FAST_GEMM_F32_MC), @@ -598,12 +615,9 @@ void pagedAttnAVGemmKernel( GEMM_MR = static_cast(FAST_GEMM_F32_MR), GEMM_NR = static_cast(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(FAST_GEMM_F32_PACKED_STRIDE_K), T); + size_t KC = std::min( static_cast(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 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((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((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((T / KC) * (MC / GEMM_MR) * (NC / GEMM_NR)); + int cost_per_thread = static_cast((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); } diff --git a/modules/dnn/src/layers/cpu_kernels/fast_gemm_kernels.simd.hpp b/modules/dnn/src/layers/cpu_kernels/fast_gemm_kernels.simd.hpp index 1f9eac249f..47f91f2071 100644 --- a/modules/dnn/src/layers/cpu_kernels/fast_gemm_kernels.simd.hpp +++ b/modules/dnn/src/layers/cpu_kernels/fast_gemm_kernels.simd.hpp @@ -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 &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 &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 &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 &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(FAST_GEMM_F32_MC), GEMM_NC = static_cast(FAST_GEMM_F32_NC), GEMM_MR = static_cast(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 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((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 &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(FAST_GEMM_F32_MC), @@ -1049,12 +1065,9 @@ void pagedAttnAVGemmKernel( GEMM_MR = static_cast(FAST_GEMM_F32_MR), GEMM_NR = static_cast(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(FAST_GEMM_F32_PACKED_STRIDE_K), T); + size_t KC = std::min( static_cast(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 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((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((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((T / KC) * (MC / GEMM_MR) * (NC / GEMM_NR)); + int cost_per_thread = static_cast((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); } diff --git a/modules/dnn/src/net.cpp b/modules/dnn/src/net.cpp index 42bfed7802..4aedd100ff 100644 --- a/modules/dnn/src/net.cpp +++ b/modules/dnn/src/net.cpp @@ -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 Net::getMainGraph() const { CV_Assert(impl); diff --git a/modules/dnn/src/net_impl.hpp b/modules/dnn/src/net_impl.hpp index b6cff99b29..eb4df4c5b8 100644 --- a/modules/dnn/src/net_impl.hpp +++ b/modules/dnn/src/net_impl.hpp @@ -25,6 +25,8 @@ #include "legacy_backend.hpp" // wrapMat BlobManager OpenCLBackendWrapper +#include "kv_cache_manager.hpp" + #include #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 layersTimings; std::string modelFileName; @@ -125,6 +129,7 @@ struct Net::Impl : public detail::NetImplBase std::vector buffers; std::vector scratchBufs; std::vector > allgraphs; + KVCacheManager kvCacheManager; Ptr mainGraph; int globGraphIdx; diff --git a/modules/dnn/test/test_layers.cpp b/modules/dnn/test/test_layers.cpp index f0322138e1..422afb6954 100644 --- a/modules/dnn/test/test_layers.cpp +++ b/modules/dnn/test/test_layers.cpp @@ -3005,6 +3005,134 @@ TEST(ConvolutionWinograd, Accuracy) normAssert(outLarge, refLarge, "Large input after small", 0.0, 0.0); } +class TESTKVCache : public testing::TestWithParam +{ +public: + void testKVCache(const std::string& layout) + { + auto engine_forced = static_cast( + 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 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 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 sz = {1, T, Nq * D}; + Y = Mat(sz, CV_32F); + } else { + std::vector sz = {1, Nq, T, D}; + Y = Mat(sz, CV_32F); + } + Y.setTo(0); + + std::vector 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 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 \ No newline at end of file