diff --git a/modules/dnn/src/kv_cache_manager.cpp b/modules/dnn/src/kv_cache_manager.cpp index e89e10726f..0b42d3d113 100644 --- a/modules/dnn/src/kv_cache_manager.cpp +++ b/modules/dnn/src/kv_cache_manager.cpp @@ -48,6 +48,7 @@ void setKVCacheManager(Ptr netimpl) manager.netimpl = netimpl; manager.opt.init(); initKVDataRecursively(netimpl->mainGraph, manager.kData, manager.vData, manager.opt); + manager.buildRoutes(); manager.isInitialized = true; netimpl->useKVCache = true; @@ -61,6 +62,78 @@ void KVCacheManager::init() CV_Assert(isInitialized); } +void KVCacheManager::buildRoutes() +{ + presentToPastRoutes.clear(); + hasRoutes = false; + + if (!netimpl || !netimpl->mainGraph) + return; + + const std::vector& gr_outputs = netimpl->mainGraph->outputs(); + for (const Arg& out_arg : gr_outputs) + { + const ArgData& out_adata = netimpl->args.at(out_arg.idx); + const std::string& out_name = out_adata.name; + + if (out_name.compare(0, 8, "present.") != 0) + continue; + + std::string past_name = "past_key_values." + out_name.substr(8); + auto it = netimpl->argnames.find(past_name); + if (it == netimpl->argnames.end()) + continue; + + Arg past_arg((int)it->second); + if (netimpl->args.at(past_arg.idx).kind != DNN_ARG_INPUT) + continue; + + presentToPastRoutes.emplace_back(out_arg.idx, past_arg.idx); + } + + hasRoutes = !presentToPastRoutes.empty(); + if (hasRoutes) + initPastTensors(); +} + +void KVCacheManager::initPastTensors() +{ + for (const auto& route : presentToPastRoutes) + { + const ArgData& past_adata = netimpl->args.at(route.second); + const MatShape& decl_shape = past_adata.shape; + if (decl_shape.dims <= 0) + continue; + + // Replace symbolic dims (stored as <=0): batch (dim 0) -> 1, all others -> 0 for empty-sequence state. + std::vector shape_vec(decl_shape.dims); + for (int d = 0; d < decl_shape.dims; d++) + shape_vec[d] = (decl_shape[d] > 0) ? decl_shape[d] : (d == 0 ? 1 : 0); + + int dtype = past_adata.type; + if (dtype < 0) + dtype = CV_32F; + + Mat& past_t = netimpl->__tensors__.at(route.second); + past_t = Mat(shape_vec, dtype, Scalar(0)); + netimpl->finalizeLayers = true; + } +} + +void KVCacheManager::applyRoutes() +{ + for (const auto& route : presentToPastRoutes) + { + const Mat& present_t = netimpl->argTensor(Arg(route.first)); + Mat& past_t = netimpl->__tensors__.at(route.second); + if (present_t.empty()) + continue; + if (past_t.shape() != present_t.shape() || past_t.type() != present_t.type()) + netimpl->finalizeLayers = true; + present_t.copyTo(past_t); + } +} + void KVCache::grow(const Mat& newData) { CV_Assert(newData.dims == 4 || newData.dims == 3); @@ -191,21 +264,16 @@ void KCache::growGenerate(const Mat& newData){ 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); - } + const int nstripes = batch_size * nHeads; + parallel_for_(Range(0, nstripes), [&](const Range& r) { + for (int i = r.start; i < r.end; i++) { + int b = i / nHeads, h = i % nHeads; + const float* src = data + (b * nHeads + h) * headDim; + float* dst = page + b * nHeads * Ps + h * Ps + t0; + for (int j = 0; j < headDim; j++) + dst[j * pageSize] = src[j]; } - } + }); nTokens += 1; } @@ -225,8 +293,10 @@ void VCache::growGenerate(const Mat& newData){ 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++){ + const int nstripes = batch_size * nHeads; + parallel_for_(Range(0, nstripes), [&](const Range& r) { + for (int i = r.start; i < r.end; i++) { + int b = i / nHeads, h = i % nHeads; 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); @@ -242,11 +312,10 @@ void VCache::growGenerate(const Mat& newData){ } } } - } + }); 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 index 89b2961499..274c435196 100644 --- a/modules/dnn/src/kv_cache_manager.hpp +++ b/modules/dnn/src/kv_cache_manager.hpp @@ -89,7 +89,14 @@ struct KVCacheManager FastGemmOpt opt; bool isInitialized = false; + // present.* output arg idx -> past_key_values.* input arg idx + std::vector> presentToPastRoutes; + bool hasRoutes = false; + void init(); + void buildRoutes(); + void applyRoutes(); + void initPastTensors(); }; void setKVCacheManager(Ptr netimpl); diff --git a/modules/dnn/src/layers/gemm_layer.cpp b/modules/dnn/src/layers/gemm_layer.cpp index ceac819619..de6761f63d 100644 --- a/modules/dnn/src/layers/gemm_layer.cpp +++ b/modules/dnn/src/layers/gemm_layer.cpp @@ -255,7 +255,7 @@ public: LayerGemmOpMode mode = getOpMode(inputs.size(), blobs.size()); // pack B if it is const - if (constB(mode)) { + if (constB(mode) && blobs[0].data != last_packed_blob_data) { fastGemmPackB(blobs[0], packed_B, trans_b, opt); // Pre-pack B in the "thin" layout when the gemm shape has a @@ -304,6 +304,7 @@ public: } } #endif + last_packed_blob_data = blobs[0].data; } if (constC(mode) && flatten_a) { @@ -578,6 +579,7 @@ private: std::vector broadcast_C; int real_ndims_C; FastGemmOpt opt; + const uchar* last_packed_blob_data = nullptr; }; Ptr GemmLayer::create(const LayerParams& params) { diff --git a/modules/dnn/src/layers/matmul_layer.cpp b/modules/dnn/src/layers/matmul_layer.cpp index aea6901167..12ceca96be 100644 --- a/modules/dnn/src/layers/matmul_layer.cpp +++ b/modules/dnn/src/layers/matmul_layer.cpp @@ -154,20 +154,24 @@ class MatMulLayerImpl CV_FINAL : public MatMulLayer { C_shape = shape(outputs[0]); helper.compute(trans_a, trans_b, A_shape, B_shape, C_shape); - if (!blobs.empty()) { - fastGemmPackB(blobs[0], packed_input_B, trans_b, opt); + // Pack only 2D weight matrices; skip higher-dim tensors (e.g. Q@K^T in attention). + const Mat* B_mat = !blobs.empty() ? &blobs[0] : + (inputs.size() >= 2 && inputs[1].dims == 2 ? &inputs[1] : nullptr); + if (B_mat && B_mat->data != last_packed_input_B_data) { + fastGemmPackB(*B_mat, packed_input_B, trans_b, opt); helper.updatePackedBOffsets(packed_input_B.size()); - if (helper.batch == 1 && blobs[0].type() == CV_32F && + if (helper.batch == 1 && B_mat->type() == CV_32F && fastGemmThinEligible(helper.M, helper.N, helper.K)) { thin_packed_B.resize(fastGemmThinPackBSize(helper.N, helper.K)); fastGemmThinPackB(helper.N, helper.K, - blobs[0].ptr(), + B_mat->ptr(), (size_t)helper.ldb0, (size_t)helper.ldb1, thin_packed_B.data()); } else { thin_packed_B.clear(); } + last_packed_input_B_data = B_mat->data; } // broadcast bias if needed @@ -289,10 +293,17 @@ class MatMulLayerImpl CV_FINAL : public MatMulLayer { a, helper.lda0, helper.lda1, thin_packed_B.data(), beta, y, helper.ldc, opt.multi_thread); - } else { + } else if (!packed_input_B.empty()) { fastGemmBatch(helper.batch, helper.A_offsets.data(), helper.packed_B_offsets.data(), helper.C_offsets.data(), helper.M, helper.N, helper.K, alpha, a, helper.lda0, helper.lda1, packed_input_B.data(), beta, y, helper.ldc, opt); + } else { + // truly dynamic B (changes every call — no packing cache available) + const auto &B = inputs[1]; + const auto *b = B.ptr(); + fastGemmBatch(helper.batch, helper.A_offsets.data(), helper.B_offsets.data(), helper.C_offsets.data(), + helper.M, helper.N, helper.K, alpha, a, helper.lda0, helper.lda1, + b, helper.ldb0, helper.ldb1, beta, y, helper.ldc, opt); } } @@ -523,6 +534,8 @@ class MatMulLayerImpl CV_FINAL : public MatMulLayer { std::vector thin_packed_B; Mat broadcast_bias; + const uchar* last_packed_input_B_data = nullptr; + FastGemmOpt opt; MatMulHelper helper; }; diff --git a/modules/dnn/src/net_impl2.cpp b/modules/dnn/src/net_impl2.cpp index bc09c4507c..4bd024f2d9 100644 --- a/modules/dnn/src/net_impl2.cpp +++ b/modules/dnn/src/net_impl2.cpp @@ -684,6 +684,10 @@ void Net::Impl::forwardMainGraph(InputArrayOfArrays inputs, OutputArrayOfArrays // [TODO] if a target or backend change or there are some other important // global changes in configuration, finalizeLayers should be set to 'true' again finalizeLayers = false; + + // Feed present.* outputs back as past_key_values.* inputs for the next step (causal-lm-with-past). + if (useKVCache && kvCacheManager.hasRoutes) + kvCacheManager.applyRoutes(); } void Net::Impl::forwardWithSingleOutput(const std::string& outname, OutputArrayOfArrays outputBlobs) diff --git a/modules/dnn/test/test_layers.cpp b/modules/dnn/test/test_layers.cpp index 422afb6954..94af725bce 100644 --- a/modules/dnn/test/test_layers.cpp +++ b/modules/dnn/test/test_layers.cpp @@ -3188,4 +3188,4 @@ TEST(Layer_Test_Softmax, NoNaN_AllNegInf) } } -}} // namespace \ No newline at end of file +}} // namespace diff --git a/samples/dnn/gemma3_inference.py b/samples/dnn/gemma3_inference.py index c46a7cc30b..362768f03a 100644 --- a/samples/dnn/gemma3_inference.py +++ b/samples/dnn/gemma3_inference.py @@ -1,7 +1,7 @@ # 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) 2025, BigVision LLC, all rights reserved. +# Copyright (C) 2026, BigVision LLC, all rights reserved. # Third party copyrights are property of their respective owners. ''' @@ -15,11 +15,17 @@ Exporting Gemma3 model to ONNX: 1. Install the required dependencies: - pip install optimum[exporters] torch transformers + pip install optimum[exporters] optimum-onnx[onnxruntime] torch transformers 2. Export the model to ONNX: - optimum-cli export onnx --model google/gemma-3-1b-it --task causal-lm gemma3_instruct_onnx/ + Without KV-cache: + + optimum-cli export onnx --model google/gemma-3-1b-it --task causal-lm gemma3_instruct_onnx/ + + With KV-cache (recommended, faster autoregressive inference): + + optimum-cli export onnx --model google/gemma-3-1b-it --task causal-lm-with-past gemma3_instruct_onnx_with_past/ Run the script: @@ -29,9 +35,18 @@ Run the script: 2. Run the script: - python gemma3_inference.py --model= \ - --tokenizer_path= \ - --prompt="What is OpenCV?" + Without KV-cache (causal-lm export): + + python gemma3_inference.py --model= \ + --tokenizer_path= \ + --prompt="What is OpenCV?" + + With KV-cache (causal-lm-with-past export): + + python gemma3_inference.py --model= \ + --tokenizer_path= \ + --prompt="What is OpenCV?" \ + --use_kv_cache The tokenizer_path should point to an OpenCV-format config.json (e.g., from opencv_extra/testdata/dnn/llm/gemma3/config.json), NOT the HuggingFace tokenizer_config.json. @@ -48,6 +63,7 @@ def parse_args(): parser.add_argument('--tokenizer_path', type=str, required=True, help='Path to Gemma3 tokenizer config.json.') parser.add_argument('--prompt', type=str, default='What is OpenCV?', help='User prompt.') parser.add_argument('--max_new_tokens', type=int, default=64, help='Maximum number of new tokens to generate.') + parser.add_argument('--use_kv_cache', action='store_true', default=False, help='Enable KV-cache for faster inference (requires causal-lm-with-past export).') parser.add_argument('--seed', type=int, default=0, help='Random seed.') return parser.parse_args() @@ -55,36 +71,55 @@ def build_gemma3_prompt(user_prompt): '''Wrap user prompt in Gemma3 chat format.''' return 'user\n' + user_prompt + '\nmodel\n' -def gemma3_inference(net, prompt, max_new_tokens, tokenizer): +def gemma3_inference(net, prompt, max_new_tokens, tokenizer, use_kv_cache=True): print("Inferencing Gemma3 model...") tokens = tokenizer.encode(prompt) # Prepend BOS token (id=2) as required by Gemma3 tokens = [2] + list(tokens) - tokens = np.array(tokens, dtype=np.int64).reshape(1, -1) + input_ids = np.array(tokens, dtype=np.int64).reshape(1, -1) # Gemma3 special token IDs eos_id = 1 # eot_id = 106 # stop_ids = (eos_id, eot_id) - for _ in range(max_new_tokens): - seq_len = tokens.shape[1] - attention_mask = np.ones((1, seq_len), dtype=np.int64) + generated = [] - net.setInput(tokens, 'input_ids') - net.setInput(attention_mask, 'attention_mask') - logits = net.forward() # (1, seq_len, vocab_size) - logits = logits[:, -1, :] # take last token logits + if use_kv_cache: + net.enableKVCache() + prompt_len = input_ids.shape[1] - new_id = int(np.argmax(logits.reshape(-1))) - tokens = np.concatenate((tokens, np.array([[new_id]], dtype=np.int64)), axis=1) + # Prefill: process full prompt once to populate KV-cache + net.setInput(input_ids, 'input_ids') + net.setInput(np.ones((1, prompt_len), dtype=np.int64), 'attention_mask') + logits = net.forward() + new_id = int(np.argmax(logits[:, -1, :].reshape(-1))) + generated = [new_id] - if new_id in stop_ids: - break + # Generate: feed one new token per step; OpenCV routes present.* -> past_key_values.* + for _ in range(max_new_tokens - 1): + if new_id in stop_ids: + break + net.setInput(np.array([[new_id]], dtype=np.int64), 'input_ids') + net.setInput(np.ones((1, prompt_len + len(generated)), dtype=np.int64), 'attention_mask') + logits = net.forward() + new_id = int(np.argmax(logits[:, -1, :].reshape(-1))) + generated.append(new_id) + else: + # Without KV-cache: feed full growing sequence each step + for _ in range(max_new_tokens): + net.setInput(input_ids, 'input_ids') + net.setInput(np.ones((1, input_ids.shape[1]), dtype=np.int64), 'attention_mask') + logits = net.forward() + new_id = int(np.argmax(logits[:, -1, :].reshape(-1))) + if new_id in stop_ids: + break + generated.append(new_id) + input_ids = np.concatenate([input_ids, [[new_id]]], axis=1) - return tokens + return np.array([tokens + generated], dtype=np.int64) if __name__ == '__main__': @@ -100,6 +135,6 @@ if __name__ == '__main__': print(f"Prompt:\n{gemma3_prompt}") prompt_len = len(tokenizer.encode(gemma3_prompt)) + 1 # +1 for BOS token - tokens = gemma3_inference(net, gemma3_prompt, args.max_new_tokens, tokenizer) + tokens = gemma3_inference(net, gemma3_prompt, args.max_new_tokens, tokenizer, args.use_kv_cache) response = tokenizer.decode(tokens[0][prompt_len:].tolist()) print(f"Response:\n{response}") diff --git a/samples/dnn/qwen_inference.py b/samples/dnn/qwen_inference.py index 6fc33ca9f3..d36e923d4d 100644 --- a/samples/dnn/qwen_inference.py +++ b/samples/dnn/qwen_inference.py @@ -9,11 +9,17 @@ Exporting Qwen2.5 model to ONNX: 1. Install the required dependencies: - pip install optimum[exporters] torch transformers + pip install optimum[exporters] optimum-onnx[onnxruntime] torch transformers 2. Export the model to ONNX: - optimum-cli export onnx --model Qwen/Qwen2.5-0.5B-Instruct --task causal-lm qwen2.5_instruct_onnx/ + Without KV-cache: + + optimum-cli export onnx --model Qwen/Qwen2.5-0.5B-Instruct --task causal-lm qwen2.5_instruct_onnx/ + + With KV-cache (recommended, faster autoregressive inference): + + optimum-cli export onnx --model Qwen/Qwen2.5-0.5B-Instruct --task causal-lm-with-past qwen2.5_instruct_onnx_with_past/ Run the script: @@ -23,9 +29,18 @@ Run the script: 2. Run the script: - python qwen_inference.py --model= \ - --tokenizer_path= \ - --prompt="What is OpenCV?" + Without KV-cache (causal-lm export): + + python qwen_inference.py --model= \ + --tokenizer_path= \ + --prompt="What is OpenCV?" + + With KV-cache (causal-lm-with-past export): + + python qwen_inference.py --model= \ + --tokenizer_path= \ + --prompt="What is OpenCV?" \ + --use_kv_cache ''' import numpy as np @@ -39,47 +54,66 @@ def parse_args(): parser.add_argument('--tokenizer_path', type=str, required=True, help='Path to Qwen2.5 tokenizer config.json.') parser.add_argument('--prompt', type=str, default='What is OpenCV?', help='User prompt.') parser.add_argument('--max_new_tokens', type=int, default=64, help='Maximum number of new tokens to generate.') + parser.add_argument('--use_kv_cache', action='store_true', default=False, help='Enable KV-cache for faster inference (requires causal-lm-with-past export).') parser.add_argument('--seed', type=int, default=0, help='Random seed.') return parser.parse_args() -def stable_softmax(logits): - exp_logits = np.exp(logits - np.max(logits, axis=-1, keepdims=True)) - return exp_logits / np.sum(exp_logits, axis=-1, keepdims=True) - def build_chatml_prompt(user_prompt): '''Wrap user prompt in Qwen2.5 ChatML format.''' return '<|im_start|>user\n' + user_prompt + '<|im_end|>\n<|im_start|>assistant\n' -def qwen_inference(net, prompt, max_new_tokens, tokenizer): +def qwen_inference(net, prompt, max_new_tokens, tokenizer, use_kv_cache=True): print("Inferencing Qwen2.5 model...") - tokens = tokenizer.encode(prompt) - tokens = np.array(tokens, dtype=np.int64).reshape(1, -1) + tokens = list(tokenizer.encode(prompt)) + input_ids = np.array(tokens, dtype=np.int64).reshape(1, -1) # Qwen2.5 special token IDs im_end_id = 151645 # <|im_end|> eos_id = 151643 # <|endoftext|> stop_ids = (im_end_id, eos_id) - for _ in range(max_new_tokens): - seq_len = tokens.shape[1] - attention_mask = np.ones((1, seq_len), dtype=np.int64) - position_ids = np.arange(seq_len, dtype=np.int64).reshape(1, -1) + generated = [] - net.setInput(tokens, 'input_ids') - net.setInput(attention_mask, 'attention_mask') - net.setInput(position_ids, 'position_ids') - logits = net.forward() # (1, seq_len, vocab_size) - logits = logits[:, -1, :] # take last token logits + if use_kv_cache: + net.enableKVCache() + prompt_len = input_ids.shape[1] - new_id = int(np.argmax(logits.reshape(-1))) - tokens = np.concatenate((tokens, np.array([[new_id]], dtype=np.int64)), axis=1) + # Prefill: process full prompt once to populate KV-cache + net.setInput(input_ids, 'input_ids') + net.setInput(np.ones((1, prompt_len), dtype=np.int64), 'attention_mask') + net.setInput(np.arange(prompt_len, dtype=np.int64).reshape(1, -1), 'position_ids') + logits = net.forward() + new_id = int(np.argmax(logits[:, -1, :].reshape(-1))) + generated = [new_id] - if new_id in stop_ids: - break + # Generate: feed one new token per step; OpenCV routes present.* -> past_key_values.* + for _ in range(max_new_tokens - 1): + if new_id in stop_ids: + break + cur_len = prompt_len + len(generated) + net.setInput(np.array([[new_id]], dtype=np.int64), 'input_ids') + net.setInput(np.ones((1, cur_len), dtype=np.int64), 'attention_mask') + net.setInput(np.array([[cur_len - 1]], dtype=np.int64), 'position_ids') + logits = net.forward() + new_id = int(np.argmax(logits[:, -1, :].reshape(-1))) + generated.append(new_id) + else: + # Without KV-cache: feed full growing sequence each step + for _ in range(max_new_tokens): + seq_len = input_ids.shape[1] + net.setInput(input_ids, 'input_ids') + net.setInput(np.ones((1, seq_len), dtype=np.int64), 'attention_mask') + net.setInput(np.arange(seq_len, dtype=np.int64).reshape(1, -1), 'position_ids') + logits = net.forward() + new_id = int(np.argmax(logits[:, -1, :].reshape(-1))) + if new_id in stop_ids: + break + generated.append(new_id) + input_ids = np.concatenate([input_ids, [[new_id]]], axis=1) - return tokens + return np.array([tokens + generated], dtype=np.int64) if __name__ == '__main__': @@ -94,6 +128,7 @@ if __name__ == '__main__': chatml_prompt = build_chatml_prompt(args.prompt) print(f"Prompt:\n{chatml_prompt}") - tokens = qwen_inference(net, chatml_prompt, args.max_new_tokens, tokenizer) - response = tokenizer.decode(tokens[0].tolist()) + prompt_len = len(tokenizer.encode(chatml_prompt)) + tokens = qwen_inference(net, chatml_prompt, args.max_new_tokens, tokenizer, args.use_kv_cache) + response = tokenizer.decode(tokens[0][prompt_len:].tolist()) print(f"Response:\n{response}")