From eb36f9b4d972ea6f3e3ba389c342e2742a56b11e Mon Sep 17 00:00:00 2001 From: wooway777 Date: Tue, 28 Jul 2026 14:08:00 +0800 Subject: [PATCH] fix: restore ep graph --- csrc/cache/kv_cache.cpp | 11 +++++++++-- csrc/cache/kv_cache.hpp | 5 ++++- csrc/engine/compiler/paged_compiler.cpp | 16 ++++++++++++++++ csrc/pybind11/cache/cache.hpp | 5 +++-- python/infinilm/cache/cache.py | 8 ++++++-- python/infinilm/llm/model_runner/model_runner.py | 4 +++- 6 files changed, 41 insertions(+), 8 deletions(-) diff --git a/csrc/cache/kv_cache.cpp b/csrc/cache/kv_cache.cpp index 88d7367e9..2a7780090 100644 --- a/csrc/cache/kv_cache.cpp +++ b/csrc/cache/kv_cache.cpp @@ -79,9 +79,11 @@ infinicore::Tensor create_layer_kv_cache( // ========================== PagedKVCacheConfig::PagedKVCacheConfig( size_t num_blocks, - size_t block_size) + size_t block_size, + size_t max_batch_size) : num_blocks_(num_blocks), - block_size_(block_size) { + block_size_(block_size), + max_batch_size_(max_batch_size) { } std::unique_ptr @@ -99,6 +101,11 @@ PagedKVCacheConfig::block_size() const { return block_size_; } +size_t +PagedKVCacheConfig::max_batch_size() const { + return max_batch_size_; +} + namespace PagedKVCache { // ========================== // PagedKVCache diff --git a/csrc/cache/kv_cache.hpp b/csrc/cache/kv_cache.hpp index 4d0a9a704..c6ad7ebb9 100644 --- a/csrc/cache/kv_cache.hpp +++ b/csrc/cache/kv_cache.hpp @@ -39,15 +39,18 @@ class PagedKVCacheConfig final : public CacheConfig { public: PagedKVCacheConfig( size_t num_blocks, - size_t block_size = 256); + size_t block_size = 256, + size_t max_batch_size = 512); std::unique_ptr unique_copy() const override; size_t num_blocks() const; size_t block_size() const; + size_t max_batch_size() const; private: size_t num_blocks_; size_t block_size_; + size_t max_batch_size_; }; namespace PagedKVCache { diff --git a/csrc/engine/compiler/paged_compiler.cpp b/csrc/engine/compiler/paged_compiler.cpp index df3fd1cb4..d711c942d 100644 --- a/csrc/engine/compiler/paged_compiler.cpp +++ b/csrc/engine/compiler/paged_compiler.cpp @@ -38,6 +38,22 @@ PagedCompiler::PagedCompiler(const std::shared_ptr &model, RankBa for (size_t b = 256; b <= 512; b += 64) { decode_batch_sizes_.push_back(b); } + + const auto *config = dynamic_cast(model_->get_cache_config()); + if (config == nullptr) { + return; + } + const size_t max_batch_size = config->max_batch_size(); + if (max_batch_size == 0) { + throw std::invalid_argument("Paged Graph max_batch_size must be greater than zero"); + } + decode_batch_sizes_.erase( + std::remove_if(decode_batch_sizes_.begin(), decode_batch_sizes_.end(), + [max_batch_size](size_t b) { return b > max_batch_size; }), + decode_batch_sizes_.end()); + if (decode_batch_sizes_.empty() || decode_batch_sizes_.back() != max_batch_size) { + decode_batch_sizes_.push_back(max_batch_size); + } } void PagedCompiler::compile() { diff --git a/csrc/pybind11/cache/cache.hpp b/csrc/pybind11/cache/cache.hpp index 492f6c302..20c8cfb68 100644 --- a/csrc/pybind11/cache/cache.hpp +++ b/csrc/pybind11/cache/cache.hpp @@ -34,9 +34,10 @@ inline void bind_cache(py::module &m) { infinilm::cache::CacheConfig, std::shared_ptr>(m, "PagedKVCacheConfig") .def( - py::init(), + py::init(), py::arg("num_blocks"), - py::arg("block_size") = 256) + py::arg("block_size") = 256, + py::arg("max_batch_size") = 512) .def( "num_blocks", &infinilm::cache::PagedKVCacheConfig::num_blocks) diff --git a/python/infinilm/cache/cache.py b/python/infinilm/cache/cache.py index 4a9bcc446..b8b8e41c4 100644 --- a/python/infinilm/cache/cache.py +++ b/python/infinilm/cache/cache.py @@ -18,5 +18,9 @@ def __init__( class PagedKVCacheConfig(CacheConfig, _infinilm.PagedKVCacheConfig): - def __init__(self, num_blocks: int, block_size: int = 256): - _infinilm.PagedKVCacheConfig.__init__(self, num_blocks, block_size) + def __init__( + self, num_blocks: int, block_size: int = 256, max_batch_size: int = 512 + ): + _infinilm.PagedKVCacheConfig.__init__( + self, num_blocks, block_size, max_batch_size + ) diff --git a/python/infinilm/llm/model_runner/model_runner.py b/python/infinilm/llm/model_runner/model_runner.py index bbf3ccd37..ef6d9e13e 100644 --- a/python/infinilm/llm/model_runner/model_runner.py +++ b/python/infinilm/llm/model_runner/model_runner.py @@ -59,7 +59,9 @@ def __init__(self, config: EngineConfig): ) elif config.cache_type == "paged": cache_config = PagedKVCacheConfig( - num_blocks=config.num_blocks, block_size=config.block_size + num_blocks=config.num_blocks, + block_size=config.block_size, + max_batch_size=config.max_batch_size, ) logger.info(f"Using Paged KV Cache with num_blocks={config.num_blocks}") else: