From 97783fc038c0a2c5e7595f25383c0b825226db6c Mon Sep 17 00:00:00 2001 From: luzhan <513964121@qq.com> Date: Fri, 27 Feb 2026 18:12:41 +0000 Subject: [PATCH 01/17] Implement prefix reuse functionality in HostPagedKVWorkerView and related components - Added AllocatePagesForSequencesWithPrefix method to HostPagedKVWorkerView for handling sequences with prefix tokens. - Introduced PrefixSequenceState struct to manage state related to prefix reuse. - Implemented methods for updating, committing, and removing prefix states. - Enhanced the HostKVPrefixCache with a harness for testing prefix cache behavior. - Updated bindings in Python to expose new functionality for prefix reuse. - Added tests for prefix reuse scenarios, including cache hits and eviction behavior. - Modified configuration structure to include parameters for prefix reuse settings. --- batchgen/batchgen_worker.py | 22 +- batchgen/config/config.py | 9 + batchgen/kv_cache/host_kv_mananger_config.py | 20 +- core/KV_Storage/host_paged_kv_backend.cpp | 869 +++++++++++++----- core/KV_Storage/host_paged_kv_backend.h | 90 +- core/KV_Storage/host_paged_kv_config_utils.h | 9 + .../KV_Storage/host_paged_kv_prefix_cache.cpp | 593 ++++++++++++ core/KV_Storage/host_paged_kv_prefix_cache.h | 175 ++++ core/KV_Storage/host_paged_kv_worker_view.h | 164 +++- core/batchgen_Binding.cpp | 385 +++++++- core/data_structures.h | 9 + core/utils.cpp | 44 + op_builder/core_engine.py | 18 +- test/paged_kv/test_host_prefix_reuse.py | 144 +++ test/paged_kv/test_prefix_cache_binding.py | 100 ++ 15 files changed, 2370 insertions(+), 281 deletions(-) create mode 100644 core/KV_Storage/host_paged_kv_prefix_cache.cpp create mode 100644 core/KV_Storage/host_paged_kv_prefix_cache.h create mode 100644 test/paged_kv/test_host_prefix_reuse.py create mode 100644 test/paged_kv/test_prefix_cache_binding.py diff --git a/batchgen/batchgen_worker.py b/batchgen/batchgen_worker.py index 34491a899..4e6ae25c3 100644 --- a/batchgen/batchgen_worker.py +++ b/batchgen/batchgen_worker.py @@ -4411,20 +4411,37 @@ def _config_prefill_for_batch(self, prefill_uuids: List[str]) -> None: if my_prefill_uuids: global_sequence_ids = [] sequence_tokens = [] + flat_prompt_tokens = [] + prompt_offsets = [0] for uuid in my_prefill_uuids: seq = self.global_batch.get_sequence(uuid) global_sequence_ids.append(seq.global_idx) sequence_tokens.append(seq.kv_token_budget) + prompt_ids = seq.input_ids[0, :seq.prompt_length].tolist() + flat_prompt_tokens.extend(int(t) for t in prompt_ids) + prompt_offsets.append(len(flat_prompt_tokens)) logging.debug( f"Rank {self.rank}: Registering {len(global_sequence_ids)} sequences for host KV" ) self.core_engine.host_paged_kv_worker_view.register_sequences(global_sequence_ids) - self.core_engine.host_paged_kv_worker_view.allocate_pages_for_sequences( - list(zip(global_sequence_ids, sequence_tokens)) + host_cfg = getattr(self.engine_config, "Host_Paged_KV_Config", None) + enable_prefix_reuse = ( + host_cfg is not None + and getattr(host_cfg, "enable_prefix_reuse", False) ) + if enable_prefix_reuse: + self.core_engine.host_paged_kv_worker_view.allocate_pages_for_sequences_with_prefix( + list(zip(global_sequence_ids, sequence_tokens)), + flat_prompt_tokens, + prompt_offsets, + ) + else: + self.core_engine.host_paged_kv_worker_view.allocate_pages_for_sequences( + list(zip(global_sequence_ids, sequence_tokens)) + ) kv_stats = self.core_engine.host_paged_kv_worker_view.get_stats() if self.rank == 0: @@ -8103,4 +8120,3 @@ def _reset_for_new_batch(self) -> None: dist.barrier() logging.info(f"Rank {self.rank}: State reset completed") - diff --git a/batchgen/config/config.py b/batchgen/config/config.py index fc36aed2f..cffd79739 100644 --- a/batchgen/config/config.py +++ b/batchgen/config/config.py @@ -176,6 +176,15 @@ class HostPagedKVConfig: num_v_heads: int = 0 # Zero for MLA. v_head_dim: int = 0 kv_dtype: str = "bfloat16" # "bfloat16 or float8_e4m3fn" + enable_prefix_reuse: bool = False + prefix_min_reuse_pages: int = 1 + prefix_min_store_pages: int = 2 + sequence_page_node_capacity: int = 0 + radix_node_capacity: int = 0 + radix_edge_capacity: int = 0 + prefix_entry_capacity: int = 0 + prefix_page_ref_capacity: int = 0 + prefix_page_budget: int = 0 @dataclass class DevicePagedKVConfig: diff --git a/batchgen/kv_cache/host_kv_mananger_config.py b/batchgen/kv_cache/host_kv_mananger_config.py index d75cc8e6d..9e5cbde2f 100644 --- a/batchgen/kv_cache/host_kv_mananger_config.py +++ b/batchgen/kv_cache/host_kv_mananger_config.py @@ -221,6 +221,15 @@ def build_host_kv_config(model_name: str, host_kv_cache_size: int) -> Any: profile.sequence_table_capacity or config.num_pages ) config.alignment_bytes = profile.alignment_bytes + config.enable_prefix_reuse = False + config.prefix_min_reuse_pages = 1 + config.prefix_min_store_pages = 2 + config.sequence_page_node_capacity = 0 + config.radix_node_capacity = 0 + config.radix_edge_capacity = 0 + config.prefix_entry_capacity = 0 + config.prefix_page_ref_capacity = 0 + config.prefix_page_budget = 0 return config @@ -336,6 +345,15 @@ def build_host_kv_config_aux(model_name: str, host_kv_cache_size: int) -> Any | profile.sequence_table_capacity or config.num_pages ) config.alignment_bytes = profile.alignment_bytes + config.enable_prefix_reuse = False + config.prefix_min_reuse_pages = 1 + config.prefix_min_store_pages = 2 + config.sequence_page_node_capacity = 0 + config.radix_node_capacity = 0 + config.radix_edge_capacity = 0 + config.prefix_entry_capacity = 0 + config.prefix_page_ref_capacity = 0 + config.prefix_page_budget = 0 return config @@ -397,4 +415,4 @@ def build_gpu_kv_config_fixed_size( num_v_heads=profile.num_v_heads, v_head_dim=profile.v_head_dim, kv_dtype=_torch_dtype_from_string(profile.kv_dtype), - ) \ No newline at end of file + ) diff --git a/core/KV_Storage/host_paged_kv_backend.cpp b/core/KV_Storage/host_paged_kv_backend.cpp index 76e5175aa..b161bf88d 100644 --- a/core/KV_Storage/host_paged_kv_backend.cpp +++ b/core/KV_Storage/host_paged_kv_backend.cpp @@ -1,4 +1,5 @@ #include "host_paged_kv_backend.h" +#include "host_paged_kv_prefix_cache.h" #include #include @@ -12,7 +13,6 @@ #include #include #include -#include #include #include #include @@ -28,7 +28,7 @@ namespace { constexpr std::uint64_t kSharedMemoryMagic = 0x484f53544b564d47ULL; // "HOSTKVMG" -constexpr std::int32_t kInvalidPageIndex = -1; +constexpr std::int32_t kInvalidIndex = -1; constexpr std::int64_t kEmptySequenceId = std::numeric_limits::min(); constexpr std::int64_t kTombstoneSequenceId = kEmptySequenceId + 1; @@ -82,7 +82,7 @@ class ScopedMutexLock { ~ScopedMutexLock() { const int rc = pthread_mutex_unlock(mu_); if (rc != 0) { - std::terminate(); // Unlock failure is irrecoverable here. + std::terminate(); } } @@ -90,12 +90,53 @@ class ScopedMutexLock { pthread_mutex_t* mu_; }; +// Helper function to perform aligned mmap. +void* mmap_aligned(size_t length, int prot, int flags, int fd, off_t offset, + size_t alignment) { + const size_t total_len = length + alignment; + void* addr = + mmap(nullptr, total_len, PROT_NONE, MAP_PRIVATE | MAP_ANONYMOUS, -1, 0); + if (addr == MAP_FAILED) { + return MAP_FAILED; + } + + const uintptr_t raw_addr = reinterpret_cast(addr); + const uintptr_t aligned_addr = + (raw_addr + alignment - 1) & ~(alignment - 1); + void* final_addr = reinterpret_cast(aligned_addr); + + void* ret = mmap(final_addr, length, prot, flags | MAP_FIXED, fd, offset); + + const size_t prefix_len = aligned_addr - raw_addr; + if (prefix_len > 0) { + munmap(addr, prefix_len); + } + + const size_t suffix_len = total_len - length - prefix_len; + if (suffix_len > 0) { + munmap(reinterpret_cast(aligned_addr + length), suffix_len); + } + + return ret; +} + +} // namespace + struct SequenceEntry { std::int64_t sequence_id = kEmptySequenceId; std::uint32_t num_pages = 0; - std::int32_t head_page = kInvalidPageIndex; - std::int32_t tail_page = kInvalidPageIndex; + std::int32_t head_node = kInvalidIndex; + std::int32_t tail_node = kInvalidIndex; +}; + +struct SeqPageNode { + std::int32_t page_idx = kInvalidIndex; + std::int32_t next_node = kInvalidIndex; }; +using RadixNode = HostKVRadixNode; +using RadixEdge = HostKVRadixEdge; +using PrefixEntry = HostKVPrefixEntry; +using PrefixPageRef = HostKVPrefixPageRef; struct SharedHeader { std::atomic init_state{ @@ -110,53 +151,36 @@ struct SharedHeader { std::uint64_t num_layers = 0; std::uint64_t page_size_tokens = 0; std::uint32_t has_v_cache = 0; - std::atomic free_stack_top{0}; - std::atomic active_sequences{0}; - pthread_mutex_t allocation_mutex{}; - pthread_mutex_t sequence_mutex{}; -}; + std::uint32_t enable_prefix_reuse = 0; -std::size_t SafeHardwareConcurrency() { - const unsigned int hint = std::thread::hardware_concurrency() / 4; - return hint == 0 ? 1 : static_cast(hint); -} + std::uint64_t sequence_page_node_capacity = 0; + std::uint64_t radix_node_capacity = 0; + std::uint64_t radix_edge_capacity = 0; + std::uint64_t prefix_entry_capacity = 0; + std::uint64_t prefix_page_ref_capacity = 0; + std::uint64_t prefix_page_budget = 0; -// Helper function to perform aligned mmap -// This ensures the virtual address is aligned to the specified alignment (e.g., 2MB for huge pages) -// which is often required for cudaHostRegister to work correctly with huge pages. -void* mmap_aligned(size_t length, int prot, int flags, int fd, off_t offset, size_t alignment) { - // Allocate extra space to ensure we can find an aligned segment - size_t total_len = length + alignment; - - // Reserve address space using anonymous mapping - void* addr = mmap(nullptr, total_len, PROT_NONE, MAP_PRIVATE | MAP_ANONYMOUS, -1, 0); - if (addr == MAP_FAILED) { - return MAP_FAILED; - } + std::atomic free_stack_top{0}; + std::atomic seq_page_node_free_top{0}; + std::atomic radix_node_free_top{0}; + std::atomic radix_edge_free_top{0}; + std::atomic prefix_entry_free_top{0}; + std::atomic prefix_page_ref_free_top{0}; - uintptr_t raw_addr = reinterpret_cast(addr); - uintptr_t aligned_addr = (raw_addr + alignment - 1) & ~(alignment - 1); - void* final_addr = reinterpret_cast(aligned_addr); + std::atomic active_sequences{0}; + std::atomic prefix_entry_count{0}; + std::atomic prefix_used_pages{0}; - // Map the file into the aligned position using MAP_FIXED - // This replaces the anonymous mapping at that location - void* ret = mmap(final_addr, length, prot, flags | MAP_FIXED, fd, offset); - - // Unmap the unused parts of the reservation - size_t prefix_len = aligned_addr - raw_addr; - if (prefix_len > 0) { - munmap(addr, prefix_len); - } - - size_t suffix_len = total_len - length - prefix_len; - if (suffix_len > 0) { - munmap(reinterpret_cast(aligned_addr + length), suffix_len); - } + std::atomic prefix_access_epoch{0}; + std::atomic prefix_hit_count{0}; + std::atomic prefix_miss_count{0}; + std::atomic prefix_evict_count{0}; - return ret; -} + std::int32_t lru_head = kInvalidIndex; + std::int32_t lru_tail = kInvalidIndex; -} // namespace + pthread_mutex_t metadata_mutex{}; +}; struct HostPagedKVBackend::SharedState { explicit SharedState(const HostPagedKVConfig& cfg, std::size_t data_bytes, @@ -168,15 +192,28 @@ struct HostPagedKVBackend::SharedState { sequence_capacity = config.sequence_table_capacity == 0 ? config.num_pages : config.sequence_table_capacity; + sequence_page_node_capacity = config.sequence_page_node_capacity; + radix_node_capacity = config.radix_node_capacity; + radix_edge_capacity = config.radix_edge_capacity; + prefix_entry_capacity = config.prefix_entry_capacity; + prefix_page_ref_capacity = config.prefix_page_ref_capacity; ComputeOffsets(); } void Initialize(bool create_region); std::vector AcquirePages(std::int64_t sequence_id, std::size_t num_pages); + PrefixAllocationBatchResult AcquirePagesForSequencesWithPrefix( + const std::vector& sequence_ids, + const std::vector& num_tokens, + const std::vector& flat_prompt_tokens, + const std::vector& prompt_offsets); void ReleaseSequence(std::int64_t sequence_id); std::vector SequencePages( std::int64_t sequence_id, std::optional max_pages) const; + void CommitSequencePrefix(std::int64_t sequence_id, + const std::vector& prompt_tokens, + std::size_t prompt_token_count); HostPagedKVStats CollectStats() const; std::byte* DataBase() { return data_base; } @@ -190,34 +227,84 @@ struct HostPagedKVBackend::SharedState { int shm_fd = -1; std::size_t total_bytes = 0; std::byte* mapping = nullptr; + SharedHeader* header = nullptr; std::int32_t* free_stack = nullptr; - std::int64_t* page_owners = nullptr; - std::int32_t* page_links = nullptr; + std::uint32_t* page_refcount = nullptr; SequenceEntry* sequence_table = nullptr; + SeqPageNode* seq_page_nodes = nullptr; + std::int32_t* seq_page_node_free_stack = nullptr; + + RadixNode* radix_nodes = nullptr; + RadixEdge* radix_edges = nullptr; + std::int32_t* radix_node_free_stack = nullptr; + std::int32_t* radix_edge_free_stack = nullptr; + + PrefixEntry* prefix_entries = nullptr; + PrefixPageRef* prefix_page_refs = nullptr; + std::int32_t* prefix_entry_free_stack = nullptr; + std::int32_t* prefix_page_ref_free_stack = nullptr; + std::byte* data_base = nullptr; bool created_region = false; std::size_t header_offset = 0; std::size_t free_stack_offset = 0; - std::size_t page_owner_offset = 0; - std::size_t page_link_offset = 0; + std::size_t page_refcount_offset = 0; std::size_t sequence_table_offset = 0; + std::size_t seq_page_nodes_offset = 0; + std::size_t seq_page_node_free_stack_offset = 0; + std::size_t radix_nodes_offset = 0; + std::size_t radix_edges_offset = 0; + std::size_t radix_node_free_stack_offset = 0; + std::size_t radix_edge_free_stack_offset = 0; + std::size_t prefix_entries_offset = 0; + std::size_t prefix_page_refs_offset = 0; + std::size_t prefix_entry_free_stack_offset = 0; + std::size_t prefix_page_ref_free_stack_offset = 0; std::size_t data_offset = 0; std::size_t total_bytes_unaligned = 0; + std::size_t sequence_capacity = 0; + std::size_t sequence_page_node_capacity = 0; + std::size_t radix_node_capacity = 0; + std::size_t radix_edge_capacity = 0; + std::size_t prefix_entry_capacity = 0; + std::size_t prefix_page_ref_capacity = 0; + HostKVPrefixCache prefix_cache_; private: void ComputeOffsets(); void MapPointers(); + void BindPrefixCache(); void ConstructSharedState(); void WaitForInitialization() const; void ValidateSharedState() const; + + std::size_t HashSequenceId(std::int64_t sequence_id) const; + SequenceEntry* FindSequenceEntryLocked(std::int64_t sequence_id) const; SequenceEntry* FindOrInsertSequenceEntryLocked(std::int64_t sequence_id, bool* is_new); - SequenceEntry* FindSequenceEntryLocked(std::int64_t sequence_id) const; - std::size_t HashSequenceId(std::int64_t sequence_id) const; + + std::int32_t PopStackIndexLocked(std::int32_t* stack, + std::atomic* top, + const char* what) const; + void PushStackIndexLocked(std::int32_t* stack, + std::atomic* top, + std::int32_t value) const; + + std::int32_t AllocateSeqPageNodeLocked(std::int32_t page_idx); + void FreeSeqPageNodeLocked(std::int32_t node_idx); + void AppendPageToSequenceLocked(SequenceEntry* entry, std::int32_t page_idx); + std::vector CollectSequencePagesLocked( + const SequenceEntry* entry, std::optional max_pages) const; + void ReleaseSequencePagesLocked(SequenceEntry* entry); + + std::int32_t PopFreePageLocked(); + void PushFreePageLocked(std::int32_t page_idx); + void IncrementPageRefLocked(std::int32_t page_idx); + void DecrementPageRefLocked(std::int32_t page_idx); }; void HostPagedKVBackend::SharedState::ComputeOffsets() { @@ -231,23 +318,54 @@ void HostPagedKVBackend::SharedState::ComputeOffsets() { free_stack_offset = offset; offset += sizeof(std::int32_t) * config.num_pages; - offset = AlignUp(offset, alignof(std::int64_t)); - page_owner_offset = offset; - offset += sizeof(std::int64_t) * config.num_pages; - - offset = AlignUp(offset, alignof(std::int32_t)); - page_link_offset = offset; - offset += sizeof(std::int32_t) * config.num_pages; + offset = AlignUp(offset, alignof(std::uint32_t)); + page_refcount_offset = offset; + offset += sizeof(std::uint32_t) * config.num_pages; offset = AlignUp(offset, alignof(SequenceEntry)); sequence_table_offset = offset; offset += sizeof(SequenceEntry) * sequence_capacity; - // Align data_offset to 2MB (huge page size) because cudaHostRegister may require - // huge-page aligned pointers when the underlying memory is backed by huge pages - // (e.g., via Transparent Huge Pages or hugetlbfs). - // Using simple page alignment (4KB) can cause "invalid argument" errors with cudaHostRegister - // on some systems when THP is active. + offset = AlignUp(offset, alignof(SeqPageNode)); + seq_page_nodes_offset = offset; + offset += sizeof(SeqPageNode) * sequence_page_node_capacity; + + offset = AlignUp(offset, alignof(std::int32_t)); + seq_page_node_free_stack_offset = offset; + offset += sizeof(std::int32_t) * sequence_page_node_capacity; + + offset = AlignUp(offset, alignof(RadixNode)); + radix_nodes_offset = offset; + offset += sizeof(RadixNode) * radix_node_capacity; + + offset = AlignUp(offset, alignof(RadixEdge)); + radix_edges_offset = offset; + offset += sizeof(RadixEdge) * radix_edge_capacity; + + offset = AlignUp(offset, alignof(std::int32_t)); + radix_node_free_stack_offset = offset; + offset += sizeof(std::int32_t) * radix_node_capacity; + + offset = AlignUp(offset, alignof(std::int32_t)); + radix_edge_free_stack_offset = offset; + offset += sizeof(std::int32_t) * radix_edge_capacity; + + offset = AlignUp(offset, alignof(PrefixEntry)); + prefix_entries_offset = offset; + offset += sizeof(PrefixEntry) * prefix_entry_capacity; + + offset = AlignUp(offset, alignof(PrefixPageRef)); + prefix_page_refs_offset = offset; + offset += sizeof(PrefixPageRef) * prefix_page_ref_capacity; + + offset = AlignUp(offset, alignof(std::int32_t)); + prefix_entry_free_stack_offset = offset; + offset += sizeof(std::int32_t) * prefix_entry_capacity; + + offset = AlignUp(offset, alignof(std::int32_t)); + prefix_page_ref_free_stack_offset = offset; + offset += sizeof(std::int32_t) * prefix_page_ref_capacity; + constexpr std::size_t kHugePageAlignment = 2 * 1024 * 1024; offset = AlignUp(offset, kHugePageAlignment); data_offset = offset; @@ -259,16 +377,67 @@ void HostPagedKVBackend::SharedState::ComputeOffsets() { void HostPagedKVBackend::SharedState::MapPointers() { header = reinterpret_cast(mapping + header_offset); free_stack = reinterpret_cast(mapping + free_stack_offset); - page_owners = reinterpret_cast(mapping + page_owner_offset); - page_links = reinterpret_cast(mapping + page_link_offset); + page_refcount = + reinterpret_cast(mapping + page_refcount_offset); sequence_table = reinterpret_cast(mapping + sequence_table_offset); + seq_page_nodes = + reinterpret_cast(mapping + seq_page_nodes_offset); + seq_page_node_free_stack = reinterpret_cast( + mapping + seq_page_node_free_stack_offset); + + radix_nodes = reinterpret_cast(mapping + radix_nodes_offset); + radix_edges = reinterpret_cast(mapping + radix_edges_offset); + radix_node_free_stack = + reinterpret_cast(mapping + radix_node_free_stack_offset); + radix_edge_free_stack = + reinterpret_cast(mapping + radix_edge_free_stack_offset); + + prefix_entries = + reinterpret_cast(mapping + prefix_entries_offset); + prefix_page_refs = + reinterpret_cast(mapping + prefix_page_refs_offset); + prefix_entry_free_stack = reinterpret_cast( + mapping + prefix_entry_free_stack_offset); + prefix_page_ref_free_stack = reinterpret_cast( + mapping + prefix_page_ref_free_stack_offset); + data_base = mapping + data_offset; } +void HostPagedKVBackend::SharedState::BindPrefixCache() { + HostKVPrefixCacheParams params; + params.enable_prefix_reuse = config.enable_prefix_reuse; + params.prefix_min_reuse_pages = config.prefix_min_reuse_pages; + params.prefix_min_store_pages = config.prefix_min_store_pages; + params.prefix_page_budget = config.prefix_page_budget; + + HostKVPrefixCache::SharedFields shared_fields; + shared_fields.radix_node_free_top = &header->radix_node_free_top; + shared_fields.radix_edge_free_top = &header->radix_edge_free_top; + shared_fields.prefix_entry_free_top = &header->prefix_entry_free_top; + shared_fields.prefix_page_ref_free_top = &header->prefix_page_ref_free_top; + shared_fields.prefix_entry_count = &header->prefix_entry_count; + shared_fields.prefix_used_pages = &header->prefix_used_pages; + shared_fields.prefix_access_epoch = &header->prefix_access_epoch; + shared_fields.prefix_hit_count = &header->prefix_hit_count; + shared_fields.prefix_miss_count = &header->prefix_miss_count; + shared_fields.prefix_evict_count = &header->prefix_evict_count; + shared_fields.lru_head = &header->lru_head; + shared_fields.lru_tail = &header->lru_tail; + + prefix_cache_.Bind( + params, radix_nodes, radix_edges, prefix_entries, prefix_page_refs, + radix_node_free_stack, radix_edge_free_stack, prefix_entry_free_stack, + prefix_page_ref_free_stack, shared_fields, + [this](std::int32_t page_idx) { IncrementPageRefLocked(page_idx); }, + [this](std::int32_t page_idx) { DecrementPageRefLocked(page_idx); }); +} + void HostPagedKVBackend::SharedState::ConstructSharedState() { std::memset(mapping, 0, total_bytes); MapPointers(); + header->magic = kSharedMemoryMagic; header->layout_fingerprint = layout_fingerprint; header->config_hash = HashHostKVConfig(config); @@ -278,20 +447,66 @@ void HostPagedKVBackend::SharedState::ConstructSharedState() { header->alignment_bytes = config.alignment_bytes; header->num_layers = config.num_layers; header->page_size_tokens = config.page_size_tokens; - header->has_v_cache = has_v_cache ? 1 : 0; + header->has_v_cache = has_v_cache ? 1U : 0U; + header->enable_prefix_reuse = config.enable_prefix_reuse ? 1U : 0U; + + header->sequence_page_node_capacity = sequence_page_node_capacity; + header->radix_node_capacity = radix_node_capacity; + header->radix_edge_capacity = radix_edge_capacity; + header->prefix_entry_capacity = prefix_entry_capacity; + header->prefix_page_ref_capacity = prefix_page_ref_capacity; + header->prefix_page_budget = config.prefix_page_budget; + header->free_stack_top.store(static_cast(config.num_pages), std::memory_order_relaxed); + header->seq_page_node_free_top.store( + static_cast(sequence_page_node_capacity), + std::memory_order_relaxed); + header->radix_node_free_top.store( + static_cast(radix_node_capacity > 0 + ? radix_node_capacity - 1 + : 0), + std::memory_order_relaxed); + header->radix_edge_free_top.store( + static_cast(radix_edge_capacity), + std::memory_order_relaxed); + header->prefix_entry_free_top.store( + static_cast(prefix_entry_capacity), + std::memory_order_relaxed); + header->prefix_page_ref_free_top.store( + static_cast(prefix_page_ref_capacity), + std::memory_order_relaxed); + header->active_sequences.store(0, std::memory_order_relaxed); + header->prefix_entry_count.store(0, std::memory_order_relaxed); + header->prefix_used_pages.store(0, std::memory_order_relaxed); + header->prefix_access_epoch.store(0, std::memory_order_relaxed); + header->prefix_hit_count.store(0, std::memory_order_relaxed); + header->prefix_miss_count.store(0, std::memory_order_relaxed); + header->prefix_evict_count.store(0, std::memory_order_relaxed); + header->lru_head = kInvalidIndex; + header->lru_tail = kInvalidIndex; for (std::size_t i = 0; i < config.num_pages; ++i) { free_stack[i] = static_cast(config.num_pages - 1 - i); - page_owners[i] = kEmptySequenceId; - page_links[i] = kInvalidPageIndex; + page_refcount[i] = 0; } + for (std::size_t i = 0; i < sequence_capacity; ++i) { sequence_table[i] = SequenceEntry(); } + for (std::size_t i = 0; i < sequence_page_node_capacity; ++i) { + seq_page_nodes[i] = SeqPageNode(); + seq_page_node_free_stack[i] = + static_cast(sequence_page_node_capacity - 1 - i); + } + + BindPrefixCache(); + prefix_cache_.InitializePools(radix_node_capacity, radix_edge_capacity, + prefix_entry_capacity, + prefix_page_ref_capacity); + pthread_mutexattr_t attr; if (const int rc = pthread_mutexattr_init(&attr); rc != 0) { throw std::system_error(rc, std::generic_category(), @@ -311,18 +526,11 @@ void HostPagedKVBackend::SharedState::ConstructSharedState() { "pthread_mutexattr_setrobust failed"); } - if (const int rc = pthread_mutex_init(&header->allocation_mutex, &attr); - rc != 0) { - pthread_mutexattr_destroy(&attr); - throw std::system_error(rc, std::generic_category(), - "pthread_mutex_init allocation_mutex failed"); - } - if (const int rc = pthread_mutex_init(&header->sequence_mutex, &attr); + if (const int rc = pthread_mutex_init(&header->metadata_mutex, &attr); rc != 0) { - pthread_mutex_destroy(&header->allocation_mutex); pthread_mutexattr_destroy(&attr); throw std::system_error(rc, std::generic_category(), - "pthread_mutex_init sequence_mutex failed"); + "pthread_mutex_init metadata_mutex failed"); } pthread_mutexattr_destroy(&attr); @@ -337,10 +545,6 @@ void HostPagedKVBackend::SharedState::WaitForInitialization() const { if (state == InitState::kReady) { return; } - if (state == InitState::kUninitialized) { - std::this_thread::sleep_for(std::chrono::milliseconds(1)); - continue; - } std::this_thread::sleep_for(std::chrono::milliseconds(1)); } } @@ -385,6 +589,20 @@ void HostPagedKVBackend::SharedState::ValidateSharedState() const { << ")"; throw std::runtime_error(oss.str()); } + if (header->sequence_page_node_capacity != sequence_page_node_capacity) { + throw std::runtime_error("Shared memory sequence_page_node_capacity mismatch"); + } + if (header->radix_node_capacity != radix_node_capacity || + header->radix_edge_capacity != radix_edge_capacity) { + throw std::runtime_error("Shared memory radix pool capacity mismatch"); + } + if (header->prefix_entry_capacity != prefix_entry_capacity || + header->prefix_page_ref_capacity != prefix_page_ref_capacity) { + throw std::runtime_error("Shared memory prefix pool capacity mismatch"); + } + if (header->prefix_page_budget != config.prefix_page_budget) { + throw std::runtime_error("Shared memory prefix_page_budget mismatch"); + } } std::size_t HostPagedKVBackend::SharedState::HashSequenceId( @@ -429,10 +647,8 @@ SequenceEntry* HostPagedKVBackend::SharedState::FindOrInsertSequenceEntryLocked( if (entry->sequence_id == kEmptySequenceId) { SequenceEntry* target = first_tombstone != nullptr ? first_tombstone : entry; - if (target->sequence_id != sequence_id) { - *target = SequenceEntry(); - target->sequence_id = sequence_id; - } + *target = SequenceEntry(); + target->sequence_id = sequence_id; if (is_new != nullptr) { *is_new = true; } @@ -448,6 +664,122 @@ SequenceEntry* HostPagedKVBackend::SharedState::FindOrInsertSequenceEntryLocked( std::to_string(sequence_capacity) + ")"); } +std::int32_t HostPagedKVBackend::SharedState::PopStackIndexLocked( + std::int32_t* stack, std::atomic* top, + const char* what) const { + const std::uint32_t current = top->load(std::memory_order_relaxed); + if (current == 0) { + throw std::runtime_error(std::string("Out of ") + what); + } + const std::uint32_t next = current - 1; + const std::int32_t value = stack[next]; + top->store(next, std::memory_order_relaxed); + return value; +} + +void HostPagedKVBackend::SharedState::PushStackIndexLocked( + std::int32_t* stack, std::atomic* top, + std::int32_t value) const { + const std::uint32_t current = top->load(std::memory_order_relaxed); + stack[current] = value; + top->store(current + 1, std::memory_order_relaxed); +} + +std::int32_t HostPagedKVBackend::SharedState::AllocateSeqPageNodeLocked( + std::int32_t page_idx) { + const std::int32_t node_idx = PopStackIndexLocked( + seq_page_node_free_stack, &header->seq_page_node_free_top, + "sequence page nodes"); + seq_page_nodes[node_idx].page_idx = page_idx; + seq_page_nodes[node_idx].next_node = kInvalidIndex; + return node_idx; +} + +void HostPagedKVBackend::SharedState::FreeSeqPageNodeLocked( + std::int32_t node_idx) { + seq_page_nodes[node_idx] = SeqPageNode(); + PushStackIndexLocked(seq_page_node_free_stack, + &header->seq_page_node_free_top, node_idx); +} + +void HostPagedKVBackend::SharedState::AppendPageToSequenceLocked( + SequenceEntry* entry, std::int32_t page_idx) { + const std::int32_t node_idx = AllocateSeqPageNodeLocked(page_idx); + if (entry->head_node == kInvalidIndex) { + entry->head_node = node_idx; + entry->tail_node = node_idx; + } else { + seq_page_nodes[entry->tail_node].next_node = node_idx; + entry->tail_node = node_idx; + } + ++entry->num_pages; +} + +std::vector HostPagedKVBackend::SharedState::CollectSequencePagesLocked( + const SequenceEntry* entry, std::optional max_pages) const { + const std::size_t available_pages = entry->num_pages; + const std::size_t limit = max_pages.has_value() + ? std::min(max_pages.value(), available_pages) + : available_pages; + if (max_pages.has_value() && max_pages.value() > available_pages) { + throw std::out_of_range( + "Requested " + std::to_string(max_pages.value()) + + " pages but only " + std::to_string(available_pages) + + " pages allocated for sequence " + + std::to_string(entry->sequence_id)); + } + + std::vector pages; + pages.reserve(limit); + std::int32_t node_idx = entry->head_node; + while (node_idx != kInvalidIndex && pages.size() < limit) { + pages.push_back(seq_page_nodes[node_idx].page_idx); + node_idx = seq_page_nodes[node_idx].next_node; + } + return pages; +} + +void HostPagedKVBackend::SharedState::ReleaseSequencePagesLocked( + SequenceEntry* entry) { + std::int32_t node_idx = entry->head_node; + while (node_idx != kInvalidIndex) { + const std::int32_t next = seq_page_nodes[node_idx].next_node; + const std::int32_t page_idx = seq_page_nodes[node_idx].page_idx; + DecrementPageRefLocked(page_idx); + FreeSeqPageNodeLocked(node_idx); + node_idx = next; + } + entry->num_pages = 0; + entry->head_node = kInvalidIndex; + entry->tail_node = kInvalidIndex; +} + +std::int32_t HostPagedKVBackend::SharedState::PopFreePageLocked() { + return PopStackIndexLocked(free_stack, &header->free_stack_top, + "free pages"); +} + +void HostPagedKVBackend::SharedState::PushFreePageLocked(std::int32_t page_idx) { + PushStackIndexLocked(free_stack, &header->free_stack_top, page_idx); +} + +void HostPagedKVBackend::SharedState::IncrementPageRefLocked( + std::int32_t page_idx) { + ++page_refcount[page_idx]; +} + +void HostPagedKVBackend::SharedState::DecrementPageRefLocked( + std::int32_t page_idx) { + if (page_refcount[page_idx] == 0) { + throw std::runtime_error("page_refcount underflow on page " + + std::to_string(page_idx)); + } + --page_refcount[page_idx]; + if (page_refcount[page_idx] == 0) { + PushFreePageLocked(page_idx); + } +} + void HostPagedKVBackend::SharedState::Initialize(bool create_region) { const std::size_t page_size = GetSystemPageSize(); total_bytes = AlignUp(total_bytes_unaligned, page_size); @@ -519,15 +851,11 @@ void HostPagedKVBackend::SharedState::Initialize(bool create_region) { } } - // Use mmap_aligned to ensure 2MB alignment for the shared memory mapping - // This is critical for cudaHostRegister to work reliably with huge pages (THP) constexpr std::size_t kHugePageSize = 2 * 1024 * 1024; - // Ensure alignment is at least the system page size const std::size_t alignment = std::max(kHugePageSize, page_size); - - void* mapped = mmap_aligned(total_bytes, PROT_READ | PROT_WRITE, - MAP_SHARED, shm_fd, 0, alignment); + void* mapped = mmap_aligned(total_bytes, PROT_READ | PROT_WRITE, MAP_SHARED, + shm_fd, 0, alignment); if (mapped == MAP_FAILED) { const int err = errno; close(shm_fd); @@ -535,12 +863,11 @@ void HostPagedKVBackend::SharedState::Initialize(bool create_region) { throw std::system_error(err, std::generic_category(), "mmap failed"); } - // Advise the kernel to use huge pages for this mapping if possible - // This complements the 2MB alignment we added to data_offset madvise(mapped, total_bytes, MADV_HUGEPAGE); mapping = static_cast(mapped); MapPointers(); + BindPrefixCache(); if (created_region) { header->init_state.store( @@ -558,130 +885,236 @@ std::vector HostPagedKVBackend::SharedState::AcquirePages( if (num_pages == 0) { return {}; } - std::vector pages(num_pages); - { - ScopedMutexLock lock(&header->allocation_mutex); - const std::uint32_t top = - header->free_stack_top.load(std::memory_order_relaxed); - if (top < num_pages) { + + std::vector pages; + pages.reserve(num_pages); + + ScopedMutexLock lock(&header->metadata_mutex); + + if (header->free_stack_top.load(std::memory_order_relaxed) < num_pages) { + throw std::runtime_error( + "Insufficient free pages available for sequence " + + std::to_string(sequence_id) + " (requested=" + + std::to_string(num_pages) + ", available=" + + std::to_string(header->free_stack_top.load(std::memory_order_relaxed)) + + ")"); + } + + if (header->seq_page_node_free_top.load(std::memory_order_relaxed) < num_pages) { + throw std::runtime_error("Insufficient sequence page nodes"); + } + + bool is_new = false; + SequenceEntry* entry = FindOrInsertSequenceEntryLocked(sequence_id, &is_new); + if (is_new) { + header->active_sequences.fetch_add(1, std::memory_order_relaxed); + } + + for (std::size_t i = 0; i < num_pages; ++i) { + const std::int32_t page_idx = PopFreePageLocked(); + IncrementPageRefLocked(page_idx); + AppendPageToSequenceLocked(entry, page_idx); + pages.push_back(page_idx); + } + + return pages; +} + +PrefixAllocationBatchResult +HostPagedKVBackend::SharedState::AcquirePagesForSequencesWithPrefix( + const std::vector& sequence_ids, + const std::vector& num_tokens, + const std::vector& flat_prompt_tokens, + const std::vector& prompt_offsets) { + if (sequence_ids.size() != num_tokens.size()) { + throw std::invalid_argument( + "sequence_ids and num_tokens must have the same length"); + } + if (prompt_offsets.size() != sequence_ids.size() + 1) { + throw std::invalid_argument( + "prompt_offsets must contain sequence_count + 1 entries"); + } + if (prompt_offsets.empty() || prompt_offsets.front() != 0 || + prompt_offsets.back() != flat_prompt_tokens.size()) { + throw std::invalid_argument("prompt_offsets boundaries are invalid"); + } + + PrefixAllocationBatchResult result; + result.allocated_pages.resize(sequence_ids.size()); + result.reused_prefix_tokens.resize(sequence_ids.size(), 0); + + for (std::size_t i = 0; i < sequence_ids.size(); ++i) { + const std::size_t begin = prompt_offsets[i]; + const std::size_t end = prompt_offsets[i + 1]; + if (begin > end || end > flat_prompt_tokens.size()) { + throw std::invalid_argument("prompt_offsets contain invalid slice"); + } + if (num_tokens[i] == 0) { + throw std::invalid_argument("num_tokens entries must be > 0"); + } + + const std::size_t required_pages = + (num_tokens[i] + config.page_size_tokens - 1) / config.page_size_tokens; + const std::size_t prompt_tokens = end - begin; + const std::size_t prompt_full_pages = + std::min(required_pages, prompt_tokens / config.page_size_tokens); + + if (!config.enable_prefix_reuse || prompt_full_pages == 0) { + result.allocated_pages[i] = AcquirePages(sequence_ids[i], required_pages); + continue; + } + + std::vector pages; + pages.reserve(required_pages); + + ScopedMutexLock lock(&header->metadata_mutex); + + const auto lookup = prefix_cache_.LookupPrefixPagesLocked( + flat_prompt_tokens.data() + begin, + prompt_full_pages * config.page_size_tokens, prompt_full_pages); + + const std::size_t reused_pages = std::min( + lookup.reused_pages, prompt_full_pages); + const std::size_t new_pages = required_pages - reused_pages; + + if (header->free_stack_top.load(std::memory_order_relaxed) < new_pages) { throw std::runtime_error( - "Insufficient free pages available for sequence " + - std::to_string(sequence_id) + - " (requested=" + std::to_string(num_pages) + - ", available=" + std::to_string(top) + ")"); + "Insufficient free pages for prefix-aware allocation of sequence " + + std::to_string(sequence_ids[i])); } - std::uint32_t new_top = top - static_cast(num_pages); - for (std::size_t i = 0; i < num_pages; ++i) { - pages[i] = free_stack[new_top + i]; + if (header->seq_page_node_free_top.load(std::memory_order_relaxed) < + required_pages) { + throw std::runtime_error( + "Insufficient sequence page nodes for prefix-aware allocation"); } - header->free_stack_top.store(new_top, std::memory_order_relaxed); - } - { - ScopedMutexLock lock(&header->sequence_mutex); bool is_new = false; SequenceEntry* entry = - FindOrInsertSequenceEntryLocked(sequence_id, &is_new); + FindOrInsertSequenceEntryLocked(sequence_ids[i], &is_new); if (is_new) { header->active_sequences.fetch_add(1, std::memory_order_relaxed); } - for (std::size_t i = 0; i < num_pages; ++i) { - const std::int32_t page = pages[i]; - page_owners[page] = sequence_id; - page_links[page] = kInvalidPageIndex; - if (entry->head_page == kInvalidPageIndex) { - entry->head_page = page; - entry->tail_page = page; - } else { - page_links[entry->tail_page] = page; - entry->tail_page = page; - } - ++entry->num_pages; + + for (std::size_t j = 0; j < reused_pages; ++j) { + const std::int32_t page_idx = lookup.pages[j]; + IncrementPageRefLocked(page_idx); + AppendPageToSequenceLocked(entry, page_idx); + pages.push_back(page_idx); + } + + for (std::size_t j = 0; j < new_pages; ++j) { + const std::int32_t page_idx = PopFreePageLocked(); + IncrementPageRefLocked(page_idx); + AppendPageToSequenceLocked(entry, page_idx); + pages.push_back(page_idx); } + + result.allocated_pages[i] = std::move(pages); + result.reused_prefix_tokens[i] = reused_pages * config.page_size_tokens; } - return pages; + + return result; } void HostPagedKVBackend::SharedState::ReleaseSequence( std::int64_t sequence_id) { - std::vector pages; - { - ScopedMutexLock lock(&header->sequence_mutex); - SequenceEntry* entry = FindSequenceEntryLocked(sequence_id); - if (entry == nullptr) { - throw std::out_of_range("Sequence ID " + - std::to_string(sequence_id) + - " not found during release"); - } - pages.reserve(entry->num_pages); - std::int32_t page = entry->head_page; - while (page != kInvalidPageIndex) { - pages.push_back(page); - const std::int32_t next = page_links[page]; - page_links[page] = kInvalidPageIndex; - page_owners[page] = kEmptySequenceId; - page = next; - } - entry->sequence_id = kTombstoneSequenceId; - entry->num_pages = 0; - entry->head_page = kInvalidPageIndex; - entry->tail_page = kInvalidPageIndex; - header->active_sequences.fetch_sub(1, std::memory_order_relaxed); - } - - if (!pages.empty()) { - ScopedMutexLock lock(&header->allocation_mutex); - std::uint32_t top = - header->free_stack_top.load(std::memory_order_relaxed); - for (std::int32_t page : pages) { - free_stack[top++] = page; - } - header->free_stack_top.store(top, std::memory_order_relaxed); + ScopedMutexLock lock(&header->metadata_mutex); + SequenceEntry* entry = FindSequenceEntryLocked(sequence_id); + if (entry == nullptr) { + throw std::out_of_range("Sequence ID " + std::to_string(sequence_id) + + " not found during release"); } + + ReleaseSequencePagesLocked(entry); + entry->sequence_id = kTombstoneSequenceId; + header->active_sequences.fetch_sub(1, std::memory_order_relaxed); } std::vector HostPagedKVBackend::SharedState::SequencePages( std::int64_t sequence_id, std::optional max_pages) const { - ScopedMutexLock lock(&header->sequence_mutex); + ScopedMutexLock lock(&header->metadata_mutex); SequenceEntry* entry = FindSequenceEntryLocked(sequence_id); if (entry == nullptr) { throw std::out_of_range("Sequence ID " + std::to_string(sequence_id) + " not found when fetching pages"); } - const std::size_t available_pages = entry->num_pages; - const std::size_t limit = max_pages.has_value() - ? std::min(max_pages.value(), available_pages) - : available_pages; - if (max_pages.has_value() && max_pages.value() > available_pages) { - throw std::out_of_range( - "Requested " + std::to_string(max_pages.value()) + - " pages but only " + std::to_string(available_pages) + - " pages allocated for sequence " + std::to_string(sequence_id)); + return CollectSequencePagesLocked(entry, max_pages); +} + +void HostPagedKVBackend::SharedState::CommitSequencePrefix( + std::int64_t sequence_id, const std::vector& prompt_tokens, + std::size_t prompt_token_count) { + if (!config.enable_prefix_reuse || prompt_tokens.empty()) { + return; } - std::vector pages; - pages.reserve(limit); - std::int32_t page = entry->head_page; - std::size_t count = 0; - while (page != kInvalidPageIndex && count < limit) { - pages.push_back(page); - page = page_links[page]; - ++count; + const std::size_t bounded_prompt_tokens = + std::min(prompt_token_count, prompt_tokens.size()); + const std::size_t full_prompt_pages = + bounded_prompt_tokens / config.page_size_tokens; + if (full_prompt_pages < config.prefix_min_store_pages) { + return; + } + + try { + ScopedMutexLock lock(&header->metadata_mutex); + SequenceEntry* entry = FindSequenceEntryLocked(sequence_id); + if (entry == nullptr || entry->num_pages == 0) { + return; + } + + const std::size_t storable_pages = + std::min(entry->num_pages, full_prompt_pages); + if (storable_pages < config.prefix_min_store_pages) { + return; + } + + const std::size_t token_count = storable_pages * config.page_size_tokens; + auto pages = CollectSequencePagesLocked(entry, storable_pages); + if (pages.size() != storable_pages) { + return; + } + prefix_cache_.CommitPrefixLocked(prompt_tokens.data(), token_count, pages); + } catch (const std::exception&) { + // Prefix commit is opportunistic. Ignore insertion failures and keep the + // inference flow unaffected. + return; } - return pages; } HostPagedKVStats HostPagedKVBackend::SharedState::CollectStats() const { HostPagedKVStats stats; stats.num_total_pages = config.num_pages; - const std::uint32_t free_count = + stats.num_free_pages = header->free_stack_top.load(std::memory_order_relaxed); - stats.num_free_pages = free_count; - stats.num_used_pages = config.num_pages - free_count; + stats.num_used_pages = stats.num_total_pages - stats.num_free_pages; stats.num_active_sequences = header->active_sequences.load(std::memory_order_relaxed); stats.sequence_table_capacity = sequence_capacity; stats.total_bytes = total_bytes; + + stats.num_prefix_entries = + header->prefix_entry_count.load(std::memory_order_relaxed); + stats.num_prefix_hits = + header->prefix_hit_count.load(std::memory_order_relaxed); + stats.num_prefix_misses = + header->prefix_miss_count.load(std::memory_order_relaxed); + stats.num_prefix_evictions = + header->prefix_evict_count.load(std::memory_order_relaxed); + stats.num_prefix_pinned_pages = + header->prefix_used_pages.load(std::memory_order_relaxed); + + std::size_t shared_pages = 0; + { + ScopedMutexLock lock(&header->metadata_mutex); + for (std::size_t i = 0; i < config.num_pages; ++i) { + if (page_refcount[i] > 1) { + ++shared_pages; + } + } + } + stats.num_shared_pages = shared_pages; + return stats; } @@ -724,76 +1157,40 @@ HostPagedKVBackend::AcquirePagesForSequences( throw std::invalid_argument( "sequence_ids and num_tokens must have the same length"); } - if (sequence_ids.empty()) { - return {}; - } + std::vector> allocations; + allocations.reserve(sequence_ids.size()); - auto allocate_one = [this](std::int64_t sequence_id, - std::size_t num_tokens_value) { - if (num_tokens_value == 0) { + for (std::size_t i = 0; i < sequence_ids.size(); ++i) { + if (num_tokens[i] == 0) { throw std::invalid_argument( "num_tokens must be greater than zero for sequence " + - std::to_string(sequence_id)); + std::to_string(sequence_ids[i])); } const std::size_t required_pages = - (num_tokens_value + config_.page_size_tokens - 1) / + (num_tokens[i] + config_.page_size_tokens - 1) / config_.page_size_tokens; - return state_->AcquirePages(sequence_id, required_pages); - }; - - if (sequence_ids.size() == 1 || SafeHardwareConcurrency() == 1) { - std::vector> allocations; - allocations.reserve(sequence_ids.size()); - for (std::size_t i = 0; i < sequence_ids.size(); ++i) { - allocations.emplace_back( - allocate_one(sequence_ids[i], num_tokens[i])); - } - return allocations; - } - - std::vector>> futures; - futures.reserve(sequence_ids.size()); - for (std::size_t i = 0; i < sequence_ids.size(); ++i) { - const std::int64_t sequence_id = sequence_ids[i]; - const std::size_t tokens = num_tokens[i]; - futures.emplace_back(std::async(std::launch::async, [=]() { - return allocate_one(sequence_id, tokens); - })); - } - - std::vector> allocations; - allocations.reserve(futures.size()); - for (auto& future : futures) { - allocations.emplace_back(future.get()); + allocations.emplace_back(state_->AcquirePages(sequence_ids[i], required_pages)); } return allocations; } +PrefixAllocationBatchResult HostPagedKVBackend::AcquirePagesForSequencesWithPrefix( + const std::vector& sequence_ids, + const std::vector& num_tokens, + const std::vector& flat_prompt_tokens, + const std::vector& prompt_offsets) { + return state_->AcquirePagesForSequencesWithPrefix( + sequence_ids, num_tokens, flat_prompt_tokens, prompt_offsets); +} + void HostPagedKVBackend::ReleaseSequence(std::int64_t sequence_id) { state_->ReleaseSequence(sequence_id); } void HostPagedKVBackend::ReleaseSequences( const std::vector& sequence_ids) { - if (sequence_ids.empty()) { - return; - } - if (sequence_ids.size() == 1 || SafeHardwareConcurrency() == 1) { - for (std::int64_t sequence_id : sequence_ids) { - state_->ReleaseSequence(sequence_id); - } - return; - } - std::vector> futures; - futures.reserve(sequence_ids.size()); for (std::int64_t sequence_id : sequence_ids) { - futures.emplace_back(std::async(std::launch::async, - [state = state_.get(), sequence_id]() { - state->ReleaseSequence(sequence_id); - })); - } - for (auto& future : futures) { - future.get(); + state_->ReleaseSequence(sequence_id); } } @@ -802,6 +1199,12 @@ std::vector HostPagedKVBackend::SequencePages( return state_->SequencePages(sequence_id, max_pages); } +void HostPagedKVBackend::CommitSequencePrefix( + std::int64_t sequence_id, const std::vector& prompt_tokens, + std::size_t prompt_token_count) { + state_->CommitSequencePrefix(sequence_id, prompt_tokens, prompt_token_count); +} + HostPagedKVStats HostPagedKVBackend::CollectStats() const { return state_->CollectStats(); } diff --git a/core/KV_Storage/host_paged_kv_backend.h b/core/KV_Storage/host_paged_kv_backend.h index 66627a405..58b779e51 100644 --- a/core/KV_Storage/host_paged_kv_backend.h +++ b/core/KV_Storage/host_paged_kv_backend.h @@ -3,6 +3,7 @@ #include #include +#include #include #include #include @@ -19,6 +20,12 @@ struct HostPagedKVStats { std::size_t num_active_sequences = 0; std::size_t sequence_table_capacity = 0; std::size_t total_bytes = 0; + std::size_t num_prefix_entries = 0; + std::size_t num_prefix_hits = 0; + std::size_t num_prefix_misses = 0; + std::size_t num_prefix_evictions = 0; + std::size_t num_prefix_pinned_pages = 0; + std::size_t num_shared_pages = 0; }; struct HostPagedKVConfig { @@ -34,6 +41,20 @@ struct HostPagedKVConfig { std::size_t v_element_size_bytes = 0; std::size_t sequence_table_capacity = 0; std::size_t alignment_bytes = 64; + bool enable_prefix_reuse = false; + std::size_t prefix_min_reuse_pages = 1; + std::size_t prefix_min_store_pages = 2; + std::size_t sequence_page_node_capacity = 0; + std::size_t radix_node_capacity = 0; + std::size_t radix_edge_capacity = 0; + std::size_t prefix_entry_capacity = 0; + std::size_t prefix_page_ref_capacity = 0; + std::size_t prefix_page_budget = 0; +}; + +struct PrefixAllocationBatchResult { + std::vector> allocated_pages; + std::vector reused_prefix_tokens; }; inline std::uint64_t HashCombine(std::uint64_t seed, std::uint64_t value) { @@ -91,6 +112,36 @@ inline HostPagedKVConfig SanitizeConfig(HostPagedKVConfig config) { config.v_element_size_bytes = config.k_element_size_bytes; } } + if (config.sequence_page_node_capacity == 0) { + config.sequence_page_node_capacity = std::max( + config.num_pages, static_cast(config.num_pages * 4)); + } + if (config.radix_node_capacity == 0) { + config.radix_node_capacity = + std::max(4096, config.num_pages / 2); + } + if (config.radix_edge_capacity == 0) { + config.radix_edge_capacity = + std::max(config.radix_node_capacity * 2, + config.radix_node_capacity + 1); + } + if (config.prefix_entry_capacity == 0) { + config.prefix_entry_capacity = + std::max(1024, config.num_pages / 64); + } + if (config.prefix_page_budget == 0) { + config.prefix_page_budget = std::max(64, config.num_pages / 4); + } + if (config.prefix_page_ref_capacity == 0) { + config.prefix_page_ref_capacity = std::max( + config.num_pages, config.prefix_entry_capacity * 8); + } + if (config.prefix_min_reuse_pages == 0) { + config.prefix_min_reuse_pages = 1; + } + if (config.prefix_min_store_pages == 0) { + config.prefix_min_store_pages = 1; + } return config; } @@ -120,7 +171,17 @@ inline std::string ToString(const HostPagedKVConfig& config) { << ", k_element_size_bytes=" << config.k_element_size_bytes << ", v_element_size_bytes=" << config.v_element_size_bytes << ", sequence_table_capacity=" << config.sequence_table_capacity - << ", alignment_bytes=" << config.alignment_bytes << ")"; + << ", alignment_bytes=" << config.alignment_bytes + << ", enable_prefix_reuse=" << config.enable_prefix_reuse + << ", prefix_min_reuse_pages=" << config.prefix_min_reuse_pages + << ", prefix_min_store_pages=" << config.prefix_min_store_pages + << ", sequence_page_node_capacity=" + << config.sequence_page_node_capacity + << ", radix_node_capacity=" << config.radix_node_capacity + << ", radix_edge_capacity=" << config.radix_edge_capacity + << ", prefix_entry_capacity=" << config.prefix_entry_capacity + << ", prefix_page_ref_capacity=" << config.prefix_page_ref_capacity + << ", prefix_page_budget=" << config.prefix_page_budget << ")"; return oss.str(); } @@ -131,7 +192,13 @@ inline std::string ToString(const HostPagedKVStats& stats) { << ", used_pages=" << stats.num_used_pages << ", active_sequences=" << stats.num_active_sequences << ", sequence_table_capacity=" << stats.sequence_table_capacity - << ", total_bytes=" << stats.total_bytes << ")"; + << ", total_bytes=" << stats.total_bytes + << ", prefix_entries=" << stats.num_prefix_entries + << ", prefix_hits=" << stats.num_prefix_hits + << ", prefix_misses=" << stats.num_prefix_misses + << ", prefix_evictions=" << stats.num_prefix_evictions + << ", prefix_pinned_pages=" << stats.num_prefix_pinned_pages + << ", shared_pages=" << stats.num_shared_pages << ")"; return oss.str(); } @@ -149,6 +216,15 @@ inline std::uint64_t HashHostKVConfig(const HostPagedKVConfig& config) { seed = HashCombine(seed, sanitized.v_element_size_bytes); seed = HashCombine(seed, sanitized.sequence_table_capacity); seed = HashCombine(seed, sanitized.alignment_bytes); + seed = HashCombine(seed, sanitized.enable_prefix_reuse ? 1ULL : 0ULL); + seed = HashCombine(seed, sanitized.prefix_min_reuse_pages); + seed = HashCombine(seed, sanitized.prefix_min_store_pages); + seed = HashCombine(seed, sanitized.sequence_page_node_capacity); + seed = HashCombine(seed, sanitized.radix_node_capacity); + seed = HashCombine(seed, sanitized.radix_edge_capacity); + seed = HashCombine(seed, sanitized.prefix_entry_capacity); + seed = HashCombine(seed, sanitized.prefix_page_ref_capacity); + seed = HashCombine(seed, sanitized.prefix_page_budget); return seed; } @@ -171,6 +247,12 @@ class HostPagedKVBackend { const std::vector& sequence_ids, const std::vector& num_tokens); + PrefixAllocationBatchResult AcquirePagesForSequencesWithPrefix( + const std::vector& sequence_ids, + const std::vector& num_tokens, + const std::vector& flat_prompt_tokens, + const std::vector& prompt_offsets); + void ReleaseSequence(std::int64_t sequence_id); void ReleaseSequences(const std::vector& sequence_ids); @@ -178,6 +260,10 @@ class HostPagedKVBackend { std::vector SequencePages( std::int64_t sequence_id, std::optional max_pages) const; + void CommitSequencePrefix(std::int64_t sequence_id, + const std::vector& prompt_tokens, + std::size_t prompt_token_count); + HostPagedKVStats CollectStats() const; std::byte* DataBase(); diff --git a/core/KV_Storage/host_paged_kv_config_utils.h b/core/KV_Storage/host_paged_kv_config_utils.h index d0f3397e5..63afddcbb 100644 --- a/core/KV_Storage/host_paged_kv_config_utils.h +++ b/core/KV_Storage/host_paged_kv_config_utils.h @@ -87,6 +87,15 @@ inline HostPagedKVConfig BuildHostPagedKVConfig( config.sequence_table_capacity = detail::DetermineSequenceTableCapacity( engine_config.kv_storage_config, config.num_pages); config.alignment_bytes = 64; + config.enable_prefix_reuse = external.enable_prefix_reuse; + config.prefix_min_reuse_pages = external.prefix_min_reuse_pages; + config.prefix_min_store_pages = external.prefix_min_store_pages; + config.sequence_page_node_capacity = external.sequence_page_node_capacity; + config.radix_node_capacity = external.radix_node_capacity; + config.radix_edge_capacity = external.radix_edge_capacity; + config.prefix_entry_capacity = external.prefix_entry_capacity; + config.prefix_page_ref_capacity = external.prefix_page_ref_capacity; + config.prefix_page_budget = external.prefix_page_budget; return config; } diff --git a/core/KV_Storage/host_paged_kv_prefix_cache.cpp b/core/KV_Storage/host_paged_kv_prefix_cache.cpp new file mode 100644 index 000000000..37fe16e6a --- /dev/null +++ b/core/KV_Storage/host_paged_kv_prefix_cache.cpp @@ -0,0 +1,593 @@ +#include "host_paged_kv_prefix_cache.h" + +#include +#include +#include +#include +#include + +namespace batchgen::kv { + +void HostKVPrefixCache::Bind(const HostKVPrefixCacheParams& params, + HostKVRadixNode* radix_nodes, + HostKVRadixEdge* radix_edges, + HostKVPrefixEntry* prefix_entries, + HostKVPrefixPageRef* prefix_page_refs, + std::int32_t* radix_node_free_stack, + std::int32_t* radix_edge_free_stack, + std::int32_t* prefix_entry_free_stack, + std::int32_t* prefix_page_ref_free_stack, + const SharedFields& shared_fields, + PageRefCallback increment_page_ref_cb, + PageRefCallback decrement_page_ref_cb) { + params_ = params; + shared_ = shared_fields; + + radix_nodes_ = radix_nodes; + radix_edges_ = radix_edges; + prefix_entries_ = prefix_entries; + prefix_page_refs_ = prefix_page_refs; + + radix_node_free_stack_ = radix_node_free_stack; + radix_edge_free_stack_ = radix_edge_free_stack; + prefix_entry_free_stack_ = prefix_entry_free_stack; + prefix_page_ref_free_stack_ = prefix_page_ref_free_stack; + + increment_page_ref_cb_ = std::move(increment_page_ref_cb); + decrement_page_ref_cb_ = std::move(decrement_page_ref_cb); +} + +void HostKVPrefixCache::InitializePools(std::size_t radix_node_capacity, + std::size_t radix_edge_capacity, + std::size_t prefix_entry_capacity, + std::size_t prefix_page_ref_capacity) { + for (std::size_t i = 0; i < radix_node_capacity; ++i) { + radix_nodes_[i] = HostKVRadixNode(); + } + if (radix_node_capacity > 0) { + radix_nodes_[0].parent_node = kHostKVInvalidIndex; + radix_nodes_[0].parent_edge = kHostKVInvalidIndex; + radix_nodes_[0].first_edge = kHostKVInvalidIndex; + radix_nodes_[0].terminal_entry = kHostKVInvalidIndex; + radix_nodes_[0].child_count = 0; + std::size_t cursor = 0; + for (std::size_t i = radix_node_capacity; i > 1; --i) { + radix_node_free_stack_[cursor++] = static_cast(i - 1); + } + } + + for (std::size_t i = 0; i < radix_edge_capacity; ++i) { + radix_edges_[i] = HostKVRadixEdge(); + radix_edge_free_stack_[i] = + static_cast(radix_edge_capacity - 1 - i); + } + + for (std::size_t i = 0; i < prefix_entry_capacity; ++i) { + prefix_entries_[i] = HostKVPrefixEntry(); + prefix_entry_free_stack_[i] = + static_cast(prefix_entry_capacity - 1 - i); + } + + for (std::size_t i = 0; i < prefix_page_ref_capacity; ++i) { + prefix_page_refs_[i] = HostKVPrefixPageRef(); + prefix_page_ref_free_stack_[i] = + static_cast(prefix_page_ref_capacity - 1 - i); + } +} + +std::int32_t HostKVPrefixCache::PopStackIndexLocked( + std::int32_t* stack, std::atomic* top, + const char* what) const { + const std::uint32_t current = top->load(std::memory_order_relaxed); + if (current == 0) { + throw std::runtime_error(std::string("Out of ") + what); + } + const std::uint32_t next = current - 1; + const std::int32_t value = stack[next]; + top->store(next, std::memory_order_relaxed); + return value; +} + +void HostKVPrefixCache::PushStackIndexLocked(std::int32_t* stack, + std::atomic* top, + std::int32_t value) const { + const std::uint32_t current = top->load(std::memory_order_relaxed); + stack[current] = value; + top->store(current + 1, std::memory_order_relaxed); +} + +std::int32_t HostKVPrefixCache::AllocateRadixNodeLocked(std::int32_t parent_node, + std::int32_t parent_edge) { + const std::int32_t node_idx = PopStackIndexLocked( + radix_node_free_stack_, shared_.radix_node_free_top, "radix nodes"); + radix_nodes_[node_idx] = HostKVRadixNode(); + radix_nodes_[node_idx].parent_node = parent_node; + radix_nodes_[node_idx].parent_edge = parent_edge; + return node_idx; +} + +void HostKVPrefixCache::FreeRadixNodeLocked(std::int32_t node_idx) { + radix_nodes_[node_idx] = HostKVRadixNode(); + PushStackIndexLocked(radix_node_free_stack_, shared_.radix_node_free_top, + node_idx); +} + +std::int32_t HostKVPrefixCache::AllocateRadixEdgeLocked( + std::int32_t child_node, std::int32_t next_sibling_edge, + const std::int32_t* label_tokens, std::size_t label_len) { + if (label_len == 0 || label_len > kHostKVRadixEdgeLabelChunk) { + throw std::invalid_argument("invalid radix edge label length"); + } + const std::int32_t edge_idx = PopStackIndexLocked( + radix_edge_free_stack_, shared_.radix_edge_free_top, "radix edges"); + radix_edges_[edge_idx] = HostKVRadixEdge(); + radix_edges_[edge_idx].child_node = child_node; + radix_edges_[edge_idx].next_sibling_edge = next_sibling_edge; + radix_edges_[edge_idx].label_len = static_cast(label_len); + std::memcpy(radix_edges_[edge_idx].label_tokens, label_tokens, + sizeof(std::int32_t) * label_len); + return edge_idx; +} + +void HostKVPrefixCache::FreeRadixEdgeLocked(std::int32_t edge_idx) { + radix_edges_[edge_idx] = HostKVRadixEdge(); + PushStackIndexLocked(radix_edge_free_stack_, shared_.radix_edge_free_top, + edge_idx); +} + +std::int32_t HostKVPrefixCache::AllocatePrefixEntryLocked() { + const std::int32_t entry_idx = PopStackIndexLocked( + prefix_entry_free_stack_, shared_.prefix_entry_free_top, + "prefix entries"); + prefix_entries_[entry_idx] = HostKVPrefixEntry(); + prefix_entries_[entry_idx].in_use = 1; + return entry_idx; +} + +void HostKVPrefixCache::FreePrefixEntryLocked(std::int32_t entry_idx) { + prefix_entries_[entry_idx] = HostKVPrefixEntry(); + PushStackIndexLocked(prefix_entry_free_stack_, shared_.prefix_entry_free_top, + entry_idx); +} + +std::int32_t HostKVPrefixCache::AllocatePrefixPageRefLocked(std::int32_t page_idx, + std::int32_t next) { + const std::int32_t ref_idx = PopStackIndexLocked( + prefix_page_ref_free_stack_, shared_.prefix_page_ref_free_top, + "prefix page refs"); + prefix_page_refs_[ref_idx].page_idx = page_idx; + prefix_page_refs_[ref_idx].next = next; + return ref_idx; +} + +void HostKVPrefixCache::FreePrefixPageRefLocked(std::int32_t page_ref_idx) { + prefix_page_refs_[page_ref_idx] = HostKVPrefixPageRef(); + PushStackIndexLocked(prefix_page_ref_free_stack_, + shared_.prefix_page_ref_free_top, page_ref_idx); +} + +std::int32_t HostKVPrefixCache::FindEdgeByFirstTokenLocked( + std::int32_t node_idx, std::int32_t first_token) const { + std::int32_t edge_idx = radix_nodes_[node_idx].first_edge; + while (edge_idx != kHostKVInvalidIndex) { + const HostKVRadixEdge& edge = radix_edges_[edge_idx]; + if (edge.label_len > 0 && edge.label_tokens[0] == first_token) { + return edge_idx; + } + edge_idx = edge.next_sibling_edge; + } + return kHostKVInvalidIndex; +} + +std::int32_t HostKVPrefixCache::AppendTokenPathLocked(std::int32_t start_node, + const std::int32_t* tokens, + std::size_t token_count) { + std::int32_t node_idx = start_node; + std::size_t pos = 0; + while (pos < token_count) { + const std::size_t chunk = + std::min(kHostKVRadixEdgeLabelChunk, token_count - pos); + const std::int32_t child_idx = + AllocateRadixNodeLocked(node_idx, kHostKVInvalidIndex); + const std::int32_t edge_idx = AllocateRadixEdgeLocked( + child_idx, radix_nodes_[node_idx].first_edge, tokens + pos, chunk); + radix_nodes_[node_idx].first_edge = edge_idx; + ++radix_nodes_[node_idx].child_count; + radix_nodes_[child_idx].parent_edge = edge_idx; + node_idx = child_idx; + pos += chunk; + } + return node_idx; +} + +std::int32_t HostKVPrefixCache::UpsertRadixPathLocked(const std::int32_t* tokens, + std::size_t token_count) { + std::int32_t node_idx = 0; + std::size_t pos = 0; + + while (pos < token_count) { + std::int32_t edge_idx = FindEdgeByFirstTokenLocked(node_idx, tokens[pos]); + if (edge_idx == kHostKVInvalidIndex) { + return AppendTokenPathLocked(node_idx, tokens + pos, token_count - pos); + } + + HostKVRadixEdge& edge = radix_edges_[edge_idx]; + const std::size_t remaining = token_count - pos; + const std::size_t compare_len = + std::min(edge.label_len, remaining); + std::size_t common = 0; + while (common < compare_len && + edge.label_tokens[common] == tokens[pos + common]) { + ++common; + } + + if (common == edge.label_len) { + pos += common; + node_idx = edge.child_node; + continue; + } + + if (common == 0) { + return AppendTokenPathLocked(node_idx, tokens + pos, token_count - pos); + } + + const std::int32_t old_child = edge.child_node; + const std::size_t old_suffix_len = edge.label_len - common; + std::int32_t old_suffix_tokens[kHostKVRadixEdgeLabelChunk] = {0}; + std::memcpy(old_suffix_tokens, edge.label_tokens + common, + sizeof(std::int32_t) * old_suffix_len); + + const std::int32_t split_node = AllocateRadixNodeLocked(node_idx, edge_idx); + const std::int32_t old_suffix_edge = AllocateRadixEdgeLocked( + old_child, kHostKVInvalidIndex, old_suffix_tokens, old_suffix_len); + + radix_nodes_[split_node].first_edge = old_suffix_edge; + radix_nodes_[split_node].child_count = 1; + + radix_nodes_[old_child].parent_node = split_node; + radix_nodes_[old_child].parent_edge = old_suffix_edge; + + edge.child_node = split_node; + edge.label_len = static_cast(common); + + pos += common; + if (pos == token_count) { + return split_node; + } + return AppendTokenPathLocked(split_node, tokens + pos, token_count - pos); + } + + return node_idx; +} + +std::vector HostKVPrefixCache::CollectPrefixEntryPagesLocked( + std::int32_t entry_idx, std::size_t max_pages) const { + std::vector pages; + if (entry_idx == kHostKVInvalidIndex || max_pages == 0) { + return pages; + } + + const HostKVPrefixEntry& entry = prefix_entries_[entry_idx]; + const std::size_t limit = std::min(entry.num_pages, max_pages); + pages.reserve(limit); + + std::int32_t ref_idx = entry.page_ref_head; + while (ref_idx != kHostKVInvalidIndex && pages.size() < limit) { + pages.push_back(prefix_page_refs_[ref_idx].page_idx); + ref_idx = prefix_page_refs_[ref_idx].next; + } + return pages; +} + +HostKVPrefixCache::LookupResult HostKVPrefixCache::LookupPrefixPagesLocked( + const std::int32_t* tokens, std::size_t token_count, std::size_t max_pages) { + LookupResult result; + if (!params_.enable_prefix_reuse || tokens == nullptr || token_count == 0 || + max_pages == 0) { + return result; + } + + auto consider_entry = [&](std::int32_t entry_idx) { + if (entry_idx == kHostKVInvalidIndex) { + return; + } + const HostKVPrefixEntry& entry = prefix_entries_[entry_idx]; + if (entry.in_use == 0) { + return; + } + const std::size_t candidate_pages = + std::min(entry.num_pages, max_pages); + if (candidate_pages < params_.prefix_min_reuse_pages || + candidate_pages <= result.reused_pages) { + return; + } + + auto pages = CollectPrefixEntryPagesLocked(entry_idx, candidate_pages); + if (pages.size() < candidate_pages) { + return; + } + + result.pages = std::move(pages); + result.reused_pages = candidate_pages; + result.entry_idx = entry_idx; + }; + + std::int32_t node_idx = 0; + std::size_t pos = 0; + consider_entry(radix_nodes_[node_idx].terminal_entry); + + while (pos < token_count) { + const std::int32_t edge_idx = + FindEdgeByFirstTokenLocked(node_idx, tokens[pos]); + if (edge_idx == kHostKVInvalidIndex) { + break; + } + const HostKVRadixEdge& edge = radix_edges_[edge_idx]; + const std::size_t remaining = token_count - pos; + const std::size_t compare_len = + std::min(edge.label_len, remaining); + + std::size_t common = 0; + while (common < compare_len && + edge.label_tokens[common] == tokens[pos + common]) { + ++common; + } + if (common != edge.label_len) { + break; + } + + pos += common; + node_idx = edge.child_node; + consider_entry(radix_nodes_[node_idx].terminal_entry); + } + + if (result.reused_pages >= params_.prefix_min_reuse_pages) { + shared_.prefix_hit_count->fetch_add(1, std::memory_order_relaxed); + TouchPrefixEntryLocked(result.entry_idx); + } else { + shared_.prefix_miss_count->fetch_add(1, std::memory_order_relaxed); + result = LookupResult(); + } + + return result; +} + +void HostKVPrefixCache::LruDetachLocked(std::int32_t entry_idx) { + HostKVPrefixEntry& entry = prefix_entries_[entry_idx]; + if (entry.lru_prev != kHostKVInvalidIndex) { + prefix_entries_[entry.lru_prev].lru_next = entry.lru_next; + } else if (*shared_.lru_head == entry_idx) { + *shared_.lru_head = entry.lru_next; + } + + if (entry.lru_next != kHostKVInvalidIndex) { + prefix_entries_[entry.lru_next].lru_prev = entry.lru_prev; + } else if (*shared_.lru_tail == entry_idx) { + *shared_.lru_tail = entry.lru_prev; + } + + entry.lru_prev = kHostKVInvalidIndex; + entry.lru_next = kHostKVInvalidIndex; +} + +void HostKVPrefixCache::LruAttachTailLocked(std::int32_t entry_idx) { + HostKVPrefixEntry& entry = prefix_entries_[entry_idx]; + entry.lru_prev = *shared_.lru_tail; + entry.lru_next = kHostKVInvalidIndex; + if (*shared_.lru_tail != kHostKVInvalidIndex) { + prefix_entries_[*shared_.lru_tail].lru_next = entry_idx; + } else { + *shared_.lru_head = entry_idx; + } + *shared_.lru_tail = entry_idx; +} + +void HostKVPrefixCache::TouchPrefixEntryLocked(std::int32_t entry_idx) { + if (entry_idx == kHostKVInvalidIndex || prefix_entries_[entry_idx].in_use == 0) { + return; + } + HostKVPrefixEntry& entry = prefix_entries_[entry_idx]; + entry.last_access_epoch = + shared_.prefix_access_epoch->fetch_add(1, std::memory_order_relaxed) + 1; + if (*shared_.lru_tail == entry_idx) { + return; + } + LruDetachLocked(entry_idx); + LruAttachTailLocked(entry_idx); +} + +void HostKVPrefixCache::PruneEmptyNodeChainLocked(std::int32_t node_idx) { + while (node_idx != 0 && node_idx != kHostKVInvalidIndex) { + HostKVRadixNode& node = radix_nodes_[node_idx]; + if (node.terminal_entry != kHostKVInvalidIndex || node.child_count != 0) { + break; + } + + const std::int32_t parent_idx = node.parent_node; + const std::int32_t parent_edge = node.parent_edge; + if (parent_idx == kHostKVInvalidIndex || + parent_edge == kHostKVInvalidIndex) { + break; + } + + std::int32_t prev_edge = kHostKVInvalidIndex; + std::int32_t cur_edge = radix_nodes_[parent_idx].first_edge; + while (cur_edge != kHostKVInvalidIndex && cur_edge != parent_edge) { + prev_edge = cur_edge; + cur_edge = radix_edges_[cur_edge].next_sibling_edge; + } + if (cur_edge == kHostKVInvalidIndex) { + break; + } + + const std::int32_t next_edge = radix_edges_[cur_edge].next_sibling_edge; + if (prev_edge == kHostKVInvalidIndex) { + radix_nodes_[parent_idx].first_edge = next_edge; + } else { + radix_edges_[prev_edge].next_sibling_edge = next_edge; + } + if (radix_nodes_[parent_idx].child_count > 0) { + --radix_nodes_[parent_idx].child_count; + } + + FreeRadixEdgeLocked(cur_edge); + const std::int32_t to_free = node_idx; + node_idx = parent_idx; + FreeRadixNodeLocked(to_free); + } +} + +bool HostKVPrefixCache::EvictOnePrefixEntryLocked() { + const std::int32_t entry_idx = *shared_.lru_head; + if (entry_idx == kHostKVInvalidIndex) { + return false; + } + + HostKVPrefixEntry& entry = prefix_entries_[entry_idx]; + if (entry.in_use == 0) { + LruDetachLocked(entry_idx); + FreePrefixEntryLocked(entry_idx); + return true; + } + + LruDetachLocked(entry_idx); + + if (entry.terminal_node != kHostKVInvalidIndex && + radix_nodes_[entry.terminal_node].terminal_entry == entry_idx) { + radix_nodes_[entry.terminal_node].terminal_entry = kHostKVInvalidIndex; + } + + std::int32_t ref_idx = entry.page_ref_head; + while (ref_idx != kHostKVInvalidIndex) { + const std::int32_t next = prefix_page_refs_[ref_idx].next; + const std::int32_t page_idx = prefix_page_refs_[ref_idx].page_idx; + decrement_page_ref_cb_(page_idx); + FreePrefixPageRefLocked(ref_idx); + ref_idx = next; + } + + if (shared_.prefix_entry_count->load(std::memory_order_relaxed) > 0) { + shared_.prefix_entry_count->fetch_sub(1, std::memory_order_relaxed); + } + const std::uint32_t used_pages = + shared_.prefix_used_pages->load(std::memory_order_relaxed); + shared_.prefix_used_pages->store( + used_pages > entry.num_pages ? used_pages - entry.num_pages : 0, + std::memory_order_relaxed); + shared_.prefix_evict_count->fetch_add(1, std::memory_order_relaxed); + + const std::int32_t terminal_node = entry.terminal_node; + FreePrefixEntryLocked(entry_idx); + if (terminal_node != kHostKVInvalidIndex) { + PruneEmptyNodeChainLocked(terminal_node); + } + + return true; +} + +bool HostKVPrefixCache::EnsurePrefixCapacityLocked( + std::size_t required_nodes, std::size_t required_edges, + std::size_t required_entries, std::size_t required_page_refs, + std::size_t extra_budget_pages) { + while (FreeRadixNodeCountLocked() < required_nodes || + FreeRadixEdgeCountLocked() < required_edges || + FreePrefixEntryCountLocked() < required_entries || + FreePrefixPageRefCountLocked() < required_page_refs || + shared_.prefix_used_pages->load(std::memory_order_relaxed) + + extra_budget_pages > + params_.prefix_page_budget) { + if (!EvictOnePrefixEntryLocked()) { + return false; + } + } + return true; +} + +std::size_t HostKVPrefixCache::FreeRadixNodeCountLocked() const { + return shared_.radix_node_free_top->load(std::memory_order_relaxed); +} + +std::size_t HostKVPrefixCache::FreeRadixEdgeCountLocked() const { + return shared_.radix_edge_free_top->load(std::memory_order_relaxed); +} + +std::size_t HostKVPrefixCache::FreePrefixEntryCountLocked() const { + return shared_.prefix_entry_free_top->load(std::memory_order_relaxed); +} + +std::size_t HostKVPrefixCache::FreePrefixPageRefCountLocked() const { + return shared_.prefix_page_ref_free_top->load(std::memory_order_relaxed); +} + +bool HostKVPrefixCache::CommitPrefixLocked(const std::int32_t* tokens, + std::size_t token_count, + const std::vector& pages) { + if (!params_.enable_prefix_reuse || tokens == nullptr || token_count == 0 || + pages.size() < params_.prefix_min_store_pages) { + return false; + } + + const std::size_t max_needed_nodes = + (token_count + kHostKVRadixEdgeLabelChunk - 1) / + kHostKVRadixEdgeLabelChunk + + 1; + const std::size_t max_needed_edges = max_needed_nodes + 1; + + if (!EnsurePrefixCapacityLocked(max_needed_nodes, max_needed_edges, 1, + pages.size(), pages.size())) { + return false; + } + + const std::int32_t terminal_node = UpsertRadixPathLocked(tokens, token_count); + if (terminal_node == kHostKVInvalidIndex) { + return false; + } + + if (radix_nodes_[terminal_node].terminal_entry != kHostKVInvalidIndex) { + TouchPrefixEntryLocked(radix_nodes_[terminal_node].terminal_entry); + return true; + } + + if (!EnsurePrefixCapacityLocked(0, 0, 1, pages.size(), pages.size())) { + return false; + } + + const std::int32_t entry_idx = AllocatePrefixEntryLocked(); + HostKVPrefixEntry& prefix_entry = prefix_entries_[entry_idx]; + prefix_entry.terminal_node = terminal_node; + prefix_entry.num_pages = static_cast(pages.size()); + prefix_entry.page_ref_head = kHostKVInvalidIndex; + prefix_entry.last_access_epoch = + shared_.prefix_access_epoch->fetch_add(1, std::memory_order_relaxed) + 1; + + std::int32_t tail_ref = kHostKVInvalidIndex; + for (std::size_t i = 0; i < pages.size(); ++i) { + const std::int32_t page_ref_idx = + AllocatePrefixPageRefLocked(pages[i], kHostKVInvalidIndex); + if (prefix_entry.page_ref_head == kHostKVInvalidIndex) { + prefix_entry.page_ref_head = page_ref_idx; + tail_ref = page_ref_idx; + } else { + prefix_page_refs_[tail_ref].next = page_ref_idx; + tail_ref = page_ref_idx; + } + increment_page_ref_cb_(pages[i]); + } + + radix_nodes_[terminal_node].terminal_entry = entry_idx; + LruAttachTailLocked(entry_idx); + shared_.prefix_entry_count->fetch_add(1, std::memory_order_relaxed); + shared_.prefix_used_pages->fetch_add(static_cast(pages.size()), + std::memory_order_relaxed); + + while (shared_.prefix_used_pages->load(std::memory_order_relaxed) > + params_.prefix_page_budget) { + if (!EvictOnePrefixEntryLocked()) { + break; + } + } + + return true; +} + +} // namespace batchgen::kv diff --git a/core/KV_Storage/host_paged_kv_prefix_cache.h b/core/KV_Storage/host_paged_kv_prefix_cache.h new file mode 100644 index 000000000..91ee08aaa --- /dev/null +++ b/core/KV_Storage/host_paged_kv_prefix_cache.h @@ -0,0 +1,175 @@ +#ifndef HOST_PAGED_KV_PREFIX_CACHE_H_ +#define HOST_PAGED_KV_PREFIX_CACHE_H_ + +#include +#include +#include +#include +#include + +namespace batchgen::kv { + +constexpr std::int32_t kHostKVInvalidIndex = -1; +constexpr std::size_t kHostKVRadixEdgeLabelChunk = 64; + +struct HostKVRadixNode { + std::int32_t parent_node = kHostKVInvalidIndex; + std::int32_t parent_edge = kHostKVInvalidIndex; + std::int32_t first_edge = kHostKVInvalidIndex; + std::int32_t terminal_entry = kHostKVInvalidIndex; + std::uint32_t child_count = 0; +}; + +struct HostKVRadixEdge { + std::int32_t child_node = kHostKVInvalidIndex; + std::int32_t next_sibling_edge = kHostKVInvalidIndex; + std::uint16_t label_len = 0; + std::int32_t label_tokens[kHostKVRadixEdgeLabelChunk] = {0}; +}; + +struct HostKVPrefixPageRef { + std::int32_t page_idx = kHostKVInvalidIndex; + std::int32_t next = kHostKVInvalidIndex; +}; + +struct HostKVPrefixEntry { + std::int32_t terminal_node = kHostKVInvalidIndex; + std::uint32_t num_pages = 0; + std::int32_t page_ref_head = kHostKVInvalidIndex; + std::int32_t lru_prev = kHostKVInvalidIndex; + std::int32_t lru_next = kHostKVInvalidIndex; + std::uint64_t last_access_epoch = 0; + std::uint8_t in_use = 0; +}; + +struct HostKVPrefixCacheParams { + bool enable_prefix_reuse = false; + std::size_t prefix_min_reuse_pages = 1; + std::size_t prefix_min_store_pages = 2; + std::size_t prefix_page_budget = 0; +}; + +class HostKVPrefixCache { + public: + struct SharedFields { + std::atomic* radix_node_free_top = nullptr; + std::atomic* radix_edge_free_top = nullptr; + std::atomic* prefix_entry_free_top = nullptr; + std::atomic* prefix_page_ref_free_top = nullptr; + + std::atomic* prefix_entry_count = nullptr; + std::atomic* prefix_used_pages = nullptr; + + std::atomic* prefix_access_epoch = nullptr; + std::atomic* prefix_hit_count = nullptr; + std::atomic* prefix_miss_count = nullptr; + std::atomic* prefix_evict_count = nullptr; + + std::int32_t* lru_head = nullptr; + std::int32_t* lru_tail = nullptr; + }; + + struct LookupResult { + std::vector pages; + std::size_t reused_pages = 0; + std::int32_t entry_idx = kHostKVInvalidIndex; + }; + + using PageRefCallback = std::function; + + HostKVPrefixCache() = default; + + void Bind(const HostKVPrefixCacheParams& params, HostKVRadixNode* radix_nodes, + HostKVRadixEdge* radix_edges, HostKVPrefixEntry* prefix_entries, + HostKVPrefixPageRef* prefix_page_refs, + std::int32_t* radix_node_free_stack, + std::int32_t* radix_edge_free_stack, + std::int32_t* prefix_entry_free_stack, + std::int32_t* prefix_page_ref_free_stack, + const SharedFields& shared_fields, + PageRefCallback increment_page_ref_cb, + PageRefCallback decrement_page_ref_cb); + + void InitializePools(std::size_t radix_node_capacity, + std::size_t radix_edge_capacity, + std::size_t prefix_entry_capacity, + std::size_t prefix_page_ref_capacity); + + LookupResult LookupPrefixPagesLocked(const std::int32_t* tokens, + std::size_t token_count, + std::size_t max_pages); + + bool CommitPrefixLocked(const std::int32_t* tokens, std::size_t token_count, + const std::vector& pages); + + private: + std::int32_t PopStackIndexLocked(std::int32_t* stack, + std::atomic* top, + const char* what) const; + void PushStackIndexLocked(std::int32_t* stack, + std::atomic* top, + std::int32_t value) const; + + std::int32_t AllocateRadixNodeLocked(std::int32_t parent_node, + std::int32_t parent_edge); + void FreeRadixNodeLocked(std::int32_t node_idx); + std::int32_t AllocateRadixEdgeLocked(std::int32_t child_node, + std::int32_t next_sibling_edge, + const std::int32_t* label_tokens, + std::size_t label_len); + void FreeRadixEdgeLocked(std::int32_t edge_idx); + + std::int32_t AllocatePrefixEntryLocked(); + void FreePrefixEntryLocked(std::int32_t entry_idx); + std::int32_t AllocatePrefixPageRefLocked(std::int32_t page_idx, + std::int32_t next); + void FreePrefixPageRefLocked(std::int32_t page_ref_idx); + + std::int32_t FindEdgeByFirstTokenLocked(std::int32_t node_idx, + std::int32_t first_token) const; + std::int32_t AppendTokenPathLocked(std::int32_t start_node, + const std::int32_t* tokens, + std::size_t token_count); + std::int32_t UpsertRadixPathLocked(const std::int32_t* tokens, + std::size_t token_count); + + std::vector CollectPrefixEntryPagesLocked( + std::int32_t entry_idx, std::size_t max_pages) const; + + void TouchPrefixEntryLocked(std::int32_t entry_idx); + void LruDetachLocked(std::int32_t entry_idx); + void LruAttachTailLocked(std::int32_t entry_idx); + + void PruneEmptyNodeChainLocked(std::int32_t node_idx); + bool EvictOnePrefixEntryLocked(); + bool EnsurePrefixCapacityLocked(std::size_t required_nodes, + std::size_t required_edges, + std::size_t required_entries, + std::size_t required_page_refs, + std::size_t extra_budget_pages); + + std::size_t FreeRadixNodeCountLocked() const; + std::size_t FreeRadixEdgeCountLocked() const; + std::size_t FreePrefixEntryCountLocked() const; + std::size_t FreePrefixPageRefCountLocked() const; + + HostKVPrefixCacheParams params_{}; + SharedFields shared_{}; + + HostKVRadixNode* radix_nodes_ = nullptr; + HostKVRadixEdge* radix_edges_ = nullptr; + HostKVPrefixEntry* prefix_entries_ = nullptr; + HostKVPrefixPageRef* prefix_page_refs_ = nullptr; + + std::int32_t* radix_node_free_stack_ = nullptr; + std::int32_t* radix_edge_free_stack_ = nullptr; + std::int32_t* prefix_entry_free_stack_ = nullptr; + std::int32_t* prefix_page_ref_free_stack_ = nullptr; + + PageRefCallback increment_page_ref_cb_; + PageRefCallback decrement_page_ref_cb_; +}; + +} // namespace batchgen::kv + +#endif // HOST_PAGED_KV_PREFIX_CACHE_H_ diff --git a/core/KV_Storage/host_paged_kv_worker_view.h b/core/KV_Storage/host_paged_kv_worker_view.h index 8c1006650..7f6fa0403 100644 --- a/core/KV_Storage/host_paged_kv_worker_view.h +++ b/core/KV_Storage/host_paged_kv_worker_view.h @@ -16,6 +16,7 @@ #include #include #include +#include #include #include #include @@ -264,6 +265,52 @@ class HostPagedKVWorkerView { return allocated_pages; } + std::pair>, std::vector> + AllocatePagesForSequencesWithPrefix( + const std::vector& sequence_ids, + const std::vector& num_tokens, + const std::vector& flat_prompt_tokens, + const std::vector& prompt_offsets) { + if (sequence_ids.size() != num_tokens.size()) { + throw std::invalid_argument( + "sequence_ids and num_tokens must have the same length"); + } + if (prompt_offsets.size() != sequence_ids.size() + 1) { + throw std::invalid_argument( + "prompt_offsets size must equal sequence_ids size + 1"); + } + if (prompt_offsets.empty() || prompt_offsets.front() != 0 || + prompt_offsets.back() != flat_prompt_tokens.size()) { + throw std::invalid_argument( + "prompt_offsets boundaries do not match flat_prompt_tokens"); + } + EnsureSequencesRegistered(sequence_ids); + for (std::size_t i = 0; i < num_tokens.size(); ++i) { + if (num_tokens[i] == 0) { + throw std::invalid_argument( + "num_tokens must be greater than zero"); + } + } + + auto alloc_result = backend_.AcquirePagesForSequencesWithPrefix( + sequence_ids, num_tokens, flat_prompt_tokens, prompt_offsets); + + for (std::size_t i = 0; i < sequence_ids.size(); ++i) { + AppendAllocatedPages(sequence_ids[i], alloc_result.allocated_pages[i]); + const std::size_t begin = prompt_offsets[i]; + const std::size_t end = prompt_offsets[i + 1]; + UpdatePrefixState(sequence_ids[i], + std::vector( + flat_prompt_tokens.begin() + begin, + flat_prompt_tokens.begin() + end), + end - begin, + alloc_result.reused_prefix_tokens[i]); + } + + return {std::move(alloc_result.allocated_pages), + std::move(alloc_result.reused_prefix_tokens)}; + } + std::vector GrowSequencePages( std::int64_t sequence_id, std::size_t num_pages) { if (num_pages == 0) { @@ -308,6 +355,7 @@ class HostPagedKVWorkerView { ResetCopyStreams(); UnregisterPinnedMemory(); page_table_.Clear(); + ClearPrefixStates(); } KVAsyncTask AsyncLoadLayerKVToDevice( @@ -837,6 +885,7 @@ class HostPagedKVWorkerView { void UnregisterSequence(std::int64_t sequence_id) { page_table_.Remove(sequence_id); + RemovePrefixState(sequence_id); } void UnregisterSequences(const std::vector& sequence_ids) { @@ -895,10 +944,16 @@ class HostPagedKVWorkerView { } } + std::vector prefix_skip_tokens(sequence_ids.size(), 0); + for (std::size_t i = 0; i < sequence_ids.size(); ++i) { + prefix_skip_tokens[i] = PrefixTokensToSkip(sequence_ids[i]); + } + return LaunchAsyncTask([this, layer_idx, sequence_ids = std::move(sequence_ids), sequence_lengths = std::move(sequence_lengths), - prepared_k, prepared_v, tokens_per_sequence]() { + prepared_k, prepared_v, tokens_per_sequence, + prefix_skip_tokens]() { c10::cuda::OptionalCUDAGuard device_guard(device_index_); const auto cuda_stream = CopyStream(CopyDirection::kDeviceToHost); @@ -933,13 +988,19 @@ class HostPagedKVWorkerView { if (tokens_to_copy == 0) { continue; } + const std::size_t skip_tokens = + std::min(prefix_skip_tokens[batch_idx], tokens_to_copy); + const std::size_t remaining_tokens = tokens_to_copy - skip_tokens; + if (remaining_tokens == 0) { + continue; + } geometry_.ValidatePageCapacity(pages, tokens_to_copy, "AsyncOffloadLayerKVToHost"); const auto* seq_k_src = k_base + batch_idx * k_seq_stride; ForEachPageChunk( - pages, 0, tokens_to_copy, + pages, skip_tokens, remaining_tokens, [&](std::int32_t page_idx, std::size_t page_offset_tokens, std::size_t chunk_tokens, std::size_t relative_token_offset) { @@ -947,7 +1008,9 @@ class HostPagedKVWorkerView { host_base, layer_idx, page_idx) + page_offset_tokens * k_token_bytes; const std::byte* src = - seq_k_src + relative_token_offset * k_token_bytes; + seq_k_src + + (skip_tokens + relative_token_offset) * + k_token_bytes; EnqueueCopy(src, dst, chunk_tokens * k_token_bytes, CopyDirection::kDeviceToHost, cuda_stream); }); @@ -956,7 +1019,7 @@ class HostPagedKVWorkerView { const auto* seq_v_src = v_base + batch_idx * v_seq_stride; ForEachPageChunk( - pages, 0, tokens_to_copy, + pages, skip_tokens, remaining_tokens, [&](std::int32_t page_idx, std::size_t page_offset_tokens, std::size_t chunk_tokens, @@ -967,7 +1030,8 @@ class HostPagedKVWorkerView { page_offset_tokens * v_token_bytes; const std::byte* src = seq_v_src + - relative_token_offset * v_token_bytes; + (skip_tokens + relative_token_offset) * + v_token_bytes; EnqueueCopy( src, dst, chunk_tokens * v_token_bytes, CopyDirection::kDeviceToHost, cuda_stream); @@ -977,6 +1041,7 @@ class HostPagedKVWorkerView { } this->SynchronizeWithEvent(cuda_stream); + this->MaybeCommitPrefix(layer_idx, sequence_ids); // LogFirstTokenPerPage(layer_idx, sequence_ids, sequence_lengths, // tokens_per_sequence, host_base); }); @@ -1231,6 +1296,93 @@ class HostPagedKVWorkerView { } } + struct PrefixSequenceState { + std::vector prompt_tokens; + std::size_t prompt_token_count = 0; + std::size_t reused_prefix_tokens = 0; + bool committed = false; + }; + + void UpdatePrefixState(std::int64_t sequence_id, + std::vector prompt_tokens, + std::size_t prompt_token_count, + std::size_t reused_prefix_tokens) { + if (!config_.enable_prefix_reuse) { + return; + } + std::lock_guard guard(prefix_state_mutex_); + auto& state = prefix_states_[sequence_id]; + state.prompt_tokens = std::move(prompt_tokens); + state.prompt_token_count = prompt_token_count; + state.reused_prefix_tokens = reused_prefix_tokens; + state.committed = false; + } + + std::size_t PrefixTokensToSkip(std::int64_t sequence_id) const { + if (!config_.enable_prefix_reuse) { + return 0; + } + std::lock_guard guard(prefix_state_mutex_); + auto it = prefix_states_.find(sequence_id); + if (it == prefix_states_.end()) { + return 0; + } + return it->second.reused_prefix_tokens; + } + + void MaybeCommitPrefix(std::size_t layer_idx, + const std::vector& sequence_ids) { + if (!config_.enable_prefix_reuse || layer_idx + 1 != config_.num_layers) { + return; + } + + struct PrefixCommitItem { + std::int64_t sequence_id = 0; + std::vector prompt_tokens; + std::size_t prompt_token_count = 0; + }; + + std::vector commit_items; + { + std::lock_guard guard(prefix_state_mutex_); + for (std::int64_t sequence_id : sequence_ids) { + auto it = prefix_states_.find(sequence_id); + if (it == prefix_states_.end()) { + continue; + } + if (it->second.committed) { + continue; + } + commit_items.push_back( + {sequence_id, it->second.prompt_tokens, + it->second.prompt_token_count}); + it->second.committed = true; + } + } + + for (const auto& item : commit_items) { + backend_.CommitSequencePrefix(item.sequence_id, item.prompt_tokens, + item.prompt_token_count); + } + } + + void RemovePrefixState(std::int64_t sequence_id) { + std::lock_guard guard(prefix_state_mutex_); + prefix_states_.erase(sequence_id); + } + + void RemovePrefixStates(const std::vector& sequence_ids) { + std::lock_guard guard(prefix_state_mutex_); + for (std::int64_t sequence_id : sequence_ids) { + prefix_states_.erase(sequence_id); + } + } + + void ClearPrefixStates() { + std::lock_guard guard(prefix_state_mutex_); + prefix_states_.clear(); + } + void AppendAllocatedPages(std::int64_t sequence_id, const std::vector& pages) { if (pages.empty()) { @@ -1957,6 +2109,8 @@ class HostPagedKVWorkerView { std::optional h2d_stream_; std::optional d2h_stream_; HostKVPageTable page_table_; + mutable std::mutex prefix_state_mutex_; + std::unordered_map prefix_states_; inline static std::atomic task_id_counter_{0}; }; diff --git a/core/batchgen_Binding.cpp b/core/batchgen_Binding.cpp index 2151754fd..d42dec1fd 100644 --- a/core/batchgen_Binding.cpp +++ b/core/batchgen_Binding.cpp @@ -18,27 +18,259 @@ * ---------------------------------------------------------------------------- */ // clang-format on -#include "KV_Storage/host_paged_kv_manager.h" -#include "KV_Storage/host_paged_kv_worker_view.h" -#include "batchgen.h" -#include "Weights_Storage/Weights_Storage.h" -#include "allocator.h" -#include "data_structures.h" -#include +#include +#include #include #include #include #include #include +#include +#include +#include + +#include #include #include #include +#include "KV_Storage/host_paged_kv_manager.h" +#include "KV_Storage/host_paged_kv_prefix_cache.h" +#include "KV_Storage/host_paged_kv_worker_view.h" +#include "Weights_Storage/Weights_Storage.h" +#include "allocator.h" +#include "batchgen.h" +#include "data_structures.h" + namespace py = pybind11; namespace kv = batchgen::kv; namespace { +struct HostKVPrefixCacheHarnessStats { + std::uint32_t prefix_entry_count = 0; + std::uint32_t prefix_used_pages = 0; + std::uint64_t prefix_access_epoch = 0; + std::uint64_t prefix_hit_count = 0; + std::uint64_t prefix_miss_count = 0; + std::uint64_t prefix_evict_count = 0; + std::int32_t lru_head = kv::kHostKVInvalidIndex; + std::int32_t lru_tail = kv::kHostKVInvalidIndex; +}; + +class HostKVPrefixCacheHarness { + public: + HostKVPrefixCacheHarness(std::size_t num_pages, + std::size_t radix_node_capacity, + std::size_t radix_edge_capacity, + std::size_t prefix_entry_capacity, + std::size_t prefix_page_ref_capacity, + bool enable_prefix_reuse = true, + std::size_t prefix_min_reuse_pages = 1, + std::size_t prefix_min_store_pages = 1, + std::size_t prefix_page_budget = 0) + : num_pages_(num_pages), + radix_node_capacity_(radix_node_capacity), + radix_edge_capacity_(radix_edge_capacity), + prefix_entry_capacity_(prefix_entry_capacity), + prefix_page_ref_capacity_(prefix_page_ref_capacity), + radix_nodes_(radix_node_capacity), + radix_edges_(radix_edge_capacity), + prefix_entries_(prefix_entry_capacity), + prefix_page_refs_(prefix_page_ref_capacity), + radix_node_free_stack_(radix_node_capacity), + radix_edge_free_stack_(radix_edge_capacity), + prefix_entry_free_stack_(prefix_entry_capacity), + prefix_page_ref_free_stack_(prefix_page_ref_capacity), + page_refcounts_(num_pages, 0) { + if (num_pages_ == 0) { + throw std::invalid_argument("num_pages must be > 0"); + } + if (radix_node_capacity_ == 0) { + throw std::invalid_argument("radix_node_capacity must be > 0"); + } + if (radix_edge_capacity_ == 0) { + throw std::invalid_argument("radix_edge_capacity must be > 0"); + } + if (prefix_entry_capacity_ == 0) { + throw std::invalid_argument("prefix_entry_capacity must be > 0"); + } + if (prefix_page_ref_capacity_ == 0) { + throw std::invalid_argument("prefix_page_ref_capacity must be > 0"); + } + + params_.enable_prefix_reuse = enable_prefix_reuse; + params_.prefix_min_reuse_pages = + std::max(1, prefix_min_reuse_pages); + params_.prefix_min_store_pages = + std::max(1, prefix_min_store_pages); + params_.prefix_page_budget = + prefix_page_budget == 0 ? num_pages_ : prefix_page_budget; + + kv::HostKVPrefixCache::SharedFields shared_fields; + shared_fields.radix_node_free_top = &radix_node_free_top_; + shared_fields.radix_edge_free_top = &radix_edge_free_top_; + shared_fields.prefix_entry_free_top = &prefix_entry_free_top_; + shared_fields.prefix_page_ref_free_top = &prefix_page_ref_free_top_; + shared_fields.prefix_entry_count = &prefix_entry_count_; + shared_fields.prefix_used_pages = &prefix_used_pages_; + shared_fields.prefix_access_epoch = &prefix_access_epoch_; + shared_fields.prefix_hit_count = &prefix_hit_count_; + shared_fields.prefix_miss_count = &prefix_miss_count_; + shared_fields.prefix_evict_count = &prefix_evict_count_; + shared_fields.lru_head = &lru_head_; + shared_fields.lru_tail = &lru_tail_; + + cache_.Bind( + params_, radix_nodes_.data(), radix_edges_.data(), + prefix_entries_.data(), prefix_page_refs_.data(), + radix_node_free_stack_.data(), radix_edge_free_stack_.data(), + prefix_entry_free_stack_.data(), prefix_page_ref_free_stack_.data(), + shared_fields, + [this](std::int32_t page_idx) { + ValidatePageIdx(page_idx); + ++page_refcounts_[page_idx]; + }, + [this](std::int32_t page_idx) { + ValidatePageIdx(page_idx); + if (page_refcounts_[page_idx] == 0) { + throw std::runtime_error( + "page_refcount underflow on page " + + std::to_string(page_idx)); + } + --page_refcounts_[page_idx]; + }); + + Reset(); + } + + std::pair, std::size_t> Lookup( + const std::vector& tokens, std::size_t max_pages) { + const auto result = + cache_.LookupPrefixPagesLocked(tokens.data(), tokens.size(), max_pages); + return {result.pages, result.reused_pages}; + } + + bool Commit(const std::vector& tokens, + const std::vector& pages) { + ValidatePageIndices(pages); + return cache_.CommitPrefixLocked(tokens.data(), tokens.size(), pages); + } + + HostKVPrefixCacheHarnessStats GetStats() const { + HostKVPrefixCacheHarnessStats stats; + stats.prefix_entry_count = + prefix_entry_count_.load(std::memory_order_relaxed); + stats.prefix_used_pages = + prefix_used_pages_.load(std::memory_order_relaxed); + stats.prefix_access_epoch = + prefix_access_epoch_.load(std::memory_order_relaxed); + stats.prefix_hit_count = prefix_hit_count_.load(std::memory_order_relaxed); + stats.prefix_miss_count = + prefix_miss_count_.load(std::memory_order_relaxed); + stats.prefix_evict_count = + prefix_evict_count_.load(std::memory_order_relaxed); + stats.lru_head = lru_head_; + stats.lru_tail = lru_tail_; + return stats; + } + + std::int32_t PageRefcount(std::int32_t page_idx) const { + ValidatePageIdx(page_idx); + return static_cast(page_refcounts_[page_idx]); + } + + std::vector PageRefcounts() const { + std::vector result; + result.reserve(page_refcounts_.size()); + for (std::uint32_t count : page_refcounts_) { + result.push_back(static_cast(count)); + } + return result; + } + + void Reset() { + std::fill(page_refcounts_.begin(), page_refcounts_.end(), 0); + + radix_node_free_top_.store( + static_cast(radix_node_capacity_ - 1), + std::memory_order_relaxed); + radix_edge_free_top_.store( + static_cast(radix_edge_capacity_), + std::memory_order_relaxed); + prefix_entry_free_top_.store( + static_cast(prefix_entry_capacity_), + std::memory_order_relaxed); + prefix_page_ref_free_top_.store( + static_cast(prefix_page_ref_capacity_), + std::memory_order_relaxed); + + prefix_entry_count_.store(0, std::memory_order_relaxed); + prefix_used_pages_.store(0, std::memory_order_relaxed); + prefix_access_epoch_.store(0, std::memory_order_relaxed); + prefix_hit_count_.store(0, std::memory_order_relaxed); + prefix_miss_count_.store(0, std::memory_order_relaxed); + prefix_evict_count_.store(0, std::memory_order_relaxed); + lru_head_ = kv::kHostKVInvalidIndex; + lru_tail_ = kv::kHostKVInvalidIndex; + + cache_.InitializePools(radix_node_capacity_, radix_edge_capacity_, + prefix_entry_capacity_, + prefix_page_ref_capacity_); + } + + private: + void ValidatePageIdx(std::int32_t page_idx) const { + if (page_idx < 0 || static_cast(page_idx) >= num_pages_) { + throw std::out_of_range("page index out of range: " + + std::to_string(page_idx)); + } + } + + void ValidatePageIndices(const std::vector& pages) const { + for (std::int32_t page_idx : pages) { + ValidatePageIdx(page_idx); + } + } + + std::size_t num_pages_ = 0; + std::size_t radix_node_capacity_ = 0; + std::size_t radix_edge_capacity_ = 0; + std::size_t prefix_entry_capacity_ = 0; + std::size_t prefix_page_ref_capacity_ = 0; + + kv::HostKVPrefixCache cache_; + kv::HostKVPrefixCacheParams params_{}; + + std::vector radix_nodes_; + std::vector radix_edges_; + std::vector prefix_entries_; + std::vector prefix_page_refs_; + + std::vector radix_node_free_stack_; + std::vector radix_edge_free_stack_; + std::vector prefix_entry_free_stack_; + std::vector prefix_page_ref_free_stack_; + + std::vector page_refcounts_; + + std::atomic radix_node_free_top_{0}; + std::atomic radix_edge_free_top_{0}; + std::atomic prefix_entry_free_top_{0}; + std::atomic prefix_page_ref_free_top_{0}; + + std::atomic prefix_entry_count_{0}; + std::atomic prefix_used_pages_{0}; + + std::atomic prefix_access_epoch_{0}; + std::atomic prefix_hit_count_{0}; + std::atomic prefix_miss_count_{0}; + std::atomic prefix_evict_count_{0}; + + std::int32_t lru_head_ = kv::kHostKVInvalidIndex; + std::int32_t lru_tail_ = kv::kHostKVInvalidIndex; +}; + template void BindHostPagedManager(py::module& m, const char* name) { py::class_(m, name) @@ -188,6 +420,29 @@ void BindHostPagedWorkerView(py::module& m, const char* name) { return self.AllocatePagesForSequences(sequence_ids, num_tokens); }) + .def( + "allocate_pages_for_sequences_with_prefix", + [](WorkerView& self, + const std::vector>& + requests, + const std::vector& flat_prompt_tokens, + const std::vector& prompt_offsets) { + std::vector sequence_ids; + std::vector num_tokens; + sequence_ids.reserve(requests.size()); + num_tokens.reserve(requests.size()); + for (const auto& request : requests) { + sequence_ids.push_back(request.first); + num_tokens.push_back(request.second); + } + auto result = self.AllocatePagesForSequencesWithPrefix( + sequence_ids, num_tokens, flat_prompt_tokens, + prompt_offsets); + return py::make_tuple(std::move(result.first), + std::move(result.second)); + }, + py::arg("requests"), py::arg("flat_prompt_tokens"), + py::arg("prompt_offsets")) .def("grow_sequence_pages", [](WorkerView& self, std::int64_t sequence_id, std::size_t num_pages) { @@ -237,7 +492,7 @@ void BindHostPagedWorkerView(py::module& m, const char* name) { return py::make_tuple(std::move(k_ptrs), v_ptrs); }, py::arg("sequence_id"), py::arg("layer_idx"), - py::arg("max_tokens") = py::none());; + py::arg("max_tokens") = py::none()); } } // namespace @@ -285,32 +540,28 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { &BatchGen::get_kv_scale, "Get the quantization scale for KV storage.") .def("get_past_key_states", - &BatchGen::get_past_key_states, - "Get the past key states for the given query global indices and max sequence length.") - .def("init_weight_storage", &BatchGen::init_weight_storage) - .def_property( - "host_paged_kv_worker_view", - &BatchGen::host_paged_kv_worker_view, - &BatchGen::set_host_paged_kv_worker_view, - "Reference to the bound HostPagedKVWorkerView instance.") - .def_property( - "gpu_paged_kv_manager", - &BatchGen::gpu_paged_kv_manager, - &BatchGen::set_gpu_paged_kv_manager, - "Python GPU paged KV manager bound to this engine."); - + &BatchGen::get_past_key_states, + "Get the past key states for the given query global indices and " + "max sequence length.") + .def("init_weight_storage", &BatchGen::init_weight_storage) + .def_property( + "host_paged_kv_worker_view", &BatchGen::host_paged_kv_worker_view, + &BatchGen::set_host_paged_kv_worker_view, + "Reference to the bound HostPagedKVWorkerView instance.") + .def_property("gpu_paged_kv_manager", &BatchGen::gpu_paged_kv_manager, + &BatchGen::set_gpu_paged_kv_manager, + "Python GPU paged KV manager bound to this engine."); + py::class_(m, "Weights_Storage") // Updated Constructor Binding - .def(py::init(), py::arg("device_id")) - + .def(py::init(), py::arg("device_id")) .def("Init", &Weights_Storage::Init, - py::arg("shm_name"), - py::arg("byte_size"), - py::arg("module_weights_shm"), - py::arg("enable_hugetlbfs") = false) + py::arg("shm_name"), py::arg("byte_size"), + py::arg("module_weights_shm"), + py::arg("enable_hugetlbfs") = false) .def("get_tensor", &Weights_Storage::get_tensor, - py::arg("module_key")); - + py::arg("module_key")); + py::class_(m, "HostPagedKVConfig") .def(py::init<>()) .def_readwrite("shm_name", &kv::HostPagedKVConfig::shm_name) @@ -330,6 +581,24 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { &kv::HostPagedKVConfig::sequence_table_capacity) .def_readwrite("alignment_bytes", &kv::HostPagedKVConfig::alignment_bytes) + .def_readwrite("enable_prefix_reuse", + &kv::HostPagedKVConfig::enable_prefix_reuse) + .def_readwrite("prefix_min_reuse_pages", + &kv::HostPagedKVConfig::prefix_min_reuse_pages) + .def_readwrite("prefix_min_store_pages", + &kv::HostPagedKVConfig::prefix_min_store_pages) + .def_readwrite("sequence_page_node_capacity", + &kv::HostPagedKVConfig::sequence_page_node_capacity) + .def_readwrite("radix_node_capacity", + &kv::HostPagedKVConfig::radix_node_capacity) + .def_readwrite("radix_edge_capacity", + &kv::HostPagedKVConfig::radix_edge_capacity) + .def_readwrite("prefix_entry_capacity", + &kv::HostPagedKVConfig::prefix_entry_capacity) + .def_readwrite("prefix_page_ref_capacity", + &kv::HostPagedKVConfig::prefix_page_ref_capacity) + .def_readwrite("prefix_page_budget", + &kv::HostPagedKVConfig::prefix_page_budget) .def("__repr__", [](const kv::HostPagedKVConfig& self) { return kv::ToString(self); @@ -345,11 +614,67 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { .def_readwrite("sequence_table_capacity", &kv::HostPagedKVStats::sequence_table_capacity) .def_readwrite("total_bytes", &kv::HostPagedKVStats::total_bytes) + .def_readwrite("num_prefix_entries", + &kv::HostPagedKVStats::num_prefix_entries) + .def_readwrite("num_prefix_hits", + &kv::HostPagedKVStats::num_prefix_hits) + .def_readwrite("num_prefix_misses", + &kv::HostPagedKVStats::num_prefix_misses) + .def_readwrite("num_prefix_evictions", + &kv::HostPagedKVStats::num_prefix_evictions) + .def_readwrite("num_prefix_pinned_pages", + &kv::HostPagedKVStats::num_prefix_pinned_pages) + .def_readwrite("num_shared_pages", + &kv::HostPagedKVStats::num_shared_pages) .def("__repr__", [](const kv::HostPagedKVStats& self) { return kv::ToString(self); }); + py::class_(m, "HostKVPrefixCacheHarnessStats") + .def(py::init<>()) + .def_readwrite("prefix_entry_count", + &HostKVPrefixCacheHarnessStats::prefix_entry_count) + .def_readwrite("prefix_used_pages", + &HostKVPrefixCacheHarnessStats::prefix_used_pages) + .def_readwrite("prefix_access_epoch", + &HostKVPrefixCacheHarnessStats::prefix_access_epoch) + .def_readwrite("prefix_hit_count", + &HostKVPrefixCacheHarnessStats::prefix_hit_count) + .def_readwrite("prefix_miss_count", + &HostKVPrefixCacheHarnessStats::prefix_miss_count) + .def_readwrite("prefix_evict_count", + &HostKVPrefixCacheHarnessStats::prefix_evict_count) + .def_readwrite("lru_head", &HostKVPrefixCacheHarnessStats::lru_head) + .def_readwrite("lru_tail", &HostKVPrefixCacheHarnessStats::lru_tail); + + py::class_(m, "HostKVPrefixCacheHarness") + .def(py::init(), + py::arg("num_pages"), py::arg("radix_node_capacity"), + py::arg("radix_edge_capacity"), py::arg("prefix_entry_capacity"), + py::arg("prefix_page_ref_capacity"), + py::arg("enable_prefix_reuse") = true, + py::arg("prefix_min_reuse_pages") = 1, + py::arg("prefix_min_store_pages") = 1, + py::arg("prefix_page_budget") = 0) + .def("lookup", + [](HostKVPrefixCacheHarness& self, + const std::vector& tokens, + std::size_t max_pages) { + auto result = self.Lookup(tokens, max_pages); + return py::make_tuple(std::move(result.first), result.second); + }, + py::arg("tokens"), py::arg("max_pages")) + .def("commit", &HostKVPrefixCacheHarness::Commit, py::arg("tokens"), + py::arg("pages")) + .def("get_stats", &HostKVPrefixCacheHarness::GetStats) + .def("page_refcount", &HostKVPrefixCacheHarness::PageRefcount, + py::arg("page_idx")) + .def("page_refcounts", &HostKVPrefixCacheHarness::PageRefcounts) + .def("reset", &HostKVPrefixCacheHarness::Reset); + py::class_(m, "KVAsyncTask") .def_property_readonly("id", &kv::KVAsyncTask::id) .def("wait", &kv::KVAsyncTask::wait) diff --git a/core/data_structures.h b/core/data_structures.h index 6a719d4b4..1e74bc094 100644 --- a/core/data_structures.h +++ b/core/data_structures.h @@ -125,6 +125,15 @@ struct HostPagedKVConfig { std::size_t num_v_heads = 0; std::size_t v_head_dim = 0; std::string kv_dtype = "bfloat16"; + bool enable_prefix_reuse = false; + std::size_t prefix_min_reuse_pages = 1; + std::size_t prefix_min_store_pages = 2; + std::size_t sequence_page_node_capacity = 0; + std::size_t radix_node_capacity = 0; + std::size_t radix_edge_capacity = 0; + std::size_t prefix_entry_capacity = 0; + std::size_t prefix_page_ref_capacity = 0; + std::size_t prefix_page_budget = 0; }; struct DevicePagedKVConfig { diff --git a/core/utils.cpp b/core/utils.cpp index 60455219d..d80254fe6 100644 --- a/core/utils.cpp +++ b/core/utils.cpp @@ -292,6 +292,50 @@ HostPagedKVConfig parse_host_paged_kv_config( cfg.attr("v_head_dim").cast(), "Host_Paged_KV_Config.v_head_dim"); config.kv_dtype = cfg.attr("kv_dtype").cast(); + if (py::hasattr(cfg, "enable_prefix_reuse")) { + config.enable_prefix_reuse = + cfg.attr("enable_prefix_reuse").cast(); + } + if (py::hasattr(cfg, "prefix_min_reuse_pages")) { + config.prefix_min_reuse_pages = CheckedSize( + cfg.attr("prefix_min_reuse_pages").cast(), + "Host_Paged_KV_Config.prefix_min_reuse_pages"); + } + if (py::hasattr(cfg, "prefix_min_store_pages")) { + config.prefix_min_store_pages = CheckedSize( + cfg.attr("prefix_min_store_pages").cast(), + "Host_Paged_KV_Config.prefix_min_store_pages"); + } + if (py::hasattr(cfg, "sequence_page_node_capacity")) { + config.sequence_page_node_capacity = CheckedSize( + cfg.attr("sequence_page_node_capacity").cast(), + "Host_Paged_KV_Config.sequence_page_node_capacity"); + } + if (py::hasattr(cfg, "radix_node_capacity")) { + config.radix_node_capacity = CheckedSize( + cfg.attr("radix_node_capacity").cast(), + "Host_Paged_KV_Config.radix_node_capacity"); + } + if (py::hasattr(cfg, "radix_edge_capacity")) { + config.radix_edge_capacity = CheckedSize( + cfg.attr("radix_edge_capacity").cast(), + "Host_Paged_KV_Config.radix_edge_capacity"); + } + if (py::hasattr(cfg, "prefix_entry_capacity")) { + config.prefix_entry_capacity = CheckedSize( + cfg.attr("prefix_entry_capacity").cast(), + "Host_Paged_KV_Config.prefix_entry_capacity"); + } + if (py::hasattr(cfg, "prefix_page_ref_capacity")) { + config.prefix_page_ref_capacity = CheckedSize( + cfg.attr("prefix_page_ref_capacity").cast(), + "Host_Paged_KV_Config.prefix_page_ref_capacity"); + } + if (py::hasattr(cfg, "prefix_page_budget")) { + config.prefix_page_budget = CheckedSize( + cfg.attr("prefix_page_budget").cast(), + "Host_Paged_KV_Config.prefix_page_budget"); + } std::cout << "Host Paged KV Config Parsed Successfully" << std::endl; return config; } diff --git a/op_builder/core_engine.py b/op_builder/core_engine.py index 9b37e8f78..c3c03fd5d 100644 --- a/op_builder/core_engine.py +++ b/op_builder/core_engine.py @@ -30,10 +30,14 @@ def sources(self): f"{BATCHGEN_CORE_ROOT}/GPU_KV_Buffer/GPU_KV_Buffer.cpp", f"{BATCHGEN_CORE_ROOT}/KV_Storage/host_paged_kv_manager.cpp", f"{BATCHGEN_CORE_ROOT}/KV_Storage/host_paged_kv_backend.cpp", + f"{BATCHGEN_CORE_ROOT}/KV_Storage/host_paged_kv_prefix_cache.cpp", f"{BATCHGEN_CORE_ROOT}/KV_Storage/host_paged_kv_worker_view.cpp", f"{BATCHGEN_CORE_ROOT}/KV_Storage/host_kv_page_table.cpp", f"{BATCHGEN_CORE_ROOT}/KV_Storage/uva_copy_kernel.cu", - f"{BATCHGEN_CORE_ROOT}/Hetero_Attn/CPU_Kernels/grouped_query_attention_cpu_avx2_omp.cpp", + ( + f"{BATCHGEN_CORE_ROOT}/Hetero_Attn/CPU_Kernels/" + "grouped_query_attention_cpu_avx2_omp.cpp" + ), f"{BATCHGEN_CORE_ROOT}/allocator.cpp", ] @@ -42,9 +46,9 @@ def include_paths(self): def cxx_args(self): """C++ compiler flags - DON'T call super() to avoid conflicts""" - CPU_ARCH = self.cpu_arch() - SIMD_WIDTH = self.simd_width() - + cpu_arch = self.cpu_arch() + simd_width = self.simd_width() + args = [ "-O2", "-std=c++17", # Must be C++17 for PyTorch @@ -56,8 +60,8 @@ def cxx_args(self): "-fprefetch-loop-arrays", "-fopenmp", "-Wno-reorder", - CPU_ARCH, - SIMD_WIDTH, + cpu_arch, + simd_width, "-D_GLIBCXX_USE_CXX11_ABI=1", ] return args @@ -84,4 +88,4 @@ def extra_ldflags(self): ] def is_compatible(self, verbose=True): - return super().is_compatible(verbose) \ No newline at end of file + return super().is_compatible(verbose) diff --git a/test/paged_kv/test_host_prefix_reuse.py b/test/paged_kv/test_host_prefix_reuse.py new file mode 100644 index 000000000..8afa691c2 --- /dev/null +++ b/test/paged_kv/test_host_prefix_reuse.py @@ -0,0 +1,144 @@ +import ctypes +import errno +import random +import string +from unittest import SkipTest + +import torch + +from batchgen.models.engine_loader import core_engine as bg + +_libc = ctypes.CDLL("libc.so.6", use_errno=True) + + +def _random_shm_name() -> str: + suffix = "".join(random.choices(string.ascii_lowercase + string.digits, k=10)) + return f"/batchgen_prefix_{suffix}" + + +def _shm_unlink(name: str) -> None: + if not name: + return + rc = _libc.shm_unlink(name.encode("utf-8")) + if rc != 0: + err = ctypes.get_errno() + if err != errno.ENOENT: + raise OSError(err, f"shm_unlink({name}) failed") + + +def _make_mla_config( + shm_name: str, enable_prefix_reuse: bool +) -> bg.HostPagedKVConfig: # type: ignore + cfg = bg.HostPagedKVConfig() + cfg.shm_name = shm_name + cfg.num_layers = 1 + cfg.num_pages = 512 + cfg.page_size_tokens = 64 + cfg.num_k_heads = 1 + cfg.k_head_dim = 576 + cfg.num_v_heads = 0 + cfg.v_head_dim = 0 + cfg.k_element_size_bytes = 2 + cfg.v_element_size_bytes = 0 + cfg.sequence_table_capacity = 1024 + cfg.alignment_bytes = 64 + + cfg.enable_prefix_reuse = enable_prefix_reuse + cfg.prefix_min_reuse_pages = 1 + cfg.prefix_min_store_pages = 1 + cfg.sequence_page_node_capacity = 2048 + cfg.radix_node_capacity = 1024 + cfg.radix_edge_capacity = 2048 + cfg.prefix_entry_capacity = 256 + cfg.prefix_page_ref_capacity = 1024 + cfg.prefix_page_budget = 128 + return cfg + + +def test_prefix_reuse_disabled_regression() -> None: + if not torch.cuda.is_available(): + raise SkipTest("CUDA is not available") + + shm_name = _random_shm_name() + cfg = _make_mla_config(shm_name, enable_prefix_reuse=False) + worker = bg.MLAHostPagedKVWorkerView(cfg) + + try: + worker.initialize(device_index=0, create_region=True) + worker.register_sequences([101]) + + prompt_tokens = list(range(64)) + pages, reused = worker.allocate_pages_for_sequences_with_prefix( + [(101, 64)], prompt_tokens, [0, 64] + ) + + assert len(pages) == 1 + assert len(pages[0]) == 1 + assert reused == [0] + + worker.release_sequence_pages([101]) + stats = worker.get_stats() + assert stats.num_used_pages == 0 + finally: + worker.shutdown() + del worker + _shm_unlink(shm_name) + + +def test_prefix_reuse_hits_after_commit() -> None: + if not torch.cuda.is_available(): + raise SkipTest("CUDA is not available") + + shm_name = _random_shm_name() + cfg = _make_mla_config(shm_name, enable_prefix_reuse=True) + worker = bg.MLAHostPagedKVWorkerView(cfg) + + try: + worker.initialize(device_index=0, create_region=True) + + prompt_tokens = [i % 32000 for i in range(64)] + + # First sequence: allocate and offload layer 0 to trigger prefix commit. + worker.register_sequences([201]) + pages_1, reused_1 = worker.allocate_pages_for_sequences_with_prefix( + [(201, 64)], prompt_tokens, [0, 64] + ) + assert reused_1 == [0] + assert len(pages_1[0]) == 1 + + k_tensor = torch.full( + (1, 64, cfg.num_k_heads, cfg.k_head_dim), + fill_value=1.0, + dtype=torch.bfloat16, + device="cuda:0", + ) + task = worker.async_offload_layer_kv_to_host( + layer_idx=0, + sequence_ids=[201], + k_tensor=k_tensor, + v_tensor=None, + sequence_lengths=[64], + ) + task.wait() + + worker.release_sequence_pages([201]) + + # Second sequence with same prompt should reuse full first page. + worker.register_sequences([202]) + pages_2, reused_2 = worker.allocate_pages_for_sequences_with_prefix( + [(202, 64)], prompt_tokens, [0, 64] + ) + + assert len(pages_2) == 1 + assert len(pages_2[0]) == 1 + assert reused_2 == [64] + + stats = worker.get_stats() + assert stats.num_prefix_hits >= 1 + assert stats.num_shared_pages >= 1 + + worker.release_sequence_pages([202]) + finally: + worker.shutdown() + del worker + _shm_unlink(shm_name) diff --git a/test/paged_kv/test_prefix_cache_binding.py b/test/paged_kv/test_prefix_cache_binding.py new file mode 100644 index 000000000..25fa3a599 --- /dev/null +++ b/test/paged_kv/test_prefix_cache_binding.py @@ -0,0 +1,100 @@ +from unittest import SkipTest + +from batchgen.models.engine_loader import core_engine as bg + + +def _make_prefix_cache( + *, + num_pages: int = 64, + radix_nodes: int = 256, + radix_edges: int = 512, + prefix_entries: int = 64, + prefix_page_refs: int = 256, + enable_prefix_reuse: bool = True, + prefix_min_reuse_pages: int = 1, + prefix_min_store_pages: int = 1, + prefix_page_budget: int = 32, +): + if not hasattr(bg, "HostKVPrefixCacheHarness"): + raise SkipTest("HostKVPrefixCacheHarness binding is not available") + return bg.HostKVPrefixCacheHarness( + num_pages, + radix_nodes, + radix_edges, + prefix_entries, + prefix_page_refs, + enable_prefix_reuse, + prefix_min_reuse_pages, + prefix_min_store_pages, + prefix_page_budget, + ) + + +def test_prefix_cache_harness_commit_and_lookup_hit() -> None: + cache = _make_prefix_cache() + tokens = list(range(128)) + + pages, reused = cache.lookup(tokens, 2) + assert pages == [] + assert reused == 0 + assert cache.get_stats().prefix_miss_count == 1 + + assert cache.commit(tokens, [3, 4]) is True + pages, reused = cache.lookup(tokens, 2) + assert pages == [3, 4] + assert reused == 2 + + stats = cache.get_stats() + assert stats.prefix_entry_count == 1 + assert stats.prefix_used_pages == 2 + assert stats.prefix_hit_count >= 1 + assert cache.page_refcount(3) == 1 + assert cache.page_refcount(4) == 1 + + +def test_prefix_cache_harness_respects_min_store_and_reuse_pages() -> None: + cache = _make_prefix_cache(prefix_min_reuse_pages=2, prefix_min_store_pages=2) + + one_page_tokens = list(range(64)) + assert cache.commit(one_page_tokens, [5]) is False + assert cache.get_stats().prefix_entry_count == 0 + + two_page_tokens = list(range(128)) + assert cache.commit(two_page_tokens, [5, 6]) is True + + pages, reused = cache.lookup(two_page_tokens, 1) + assert pages == [] + assert reused == 0 + + pages, reused = cache.lookup(two_page_tokens, 2) + assert pages == [5, 6] + assert reused == 2 + + +def test_prefix_cache_harness_lru_evict_releases_page_refs() -> None: + cache = _make_prefix_cache(num_pages=8, prefix_page_budget=2) + tokens_a = [1000 + i for i in range(128)] + tokens_b = [2000 + i for i in range(128)] + + assert cache.commit(tokens_a, [0, 1]) is True + assert cache.page_refcount(0) == 1 + assert cache.page_refcount(1) == 1 + + assert cache.commit(tokens_b, [2, 3]) is True + stats = cache.get_stats() + assert stats.prefix_evict_count >= 1 + assert stats.prefix_entry_count == 1 + assert stats.prefix_used_pages <= 2 + + assert cache.page_refcount(0) == 0 + assert cache.page_refcount(1) == 0 + assert cache.page_refcount(2) == 1 + assert cache.page_refcount(3) == 1 + + pages, reused = cache.lookup(tokens_a, 2) + assert pages == [] + assert reused == 0 + + pages, reused = cache.lookup(tokens_b, 2) + assert pages == [2, 3] + assert reused == 2 From b4a2ae3dfaeecfbcf4aced0f1ac064ddb630f0af Mon Sep 17 00:00:00 2001 From: luzhan <513964121@qq.com> Date: Mon, 2 Mar 2026 21:10:49 +0000 Subject: [PATCH 02/17] feat: enhance prefix cache with decode token support and capacity defaults --- batchgen/batchgen_worker.py | 27 ++- batchgen/config/config.py | 30 +++ batchgen/config/engine_config_parser.py | 1 + batchgen/kv_cache/host_kv_mananger_config.py | 40 ++-- core/KV_Storage/host_paged_kv_backend.cpp | 204 ++++++++++++------ core/KV_Storage/host_paged_kv_backend.h | 7 +- .../KV_Storage/host_paged_kv_prefix_cache.cpp | 48 +++++ core/KV_Storage/host_paged_kv_prefix_cache.h | 2 + core/KV_Storage/host_paged_kv_worker_view.h | 78 +++++-- core/batchgen_Binding.cpp | 7 +- test/paged_kv/test_host_prefix_reuse.py | 113 +++++++++- test/paged_kv/test_prefix_cache_binding.py | 37 ++++ 12 files changed, 492 insertions(+), 102 deletions(-) diff --git a/batchgen/batchgen_worker.py b/batchgen/batchgen_worker.py index 4e6ae25c3..fce3a9294 100644 --- a/batchgen/batchgen_worker.py +++ b/batchgen/batchgen_worker.py @@ -976,6 +976,8 @@ def _append_decode_kv_to_host_async( # Build sequence info sequence_ids = [] sequence_lengths = [] + decode_token_ids = [] + has_decode_tokens = True # DIAGNOSTIC: Track host KV append positions for debugging append_diag = [] @@ -987,6 +989,14 @@ def _append_decode_kv_to_host_async( # Write position is current position (0-indexed) write_pos = seq.current_context_length - 1 sequence_lengths.append(write_pos) + if ( + seq.input_ids is None + or write_pos < 0 + or write_pos >= seq.input_ids.shape[1] + ): + has_decode_tokens = False + else: + decode_token_ids.append(int(seq.input_ids[0, write_pos].item())) # Track for debugging (only first few sequences) if len(append_diag) < 3 and seq.decoded_length > 1: @@ -1039,12 +1049,14 @@ def _append_decode_kv_to_host_async( self._kv_offload_event.synchronize() # Launch async append + decode_token_ids_arg = decode_token_ids if has_decode_tokens else None task = worker_view.async_append_decode_kv_to_host( layer_idx=layer_idx, sequence_ids=sequence_ids, k_tensor=k_tensor, v_tensor=v_tensor, # GQA models (GPT-OSS) have separate V; MLA models pass None sequence_lengths=sequence_lengths, + decode_token_ids=decode_token_ids_arg, ) # CRITICAL FIX: Store tensor references alongside task to prevent GC @@ -6761,12 +6773,23 @@ def _append_decode_kv_to_host_async( sequence_ids = [] sequence_lengths = [] + decode_token_ids = [] + has_decode_tokens = True for local_idx in batch: uuid = self._local_to_uuid_map[local_idx] seq = self.global_batch.get_sequence(uuid) sequence_ids.append(seq.global_idx) - sequence_lengths.append(seq.current_context_length - 1) + write_pos = seq.current_context_length - 1 + sequence_lengths.append(write_pos) + if ( + seq.input_ids is None + or write_pos < 0 + or write_pos >= seq.input_ids.shape[1] + ): + has_decode_tokens = False + else: + decode_token_ids.append(int(seq.input_ids[0, write_pos].item())) if k_tensor.dim() == 3: k_tensor = k_tensor.unsqueeze(2) @@ -6805,12 +6828,14 @@ def _append_decode_kv_to_host_async( self._kv_offload_event.record(torch.cuda.current_stream(self.torch_device)) self._kv_offload_event.synchronize() + decode_token_ids_arg = decode_token_ids if has_decode_tokens else None task = worker_view.async_append_decode_kv_to_host( layer_idx=layer_idx, sequence_ids=sequence_ids, k_tensor=k_tensor, v_tensor=v_tensor, # GQA models (GPT-OSS) have separate V; MLA models pass None sequence_lengths=sequence_lengths, + decode_token_ids=decode_token_ids_arg, ) # CRITICAL FIX: Store tensor references alongside task to prevent GC diff --git a/batchgen/config/config.py b/batchgen/config/config.py index cffd79739..a18d838a4 100644 --- a/batchgen/config/config.py +++ b/batchgen/config/config.py @@ -179,6 +179,8 @@ class HostPagedKVConfig: enable_prefix_reuse: bool = False prefix_min_reuse_pages: int = 1 prefix_min_store_pages: int = 2 + # Capacity knobs for prefix cache internals. + # 0 means "auto": values are derived from num_pages_per_layer at runtime. sequence_page_node_capacity: int = 0 radix_node_capacity: int = 0 radix_edge_capacity: int = 0 @@ -186,6 +188,34 @@ class HostPagedKVConfig: prefix_page_ref_capacity: int = 0 prefix_page_budget: int = 0 + def __post_init__(self) -> None: + self.apply_runtime_defaults() + + def apply_runtime_defaults(self) -> None: + """Fill auto (0) capacity fields using num_pages_per_layer.""" + num_pages = max(int(self.num_pages_per_layer), 0) + if num_pages == 0: + return + + if self.sequence_page_node_capacity <= 0: + self.sequence_page_node_capacity = max(num_pages, num_pages * 4) + if self.radix_node_capacity <= 0: + self.radix_node_capacity = max(4096, num_pages // 2) + if self.radix_edge_capacity <= 0: + self.radix_edge_capacity = max( + self.radix_node_capacity * 2, + self.radix_node_capacity + 1, + ) + if self.prefix_entry_capacity <= 0: + self.prefix_entry_capacity = max(1024, num_pages // 64) + if self.prefix_page_budget <= 0: + self.prefix_page_budget = max(128, num_pages // 2) + if self.prefix_page_ref_capacity <= 0: + self.prefix_page_ref_capacity = max( + num_pages, + self.prefix_entry_capacity * 8, + ) + @dataclass class DevicePagedKVConfig: num_layers: int = 0 diff --git a/batchgen/config/engine_config_parser.py b/batchgen/config/engine_config_parser.py index ea55e2dc9..2bc041498 100644 --- a/batchgen/config/engine_config_parser.py +++ b/batchgen/config/engine_config_parser.py @@ -209,6 +209,7 @@ def _parse_host_paged_kv_config( if key not in valid_fields: raise ValueError(f"Unknown key in Host_Paged_KV_Config: {key}") setattr(host_config, key, value) + host_config.apply_runtime_defaults() def _parse_device_paged_kv_config( diff --git a/batchgen/kv_cache/host_kv_mananger_config.py b/batchgen/kv_cache/host_kv_mananger_config.py index 9e5cbde2f..fe0f23fb5 100644 --- a/batchgen/kv_cache/host_kv_mananger_config.py +++ b/batchgen/kv_cache/host_kv_mananger_config.py @@ -177,6 +177,32 @@ def _resolve_profile(model_name: str) -> _HostKVModelProfile: return _PROFILE_REGISTRY[_PROFILE_ALIASES[alias]] +def _apply_capacity_defaults(config: Any) -> None: + """Populate HostPagedKVConfig capacity fields when they are zero.""" + num_pages = int(config.num_pages) + if num_pages <= 0: + return + + if int(config.sequence_page_node_capacity) <= 0: + config.sequence_page_node_capacity = max(num_pages, num_pages * 4) + if int(config.radix_node_capacity) <= 0: + config.radix_node_capacity = max(4096, num_pages // 2) + if int(config.radix_edge_capacity) <= 0: + config.radix_edge_capacity = max( + int(config.radix_node_capacity) * 2, + int(config.radix_node_capacity) + 1, + ) + if int(config.prefix_entry_capacity) <= 0: + config.prefix_entry_capacity = max(1024, num_pages // 64) + if int(config.prefix_page_budget) <= 0: + config.prefix_page_budget = max(128, num_pages // 2) + if int(config.prefix_page_ref_capacity) <= 0: + config.prefix_page_ref_capacity = max( + num_pages, + int(config.prefix_entry_capacity) * 8, + ) + + def build_host_kv_config(model_name: str, host_kv_cache_size: int) -> Any: """Builds a core HostPagedKVConfig for the given model and host budget.""" @@ -224,12 +250,7 @@ def build_host_kv_config(model_name: str, host_kv_cache_size: int) -> Any: config.enable_prefix_reuse = False config.prefix_min_reuse_pages = 1 config.prefix_min_store_pages = 2 - config.sequence_page_node_capacity = 0 - config.radix_node_capacity = 0 - config.radix_edge_capacity = 0 - config.prefix_entry_capacity = 0 - config.prefix_page_ref_capacity = 0 - config.prefix_page_budget = 0 + _apply_capacity_defaults(config) return config @@ -348,12 +369,7 @@ def build_host_kv_config_aux(model_name: str, host_kv_cache_size: int) -> Any | config.enable_prefix_reuse = False config.prefix_min_reuse_pages = 1 config.prefix_min_store_pages = 2 - config.sequence_page_node_capacity = 0 - config.radix_node_capacity = 0 - config.radix_edge_capacity = 0 - config.prefix_entry_capacity = 0 - config.prefix_page_ref_capacity = 0 - config.prefix_page_budget = 0 + _apply_capacity_defaults(config) return config diff --git a/core/KV_Storage/host_paged_kv_backend.cpp b/core/KV_Storage/host_paged_kv_backend.cpp index b161bf88d..dc88d09b7 100644 --- a/core/KV_Storage/host_paged_kv_backend.cpp +++ b/core/KV_Storage/host_paged_kv_backend.cpp @@ -285,7 +285,8 @@ struct HostPagedKVBackend::SharedState { std::size_t HashSequenceId(std::int64_t sequence_id) const; SequenceEntry* FindSequenceEntryLocked(std::int64_t sequence_id) const; SequenceEntry* FindOrInsertSequenceEntryLocked(std::int64_t sequence_id, - bool* is_new); + bool* is_new, + std::int64_t* previous_marker); std::int32_t PopStackIndexLocked(std::int32_t* stack, std::atomic* top, @@ -633,7 +634,7 @@ SequenceEntry* HostPagedKVBackend::SharedState::FindSequenceEntryLocked( } SequenceEntry* HostPagedKVBackend::SharedState::FindOrInsertSequenceEntryLocked( - std::int64_t sequence_id, bool* is_new) { + std::int64_t sequence_id, bool* is_new, std::int64_t* previous_marker) { std::size_t index = HashSequenceId(sequence_id); SequenceEntry* first_tombstone = nullptr; for (std::size_t probe = 0; probe < sequence_capacity; ++probe) { @@ -642,11 +643,17 @@ SequenceEntry* HostPagedKVBackend::SharedState::FindOrInsertSequenceEntryLocked( if (is_new != nullptr) { *is_new = false; } + if (previous_marker != nullptr) { + *previous_marker = sequence_id; + } return entry; } if (entry->sequence_id == kEmptySequenceId) { SequenceEntry* target = first_tombstone != nullptr ? first_tombstone : entry; + if (previous_marker != nullptr) { + *previous_marker = target->sequence_id; + } *target = SequenceEntry(); target->sequence_id = sequence_id; if (is_new != nullptr) { @@ -905,7 +912,8 @@ std::vector HostPagedKVBackend::SharedState::AcquirePages( } bool is_new = false; - SequenceEntry* entry = FindOrInsertSequenceEntryLocked(sequence_id, &is_new); + SequenceEntry* entry = + FindOrInsertSequenceEntryLocked(sequence_id, &is_new, nullptr); if (is_new) { header->active_sequences.fetch_add(1, std::memory_order_relaxed); } @@ -943,74 +951,148 @@ HostPagedKVBackend::SharedState::AcquirePagesForSequencesWithPrefix( result.allocated_pages.resize(sequence_ids.size()); result.reused_prefix_tokens.resize(sequence_ids.size(), 0); - for (std::size_t i = 0; i < sequence_ids.size(); ++i) { - const std::size_t begin = prompt_offsets[i]; - const std::size_t end = prompt_offsets[i + 1]; - if (begin > end || end > flat_prompt_tokens.size()) { - throw std::invalid_argument("prompt_offsets contain invalid slice"); - } - if (num_tokens[i] == 0) { - throw std::invalid_argument("num_tokens entries must be > 0"); - } + struct SequenceRollbackRecord { + SequenceEntry* entry = nullptr; + std::int64_t previous_marker = kEmptySequenceId; + std::uint32_t previous_num_pages = 0; + std::int32_t previous_head_node = kInvalidIndex; + std::int32_t previous_tail_node = kInvalidIndex; + bool was_new = false; + }; - const std::size_t required_pages = - (num_tokens[i] + config.page_size_tokens - 1) / config.page_size_tokens; - const std::size_t prompt_tokens = end - begin; - const std::size_t prompt_full_pages = - std::min(required_pages, prompt_tokens / config.page_size_tokens); - - if (!config.enable_prefix_reuse || prompt_full_pages == 0) { - result.allocated_pages[i] = AcquirePages(sequence_ids[i], required_pages); - continue; - } + ScopedMutexLock lock(&header->metadata_mutex); + std::vector rollback_records; + rollback_records.reserve(sequence_ids.size()); - std::vector pages; - pages.reserve(required_pages); + try { + for (std::size_t i = 0; i < sequence_ids.size(); ++i) { + const std::size_t begin = prompt_offsets[i]; + const std::size_t end = prompt_offsets[i + 1]; + if (begin > end || end > flat_prompt_tokens.size()) { + throw std::invalid_argument("prompt_offsets contain invalid slice"); + } + if (num_tokens[i] == 0) { + throw std::invalid_argument("num_tokens entries must be > 0"); + } - ScopedMutexLock lock(&header->metadata_mutex); + const std::size_t required_pages = + (num_tokens[i] + config.page_size_tokens - 1) / + config.page_size_tokens; + const std::size_t prompt_tokens = end - begin; + const std::size_t prompt_full_pages = + std::min(required_pages, prompt_tokens / config.page_size_tokens); + + std::size_t reused_pages = 0; + std::vector reused_page_candidates; + if (config.enable_prefix_reuse && prompt_full_pages > 0) { + auto lookup = prefix_cache_.LookupPrefixPagesLocked( + flat_prompt_tokens.data() + begin, + prompt_full_pages * config.page_size_tokens, + prompt_full_pages); + reused_pages = + std::min(lookup.reused_pages, prompt_full_pages); + reused_page_candidates = std::move(lookup.pages); + if (reused_page_candidates.size() < reused_pages) { + throw std::runtime_error( + "Prefix lookup returned fewer pages than expected"); + } + } - const auto lookup = prefix_cache_.LookupPrefixPagesLocked( - flat_prompt_tokens.data() + begin, - prompt_full_pages * config.page_size_tokens, prompt_full_pages); + const std::size_t new_pages = required_pages - reused_pages; + if (header->free_stack_top.load(std::memory_order_relaxed) < + new_pages) { + throw std::runtime_error( + "Insufficient free pages for prefix-aware allocation of " + "sequence " + + std::to_string(sequence_ids[i])); + } + if (header->seq_page_node_free_top.load(std::memory_order_relaxed) < + required_pages) { + throw std::runtime_error( + "Insufficient sequence page nodes for prefix-aware allocation"); + } - const std::size_t reused_pages = std::min( - lookup.reused_pages, prompt_full_pages); - const std::size_t new_pages = required_pages - reused_pages; + bool is_new = false; + std::int64_t previous_marker = kEmptySequenceId; + SequenceEntry* entry = FindOrInsertSequenceEntryLocked( + sequence_ids[i], &is_new, &previous_marker); + + SequenceRollbackRecord rollback_record; + rollback_record.entry = entry; + rollback_record.previous_marker = previous_marker; + rollback_record.previous_num_pages = entry->num_pages; + rollback_record.previous_head_node = entry->head_node; + rollback_record.previous_tail_node = entry->tail_node; + rollback_record.was_new = is_new; + rollback_records.push_back(rollback_record); + + if (is_new) { + header->active_sequences.fetch_add(1, std::memory_order_relaxed); + } - if (header->free_stack_top.load(std::memory_order_relaxed) < new_pages) { - throw std::runtime_error( - "Insufficient free pages for prefix-aware allocation of sequence " + - std::to_string(sequence_ids[i])); - } - if (header->seq_page_node_free_top.load(std::memory_order_relaxed) < - required_pages) { - throw std::runtime_error( - "Insufficient sequence page nodes for prefix-aware allocation"); - } + std::vector pages; + pages.reserve(required_pages); - bool is_new = false; - SequenceEntry* entry = - FindOrInsertSequenceEntryLocked(sequence_ids[i], &is_new); - if (is_new) { - header->active_sequences.fetch_add(1, std::memory_order_relaxed); - } + for (std::size_t j = 0; j < reused_pages; ++j) { + const std::int32_t page_idx = reused_page_candidates[j]; + IncrementPageRefLocked(page_idx); + AppendPageToSequenceLocked(entry, page_idx); + pages.push_back(page_idx); + } - for (std::size_t j = 0; j < reused_pages; ++j) { - const std::int32_t page_idx = lookup.pages[j]; - IncrementPageRefLocked(page_idx); - AppendPageToSequenceLocked(entry, page_idx); - pages.push_back(page_idx); - } + for (std::size_t j = 0; j < new_pages; ++j) { + const std::int32_t page_idx = PopFreePageLocked(); + IncrementPageRefLocked(page_idx); + AppendPageToSequenceLocked(entry, page_idx); + pages.push_back(page_idx); + } - for (std::size_t j = 0; j < new_pages; ++j) { - const std::int32_t page_idx = PopFreePageLocked(); - IncrementPageRefLocked(page_idx); - AppendPageToSequenceLocked(entry, page_idx); - pages.push_back(page_idx); + result.allocated_pages[i] = std::move(pages); + result.reused_prefix_tokens[i] = reused_pages * config.page_size_tokens; } + } catch (...) { + for (auto it = rollback_records.rbegin(); it != rollback_records.rend(); + ++it) { + SequenceRollbackRecord& record = *it; + SequenceEntry* entry = record.entry; + if (entry == nullptr) { + continue; + } - result.allocated_pages[i] = std::move(pages); - result.reused_prefix_tokens[i] = reused_pages * config.page_size_tokens; + std::int32_t first_added_node = kInvalidIndex; + if (entry->num_pages > record.previous_num_pages) { + if (record.previous_num_pages == 0) { + first_added_node = entry->head_node; + } else { + first_added_node = + seq_page_nodes[record.previous_tail_node].next_node; + } + } + + std::int32_t node_idx = first_added_node; + while (node_idx != kInvalidIndex) { + const std::int32_t next_node = seq_page_nodes[node_idx].next_node; + const std::int32_t page_idx = seq_page_nodes[node_idx].page_idx; + DecrementPageRefLocked(page_idx); + FreeSeqPageNodeLocked(node_idx); + node_idx = next_node; + } + + if (record.previous_tail_node != kInvalidIndex) { + seq_page_nodes[record.previous_tail_node].next_node = + kInvalidIndex; + } + + entry->num_pages = record.previous_num_pages; + entry->head_node = record.previous_head_node; + entry->tail_node = record.previous_tail_node; + entry->sequence_id = record.previous_marker; + + if (record.was_new) { + header->active_sequences.fetch_sub(1, std::memory_order_relaxed); + } + } + throw; } return result; @@ -1101,7 +1183,7 @@ HostPagedKVStats HostPagedKVBackend::SharedState::CollectStats() const { header->prefix_miss_count.load(std::memory_order_relaxed); stats.num_prefix_evictions = header->prefix_evict_count.load(std::memory_order_relaxed); - stats.num_prefix_pinned_pages = + stats.num_cache_entry_pages = header->prefix_used_pages.load(std::memory_order_relaxed); std::size_t shared_pages = 0; diff --git a/core/KV_Storage/host_paged_kv_backend.h b/core/KV_Storage/host_paged_kv_backend.h index 58b779e51..28ad544c7 100644 --- a/core/KV_Storage/host_paged_kv_backend.h +++ b/core/KV_Storage/host_paged_kv_backend.h @@ -24,7 +24,7 @@ struct HostPagedKVStats { std::size_t num_prefix_hits = 0; std::size_t num_prefix_misses = 0; std::size_t num_prefix_evictions = 0; - std::size_t num_prefix_pinned_pages = 0; + std::size_t num_cache_entry_pages = 0; std::size_t num_shared_pages = 0; }; @@ -130,7 +130,8 @@ inline HostPagedKVConfig SanitizeConfig(HostPagedKVConfig config) { std::max(1024, config.num_pages / 64); } if (config.prefix_page_budget == 0) { - config.prefix_page_budget = std::max(64, config.num_pages / 4); + config.prefix_page_budget = + std::max(128, config.num_pages / 2); } if (config.prefix_page_ref_capacity == 0) { config.prefix_page_ref_capacity = std::max( @@ -197,7 +198,7 @@ inline std::string ToString(const HostPagedKVStats& stats) { << ", prefix_hits=" << stats.num_prefix_hits << ", prefix_misses=" << stats.num_prefix_misses << ", prefix_evictions=" << stats.num_prefix_evictions - << ", prefix_pinned_pages=" << stats.num_prefix_pinned_pages + << ", cache_entry_pages=" << stats.num_cache_entry_pages << ", shared_pages=" << stats.num_shared_pages << ")"; return oss.str(); } diff --git a/core/KV_Storage/host_paged_kv_prefix_cache.cpp b/core/KV_Storage/host_paged_kv_prefix_cache.cpp index 37fe16e6a..aace747f5 100644 --- a/core/KV_Storage/host_paged_kv_prefix_cache.cpp +++ b/core/KV_Storage/host_paged_kv_prefix_cache.cpp @@ -179,6 +179,43 @@ std::int32_t HostKVPrefixCache::FindEdgeByFirstTokenLocked( return kHostKVInvalidIndex; } +std::int32_t HostKVPrefixCache::FindExactPathNodeLocked( + const std::int32_t* tokens, std::size_t token_count) const { + if (tokens == nullptr) { + return kHostKVInvalidIndex; + } + + std::int32_t node_idx = 0; + std::size_t pos = 0; + while (pos < token_count) { + const std::int32_t edge_idx = + FindEdgeByFirstTokenLocked(node_idx, tokens[pos]); + if (edge_idx == kHostKVInvalidIndex) { + return kHostKVInvalidIndex; + } + + const HostKVRadixEdge& edge = radix_edges_[edge_idx]; + const std::size_t remaining = token_count - pos; + if (remaining < edge.label_len) { + return kHostKVInvalidIndex; + } + + std::size_t common = 0; + while (common < edge.label_len && + edge.label_tokens[common] == tokens[pos + common]) { + ++common; + } + if (common != edge.label_len) { + return kHostKVInvalidIndex; + } + + pos += common; + node_idx = edge.child_node; + } + + return node_idx; +} + std::int32_t HostKVPrefixCache::AppendTokenPathLocked(std::int32_t start_node, const std::int32_t* tokens, std::size_t token_count) { @@ -527,6 +564,17 @@ bool HostKVPrefixCache::CommitPrefixLocked(const std::int32_t* tokens, return false; } + const std::int32_t existing_terminal_node = + FindExactPathNodeLocked(tokens, token_count); + if (existing_terminal_node != kHostKVInvalidIndex) { + const std::int32_t existing_entry_idx = + radix_nodes_[existing_terminal_node].terminal_entry; + if (existing_entry_idx != kHostKVInvalidIndex) { + TouchPrefixEntryLocked(existing_entry_idx); + return true; + } + } + const std::size_t max_needed_nodes = (token_count + kHostKVRadixEdgeLabelChunk - 1) / kHostKVRadixEdgeLabelChunk + diff --git a/core/KV_Storage/host_paged_kv_prefix_cache.h b/core/KV_Storage/host_paged_kv_prefix_cache.h index 91ee08aaa..7306d8caa 100644 --- a/core/KV_Storage/host_paged_kv_prefix_cache.h +++ b/core/KV_Storage/host_paged_kv_prefix_cache.h @@ -127,6 +127,8 @@ class HostKVPrefixCache { std::int32_t FindEdgeByFirstTokenLocked(std::int32_t node_idx, std::int32_t first_token) const; + std::int32_t FindExactPathNodeLocked(const std::int32_t* tokens, + std::size_t token_count) const; std::int32_t AppendTokenPathLocked(std::int32_t start_node, const std::int32_t* tokens, std::size_t token_count); diff --git a/core/KV_Storage/host_paged_kv_worker_view.h b/core/KV_Storage/host_paged_kv_worker_view.h index 7f6fa0403..e4af93495 100644 --- a/core/KV_Storage/host_paged_kv_worker_view.h +++ b/core/KV_Storage/host_paged_kv_worker_view.h @@ -303,7 +303,6 @@ class HostPagedKVWorkerView { std::vector( flat_prompt_tokens.begin() + begin, flat_prompt_tokens.begin() + end), - end - begin, alloc_result.reused_prefix_tokens[i]); } @@ -1050,7 +1049,9 @@ class HostPagedKVWorkerView { KVAsyncTask AsyncAppendDecodeKVToHost( std::size_t layer_idx, std::vector sequence_ids, torch::Tensor k_tensor, std::optional v_tensor, - SequenceLengths sequence_lengths) { + SequenceLengths sequence_lengths, + std::optional> decode_token_ids = + std::nullopt) { geometry_.EnsureLayerBounds(layer_idx, "AsyncAppendDecodeKVToHost"); EnsureDeviceReady(); const std::size_t batch = sequence_ids.size(); @@ -1061,6 +1062,12 @@ class HostPagedKVWorkerView { ValidateKTensorShape(k_tensor, batch); ValidateSequenceLengthsInput(sequence_lengths, batch, "AsyncAppendDecodeKVToHost"); + if (decode_token_ids.has_value() && + decode_token_ids->size() != batch) { + throw std::invalid_argument( + "AsyncAppendDecodeKVToHost: decode_token_ids size must match " + "sequence_ids size"); + } torch::Tensor prepared_k = k_tensor; std::optional prepared_v; if (v_tensor.has_value()) { @@ -1082,7 +1089,9 @@ class HostPagedKVWorkerView { return LaunchAsyncTask([this, layer_idx, sequence_ids = std::move(sequence_ids), sequence_lengths = std::move(sequence_lengths), - prepared_k, prepared_v]() mutable { + prepared_k, prepared_v, + decode_token_ids = std::move(decode_token_ids)]() + mutable { c10::cuda::OptionalCUDAGuard device_guard(device_index_); const auto cuda_stream = CopyStream(CopyDirection::kDeviceToHost); @@ -1142,6 +1151,11 @@ class HostPagedKVWorkerView { } this->SynchronizeWithEvent(cuda_stream); + if (decode_token_ids.has_value() && + layer_idx + 1 == config_.num_layers) { + this->AppendDecodeTokens(sequence_ids, *decode_token_ids); + } + this->MaybeCommitPrefix(layer_idx, sequence_ids); }); } @@ -1297,25 +1311,42 @@ class HostPagedKVWorkerView { } struct PrefixSequenceState { - std::vector prompt_tokens; - std::size_t prompt_token_count = 0; + std::vector tokens; std::size_t reused_prefix_tokens = 0; - bool committed = false; + std::size_t committed_token_count = 0; }; void UpdatePrefixState(std::int64_t sequence_id, std::vector prompt_tokens, - std::size_t prompt_token_count, std::size_t reused_prefix_tokens) { if (!config_.enable_prefix_reuse) { return; } std::lock_guard guard(prefix_state_mutex_); auto& state = prefix_states_[sequence_id]; - state.prompt_tokens = std::move(prompt_tokens); - state.prompt_token_count = prompt_token_count; + state.tokens = std::move(prompt_tokens); state.reused_prefix_tokens = reused_prefix_tokens; - state.committed = false; + state.committed_token_count = 0; + } + + void AppendDecodeTokens(const std::vector& sequence_ids, + const std::vector& decode_token_ids) { + if (!config_.enable_prefix_reuse) { + return; + } + if (sequence_ids.size() != decode_token_ids.size()) { + throw std::invalid_argument( + "AppendDecodeTokens: decode_token_ids size must match " + "sequence_ids size"); + } + std::lock_guard guard(prefix_state_mutex_); + for (std::size_t i = 0; i < sequence_ids.size(); ++i) { + auto it = prefix_states_.find(sequence_ids[i]); + if (it == prefix_states_.end()) { + continue; + } + it->second.tokens.push_back(decode_token_ids[i]); + } } std::size_t PrefixTokensToSkip(std::int64_t sequence_id) const { @@ -1336,13 +1367,13 @@ class HostPagedKVWorkerView { return; } - struct PrefixCommitItem { + struct CacheCommitItem { std::int64_t sequence_id = 0; - std::vector prompt_tokens; - std::size_t prompt_token_count = 0; + std::vector tokens; + std::size_t token_count = 0; }; - std::vector commit_items; + std::vector commit_items; { std::lock_guard guard(prefix_state_mutex_); for (std::int64_t sequence_id : sequence_ids) { @@ -1350,19 +1381,26 @@ class HostPagedKVWorkerView { if (it == prefix_states_.end()) { continue; } - if (it->second.committed) { + const std::size_t token_count = it->second.tokens.size(); + const std::size_t full_prompt_pages = + token_count / config_.page_size_tokens; + if (full_prompt_pages < config_.prefix_min_store_pages) { + continue; + } + const std::size_t full_page_token_count = + full_prompt_pages * config_.page_size_tokens; + if (full_page_token_count <= it->second.committed_token_count) { continue; } commit_items.push_back( - {sequence_id, it->second.prompt_tokens, - it->second.prompt_token_count}); - it->second.committed = true; + {sequence_id, it->second.tokens, full_page_token_count}); + it->second.committed_token_count = full_page_token_count; } } for (const auto& item : commit_items) { - backend_.CommitSequencePrefix(item.sequence_id, item.prompt_tokens, - item.prompt_token_count); + backend_.CommitSequencePrefix(item.sequence_id, item.tokens, + item.token_count); } } diff --git a/core/batchgen_Binding.cpp b/core/batchgen_Binding.cpp index d42dec1fd..1f9ac6ec6 100644 --- a/core/batchgen_Binding.cpp +++ b/core/batchgen_Binding.cpp @@ -374,7 +374,8 @@ void BindHostPagedWorkerView(py::module& m, const char* name) { &WorkerView::AsyncAppendDecodeKVToHost, py::arg("layer_idx"), py::arg("sequence_ids"), py::arg("k_tensor"), py::arg("v_tensor") = py::none(), - py::arg("sequence_lengths")) + py::arg("sequence_lengths"), + py::arg("decode_token_ids") = py::none()) .def( "async_load_layer_kv_to_device", [](WorkerView& self, torch::Tensor sequence_ids, @@ -622,8 +623,8 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { &kv::HostPagedKVStats::num_prefix_misses) .def_readwrite("num_prefix_evictions", &kv::HostPagedKVStats::num_prefix_evictions) - .def_readwrite("num_prefix_pinned_pages", - &kv::HostPagedKVStats::num_prefix_pinned_pages) + .def_readwrite("num_cache_entry_pages", + &kv::HostPagedKVStats::num_cache_entry_pages) .def_readwrite("num_shared_pages", &kv::HostPagedKVStats::num_shared_pages) .def("__repr__", diff --git a/test/paged_kv/test_host_prefix_reuse.py b/test/paged_kv/test_host_prefix_reuse.py index 8afa691c2..2532f98d4 100644 --- a/test/paged_kv/test_host_prefix_reuse.py +++ b/test/paged_kv/test_host_prefix_reuse.py @@ -27,13 +27,13 @@ def _shm_unlink(name: str) -> None: def _make_mla_config( - shm_name: str, enable_prefix_reuse: bool + shm_name: str, enable_prefix_reuse: bool, page_size_tokens: int = 64 ) -> bg.HostPagedKVConfig: # type: ignore cfg = bg.HostPagedKVConfig() cfg.shm_name = shm_name cfg.num_layers = 1 cfg.num_pages = 512 - cfg.page_size_tokens = 64 + cfg.page_size_tokens = page_size_tokens cfg.num_k_heads = 1 cfg.k_head_dim = 576 cfg.num_v_heads = 0 @@ -142,3 +142,112 @@ def test_prefix_reuse_hits_after_commit() -> None: worker.shutdown() del worker _shm_unlink(shm_name) + + +def test_prefix_reuse_includes_decode_tokens() -> None: + if not torch.cuda.is_available(): + raise SkipTest("CUDA is not available") + + shm_name = _random_shm_name() + cfg = _make_mla_config( + shm_name, enable_prefix_reuse=True, page_size_tokens=8 + ) + worker = bg.MLAHostPagedKVWorkerView(cfg) + + try: + worker.initialize(device_index=0, create_region=True) + + prompt_tokens = [100 + i for i in range(cfg.page_size_tokens)] + decode_tokens = [200 + i for i in range(cfg.page_size_tokens)] + full_tokens = prompt_tokens + decode_tokens + + worker.register_sequences([301]) + pages_1, reused_1 = worker.allocate_pages_for_sequences_with_prefix( + [(301, len(full_tokens))], prompt_tokens, [0, len(prompt_tokens)] + ) + assert reused_1 == [0] + assert len(pages_1[0]) == 2 + + prefill_k = torch.full( + (1, len(prompt_tokens), cfg.num_k_heads, cfg.k_head_dim), + fill_value=1.0, + dtype=torch.bfloat16, + device="cuda:0", + ) + worker.async_offload_layer_kv_to_host( + layer_idx=0, + sequence_ids=[301], + k_tensor=prefill_k, + v_tensor=None, + sequence_lengths=[len(prompt_tokens)], + ).wait() + + for step, token_id in enumerate(decode_tokens): + decode_k = torch.full( + (1, 1, cfg.num_k_heads, cfg.k_head_dim), + fill_value=2.0 + step, + dtype=torch.bfloat16, + device="cuda:0", + ) + worker.async_append_decode_kv_to_host( + layer_idx=0, + sequence_ids=[301], + k_tensor=decode_k, + v_tensor=None, + sequence_lengths=[len(prompt_tokens) + step], + decode_token_ids=[token_id], + ).wait() + + worker.release_sequence_pages([301]) + + worker.register_sequences([302]) + pages_2, reused_2 = worker.allocate_pages_for_sequences_with_prefix( + [(302, len(full_tokens))], full_tokens, [0, len(full_tokens)] + ) + assert len(pages_2[0]) == 2 + assert reused_2 == [len(full_tokens)] + + stats = worker.get_stats() + assert stats.num_prefix_hits >= 1 + assert stats.num_shared_pages >= 1 + + worker.release_sequence_pages([302]) + finally: + worker.shutdown() + del worker + _shm_unlink(shm_name) + + +def test_prefix_batch_allocation_failure_rolls_back() -> None: + if not torch.cuda.is_available(): + raise SkipTest("CUDA is not available") + + shm_name = _random_shm_name() + cfg = _make_mla_config(shm_name, enable_prefix_reuse=True) + cfg.num_pages = 2 + cfg.sequence_page_node_capacity = 16 + worker = bg.MLAHostPagedKVWorkerView(cfg) + + try: + worker.initialize(device_index=0, create_region=True) + worker.register_sequences([401, 402]) + + flat_prompt_tokens = list(range(128)) + raised = False + try: + worker.allocate_pages_for_sequences_with_prefix( + [(401, 64), (402, 128)], + flat_prompt_tokens, + [0, 64, 128], + ) + except RuntimeError: + raised = True + assert raised + + stats = worker.get_stats() + assert stats.num_used_pages == 0 + assert stats.num_active_sequences == 0 + finally: + worker.shutdown() + del worker + _shm_unlink(shm_name) diff --git a/test/paged_kv/test_prefix_cache_binding.py b/test/paged_kv/test_prefix_cache_binding.py index 25fa3a599..3733d6c79 100644 --- a/test/paged_kv/test_prefix_cache_binding.py +++ b/test/paged_kv/test_prefix_cache_binding.py @@ -98,3 +98,40 @@ def test_prefix_cache_harness_lru_evict_releases_page_refs() -> None: pages, reused = cache.lookup(tokens_b, 2) assert pages == [2, 3] assert reused == 2 + + +def test_prefix_cache_harness_supports_decode_extension_commit() -> None: + cache = _make_prefix_cache(prefix_min_store_pages=1, prefix_page_budget=16) + prompt_tokens = [3000 + i for i in range(64)] + decode_tokens = [4000 + i for i in range(64)] + + assert cache.commit(prompt_tokens, [9]) is True + assert cache.commit(prompt_tokens + decode_tokens, [9, 10]) is True + + pages, reused = cache.lookup(prompt_tokens, 1) + assert pages == [9] + assert reused == 1 + + pages, reused = cache.lookup(prompt_tokens + decode_tokens, 2) + assert pages == [9, 10] + assert reused == 2 + + +def test_prefix_cache_harness_duplicate_commit_avoids_spurious_evict() -> None: + cache = _make_prefix_cache(num_pages=8, prefix_page_budget=2) + tokens = [5000 + i for i in range(128)] + + assert cache.commit(tokens, [0, 1]) is True + stats_before = cache.get_stats() + assert stats_before.prefix_entry_count == 1 + assert stats_before.prefix_evict_count == 0 + assert cache.page_refcount(0) == 1 + assert cache.page_refcount(1) == 1 + + # Re-committing the exact same prefix should only touch recency metadata. + assert cache.commit(tokens, [0, 1]) is True + stats_after = cache.get_stats() + assert stats_after.prefix_entry_count == 1 + assert stats_after.prefix_evict_count == 0 + assert cache.page_refcount(0) == 1 + assert cache.page_refcount(1) == 1 From 23debbcf4b176f5ffacc0f9865ff6e12742aa0ab Mon Sep 17 00:00:00 2001 From: luzhan <513964121@qq.com> Date: Tue, 3 Mar 2026 14:40:04 +0000 Subject: [PATCH 03/17] feat: enhance prefix cache logging for reused and skipped tokens --- batchgen/batchgen_worker.py | 26 ++++++++++++- core/KV_Storage/host_paged_kv_worker_view.h | 42 +++++++++++++++++++++ 2 files changed, 67 insertions(+), 1 deletion(-) diff --git a/batchgen/batchgen_worker.py b/batchgen/batchgen_worker.py index fce3a9294..d16f4afc5 100644 --- a/batchgen/batchgen_worker.py +++ b/batchgen/batchgen_worker.py @@ -1006,6 +1006,11 @@ def _append_decode_kv_to_host_async( 'decoded_len': seq.decoded_length, 'write_pos': write_pos, }) + if not has_decode_tokens: + logging.debug( + f"Rank {self.rank}: PrefixCache decode token capture unavailable " + f"(layer={layer_idx}, batch_size={len(batch)}), skip decode token cache update" + ) # Log append positions for resumed sequences (layer 0 only to reduce spam) if layer_idx == 0 and append_diag and BATCHGEN_CB_DEBUG: @@ -4445,11 +4450,25 @@ def _config_prefill_for_batch(self, prefill_uuids: List[str]) -> None: and getattr(host_cfg, "enable_prefix_reuse", False) ) if enable_prefix_reuse: - self.core_engine.host_paged_kv_worker_view.allocate_pages_for_sequences_with_prefix( + _, reused_prefix_tokens = self.core_engine.host_paged_kv_worker_view.allocate_pages_for_sequences_with_prefix( list(zip(global_sequence_ids, sequence_tokens)), flat_prompt_tokens, prompt_offsets, ) + reused_sequences = sum(1 for tokens in reused_prefix_tokens if tokens > 0) + reused_tokens_total = sum(int(tokens) for tokens in reused_prefix_tokens) + page_size = max(1, int(getattr(host_cfg, "page_size_tokens", 1))) + if reused_sequences > 0: + logging.info( + f"Rank {self.rank}: PrefixCache allocation summary " + f"(batch={len(global_sequence_ids)}, reused_sequences={reused_sequences}, " + f"reused_tokens={reused_tokens_total}, reused_pages={reused_tokens_total // page_size})" + ) + else: + logging.debug( + f"Rank {self.rank}: PrefixCache allocation summary " + f"(batch={len(global_sequence_ids)}, reused_sequences=0, reused_tokens=0, reused_pages=0)" + ) else: self.core_engine.host_paged_kv_worker_view.allocate_pages_for_sequences( list(zip(global_sequence_ids, sequence_tokens)) @@ -6790,6 +6809,11 @@ def _append_decode_kv_to_host_async( has_decode_tokens = False else: decode_token_ids.append(int(seq.input_ids[0, write_pos].item())) + if not has_decode_tokens: + logging.debug( + f"Rank {self.rank}: PrefixCache decode token capture unavailable " + f"(layer={layer_idx}, batch_size={len(batch)}), skip decode token cache update" + ) if k_tensor.dim() == 3: k_tensor = k_tensor.unsqueeze(2) diff --git a/core/KV_Storage/host_paged_kv_worker_view.h b/core/KV_Storage/host_paged_kv_worker_view.h index e4af93495..c8cf1dd6e 100644 --- a/core/KV_Storage/host_paged_kv_worker_view.h +++ b/core/KV_Storage/host_paged_kv_worker_view.h @@ -295,6 +295,8 @@ class HostPagedKVWorkerView { auto alloc_result = backend_.AcquirePagesForSequencesWithPrefix( sequence_ids, num_tokens, flat_prompt_tokens, prompt_offsets); + std::size_t reused_sequence_count = 0; + std::size_t reused_token_total = 0; for (std::size_t i = 0; i < sequence_ids.size(); ++i) { AppendAllocatedPages(sequence_ids[i], alloc_result.allocated_pages[i]); const std::size_t begin = prompt_offsets[i]; @@ -304,6 +306,23 @@ class HostPagedKVWorkerView { flat_prompt_tokens.begin() + begin, flat_prompt_tokens.begin() + end), alloc_result.reused_prefix_tokens[i]); + reused_token_total += alloc_result.reused_prefix_tokens[i]; + if (alloc_result.reused_prefix_tokens[i] > 0) { + ++reused_sequence_count; + } + } + if (config_.enable_prefix_reuse) { + if (reused_sequence_count > 0) { + logger_->info( + "PrefixCache allocation completed (sequences={}, reused_sequences={}, reused_tokens={}, reused_pages={})", + sequence_ids.size(), reused_sequence_count, + reused_token_total, + reused_token_total / config_.page_size_tokens); + } else { + logger_->debug( + "PrefixCache allocation completed (sequences={}, reused_sequences=0, reused_tokens=0, reused_pages=0)", + sequence_ids.size()); + } } return {std::move(alloc_result.allocated_pages), @@ -944,8 +963,20 @@ class HostPagedKVWorkerView { } std::vector prefix_skip_tokens(sequence_ids.size(), 0); + std::size_t skipped_sequence_count = 0; + std::size_t skipped_token_total = 0; for (std::size_t i = 0; i < sequence_ids.size(); ++i) { prefix_skip_tokens[i] = PrefixTokensToSkip(sequence_ids[i]); + skipped_token_total += prefix_skip_tokens[i]; + if (prefix_skip_tokens[i] > 0) { + ++skipped_sequence_count; + } + } + if (config_.enable_prefix_reuse && skipped_sequence_count > 0) { + logger_->debug( + "PrefixCache prefill skip planned (layer={}, sequences={}, skipped_sequences={}, skipped_tokens={})", + layer_idx, sequence_ids.size(), skipped_sequence_count, + skipped_token_total); } return LaunchAsyncTask([this, layer_idx, @@ -1327,6 +1358,9 @@ class HostPagedKVWorkerView { state.tokens = std::move(prompt_tokens); state.reused_prefix_tokens = reused_prefix_tokens; state.committed_token_count = 0; + logger_->debug( + "PrefixCache state initialized (sequence_id={}, token_count={}, reused_tokens={})", + sequence_id, state.tokens.size(), reused_prefix_tokens); } void AppendDecodeTokens(const std::vector& sequence_ids, @@ -1374,6 +1408,7 @@ class HostPagedKVWorkerView { }; std::vector commit_items; + std::size_t committed_token_total = 0; { std::lock_guard guard(prefix_state_mutex_); for (std::int64_t sequence_id : sequence_ids) { @@ -1395,8 +1430,15 @@ class HostPagedKVWorkerView { commit_items.push_back( {sequence_id, it->second.tokens, full_page_token_count}); it->second.committed_token_count = full_page_token_count; + committed_token_total += full_page_token_count; } } + if (!commit_items.empty()) { + logger_->info( + "PrefixCache commit scheduled (layer={}, commit_count={}, committed_tokens={}, committed_pages={})", + layer_idx, commit_items.size(), committed_token_total, + committed_token_total / config_.page_size_tokens); + } for (const auto& item : commit_items) { backend_.CommitSequencePrefix(item.sequence_id, item.tokens, From 9490730e3731e5e617f4f8afc8b808ef5d40dd57 Mon Sep 17 00:00:00 2001 From: luzhan <513964121@qq.com> Date: Tue, 3 Mar 2026 17:21:36 +0000 Subject: [PATCH 04/17] feat: add decode token handling in BatchGenWorker for improved sequence processing --- batchgen/batchgen_worker.py | 38 +++++++++++++++++++++++++++++++++++-- 1 file changed, 36 insertions(+), 2 deletions(-) diff --git a/batchgen/batchgen_worker.py b/batchgen/batchgen_worker.py index d16f4afc5..5c786f070 100644 --- a/batchgen/batchgen_worker.py +++ b/batchgen/batchgen_worker.py @@ -989,7 +989,24 @@ def _append_decode_kv_to_host_async( # Write position is current position (0-indexed) write_pos = seq.current_context_length - 1 sequence_lengths.append(write_pos) - if ( + decode_pos = write_pos - seq.prompt_length + if decode_pos >= 0: + query_entry = self.query_book.get(local_idx) + decoded_tokens = ( + None if query_entry is None else query_entry.decoded_tokens + ) + if ( + decode_pos >= seq.decoded_length + or decoded_tokens is None + or decoded_tokens.dim() < 2 + or decode_pos >= decoded_tokens.shape[1] + ): + has_decode_tokens = False + else: + decode_token_ids.append( + int(decoded_tokens[0, decode_pos].item()) + ) + elif ( seq.input_ids is None or write_pos < 0 or write_pos >= seq.input_ids.shape[1] @@ -6801,7 +6818,24 @@ def _append_decode_kv_to_host_async( sequence_ids.append(seq.global_idx) write_pos = seq.current_context_length - 1 sequence_lengths.append(write_pos) - if ( + decode_pos = write_pos - seq.prompt_length + if decode_pos >= 0: + query_entry = self.query_book.get(local_idx) + decoded_tokens = ( + None if query_entry is None else query_entry.decoded_tokens + ) + if ( + decode_pos >= seq.decoded_length + or decoded_tokens is None + or decoded_tokens.dim() < 2 + or decode_pos >= decoded_tokens.shape[1] + ): + has_decode_tokens = False + else: + decode_token_ids.append( + int(decoded_tokens[0, decode_pos].item()) + ) + elif ( seq.input_ids is None or write_pos < 0 or write_pos >= seq.input_ids.shape[1] From 32f8dcb6d47d6d78b45d341262bbbde810eb125f Mon Sep 17 00:00:00 2001 From: luzhan <513964121@qq.com> Date: Thu, 5 Mar 2026 19:41:29 +0000 Subject: [PATCH 05/17] feat: update sequence handling in HostPagedKVWorkerView for improved performance --- .gitignore | 1 + core/KV_Storage/host_paged_kv_worker_view.h | 8 ++++---- 2 files changed, 5 insertions(+), 4 deletions(-) diff --git a/.gitignore b/.gitignore index 4ccc9be49..c7f0d0c67 100644 --- a/.gitignore +++ b/.gitignore @@ -4,6 +4,7 @@ log_*.log log *.log +tmp/ *.pt runner.sh diff --git a/core/KV_Storage/host_paged_kv_worker_view.h b/core/KV_Storage/host_paged_kv_worker_view.h index c8cf1dd6e..564e59401 100644 --- a/core/KV_Storage/host_paged_kv_worker_view.h +++ b/core/KV_Storage/host_paged_kv_worker_view.h @@ -910,10 +910,10 @@ class HostPagedKVWorkerView { if (sequence_ids.empty()) { return; } - std::for_each(sequence_ids.begin(), sequence_ids.end(), - [this](std::int64_t sequence_id) { - UnregisterSequence(sequence_id); - }); + for (std::int64_t sequence_id : sequence_ids) { + page_table_.Remove(sequence_id); + } + RemovePrefixStates(sequence_ids); } void ReleaseSequencePages(const std::vector& sequence_ids) { From 9955bd85dab9c95cce52c9f810bed46a712a6f79 Mon Sep 17 00:00:00 2001 From: luzhan <513964121@qq.com> Date: Fri, 6 Mar 2026 11:39:39 +0000 Subject: [PATCH 06/17] fix: restore namespace scope in host paged kv backend --- core/KV_Storage/host_paged_kv_backend.cpp | 2 -- 1 file changed, 2 deletions(-) diff --git a/core/KV_Storage/host_paged_kv_backend.cpp b/core/KV_Storage/host_paged_kv_backend.cpp index aa9a7fcae..87f194c29 100644 --- a/core/KV_Storage/host_paged_kv_backend.cpp +++ b/core/KV_Storage/host_paged_kv_backend.cpp @@ -213,8 +213,6 @@ void TouchPagesMultiThreaded(void* ptr, std::size_t size, std::size_t stride) { } } -} // namespace - struct HostPagedKVBackend::SharedState { explicit SharedState(const HostPagedKVConfig& cfg, std::size_t data_bytes, std::uint64_t fingerprint, bool has_v) From 8311d06fa1f0ee045f96606c26aa77383cbdfe0b Mon Sep 17 00:00:00 2001 From: luzhan <513964121@qq.com> Date: Fri, 6 Mar 2026 12:31:51 +0000 Subject: [PATCH 07/17] feat: allow prefix reuse via env --- batchgen/kv_cache/host_kv_mananger_config.py | 13 +++++++++++-- 1 file changed, 11 insertions(+), 2 deletions(-) diff --git a/batchgen/kv_cache/host_kv_mananger_config.py b/batchgen/kv_cache/host_kv_mananger_config.py index fe0f23fb5..b2d0d92cf 100644 --- a/batchgen/kv_cache/host_kv_mananger_config.py +++ b/batchgen/kv_cache/host_kv_mananger_config.py @@ -1,6 +1,7 @@ from __future__ import annotations from dataclasses import dataclass +import os from typing import Any, Dict, Sequence import torch @@ -203,6 +204,14 @@ def _apply_capacity_defaults(config: Any) -> None: ) +def _env_flag(name: str, default: bool = False) -> bool: + """Parse a boolean feature flag from the environment.""" + value = os.environ.get(name) + if value is None: + return default + return value.strip().lower() not in {"", "0", "false", "no", "off"} + + def build_host_kv_config(model_name: str, host_kv_cache_size: int) -> Any: """Builds a core HostPagedKVConfig for the given model and host budget.""" @@ -247,7 +256,7 @@ def build_host_kv_config(model_name: str, host_kv_cache_size: int) -> Any: profile.sequence_table_capacity or config.num_pages ) config.alignment_bytes = profile.alignment_bytes - config.enable_prefix_reuse = False + config.enable_prefix_reuse = _env_flag("BATCHGEN_ENABLE_PREFIX_REUSE", False) config.prefix_min_reuse_pages = 1 config.prefix_min_store_pages = 2 _apply_capacity_defaults(config) @@ -366,7 +375,7 @@ def build_host_kv_config_aux(model_name: str, host_kv_cache_size: int) -> Any | profile.sequence_table_capacity or config.num_pages ) config.alignment_bytes = profile.alignment_bytes - config.enable_prefix_reuse = False + config.enable_prefix_reuse = _env_flag("BATCHGEN_ENABLE_PREFIX_REUSE", False) config.prefix_min_reuse_pages = 1 config.prefix_min_store_pages = 2 _apply_capacity_defaults(config) From db58e5095282a564fd371e92354e774c825c378d Mon Sep 17 00:00:00 2001 From: luzhan <513964121@qq.com> Date: Fri, 6 Mar 2026 12:36:50 +0000 Subject: [PATCH 08/17] feat: add prefix cache server flag --- batchgen/batchgen_worker.py | 3 +++ batchgen/kv_cache/host_kv_mananger_config.py | 25 ++++++++++---------- batchgen/server/server_args.py | 15 ++++++++++++ batchgen/server/worker_manager.py | 4 ++++ 4 files changed, 34 insertions(+), 13 deletions(-) diff --git a/batchgen/batchgen_worker.py b/batchgen/batchgen_worker.py index 48a572f7a..dc0968298 100644 --- a/batchgen/batchgen_worker.py +++ b/batchgen/batchgen_worker.py @@ -252,6 +252,7 @@ class BatchGenWorkerArgs: disable_cuda_graphs: bool = False # Disable CUDA graph capture for decode attention cuda_graph_max_bucket_size: int = 128 # Max batch size per rank for CUDA graph capture cuda_graph_num_buckets: int = 16 # Number of CUDA graph bucket sizes + enable_prefix_cache: bool = True # Enable host KV prefix cache reuse # Dynamic host KV reservation host_kv_chunk_size: int = 8192 # Initial host KV chunk size in tokens host_kv_eviction_watermark: int = 10 # Trigger eviction when free < this % @@ -308,6 +309,7 @@ def __init__(self, args: BatchGenWorkerArgs): if args.global_rank == 0: logging.info( f"Dynamic Host KV Config: chunk_size={args.host_kv_chunk_size}, " + f"prefix_cache={args.enable_prefix_cache}, " f"eviction_watermark={args.host_kv_eviction_watermark}%, " f"eviction_enabled={args.enable_host_kv_eviction}, " f"adaptive_chunk={args.adaptive_chunk}" @@ -409,6 +411,7 @@ def __init__(self, args: BatchGenWorkerArgs): worker_kv_config = build_host_kv_config( model_name=args.model_name, host_kv_cache_size=args.global_host_kv_cache_size_gb * (1024**3), + enable_prefix_reuse=args.enable_prefix_cache, ) if args.fast_init: worker_kv_config.enable_memfd = True diff --git a/batchgen/kv_cache/host_kv_mananger_config.py b/batchgen/kv_cache/host_kv_mananger_config.py index b2d0d92cf..acb3b2de0 100644 --- a/batchgen/kv_cache/host_kv_mananger_config.py +++ b/batchgen/kv_cache/host_kv_mananger_config.py @@ -1,7 +1,6 @@ from __future__ import annotations from dataclasses import dataclass -import os from typing import Any, Dict, Sequence import torch @@ -204,15 +203,11 @@ def _apply_capacity_defaults(config: Any) -> None: ) -def _env_flag(name: str, default: bool = False) -> bool: - """Parse a boolean feature flag from the environment.""" - value = os.environ.get(name) - if value is None: - return default - return value.strip().lower() not in {"", "0", "false", "no", "off"} - - -def build_host_kv_config(model_name: str, host_kv_cache_size: int) -> Any: +def build_host_kv_config( + model_name: str, + host_kv_cache_size: int, + enable_prefix_reuse: bool = True, +) -> Any: """Builds a core HostPagedKVConfig for the given model and host budget.""" if host_kv_cache_size is None: @@ -256,7 +251,7 @@ def build_host_kv_config(model_name: str, host_kv_cache_size: int) -> Any: profile.sequence_table_capacity or config.num_pages ) config.alignment_bytes = profile.alignment_bytes - config.enable_prefix_reuse = _env_flag("BATCHGEN_ENABLE_PREFIX_REUSE", False) + config.enable_prefix_reuse = bool(enable_prefix_reuse) config.prefix_min_reuse_pages = 1 config.prefix_min_store_pages = 2 _apply_capacity_defaults(config) @@ -345,7 +340,11 @@ def build_gpu_kv_config_aux( ) -def build_host_kv_config_aux(model_name: str, host_kv_cache_size: int) -> Any | None: +def build_host_kv_config_aux( + model_name: str, + host_kv_cache_size: int, + enable_prefix_reuse: bool = True, +) -> Any | None: """Builds a HostPagedKVConfig for the DSA indexer host cache, or None.""" profile = _resolve_indexer_profile(model_name) @@ -375,7 +374,7 @@ def build_host_kv_config_aux(model_name: str, host_kv_cache_size: int) -> Any | profile.sequence_table_capacity or config.num_pages ) config.alignment_bytes = profile.alignment_bytes - config.enable_prefix_reuse = _env_flag("BATCHGEN_ENABLE_PREFIX_REUSE", False) + config.enable_prefix_reuse = bool(enable_prefix_reuse) config.prefix_min_reuse_pages = 1 config.prefix_min_store_pages = 2 _apply_capacity_defaults(config) diff --git a/batchgen/server/server_args.py b/batchgen/server/server_args.py index 38f02cfa0..b3734894d 100644 --- a/batchgen/server/server_args.py +++ b/batchgen/server/server_args.py @@ -92,6 +92,7 @@ class ServerArgs: disable_cuda_graphs: bool = False # Disable CUDA graph capture for decode attention cuda_graph_max_bucket_size: int = 128 # Max batch size per rank for CUDA graph capture cuda_graph_num_buckets: int = 16 # Number of CUDA graph bucket sizes + enable_prefix_cache: bool = True # Enable host KV prefix cache reuse # Dynamic host KV reservation settings host_kv_chunk_size: int = 8192 # Initial host KV chunk size in tokens (default: 8K) host_kv_eviction_watermark: int = 10 # Trigger host KV eviction when free pages < this % @@ -328,6 +329,19 @@ def _build_parser() -> argparse.ArgumentParser: default=16, help="Maximum number of CUDA graph bucket sizes (default: 16). More buckets = longer capture time but less padding waste.", ) + parser.set_defaults(enable_prefix_cache=True) + parser.add_argument( + "--enable-prefix-cache", + dest="enable_prefix_cache", + action="store_true", + help="Enable host KV prefix cache reuse (default: enabled)", + ) + parser.add_argument( + "--disable-prefix-cache", + dest="enable_prefix_cache", + action="store_false", + help="Disable host KV prefix cache reuse", + ) # Dynamic host KV reservation parser.add_argument( "--host-kv-chunk-size", @@ -518,6 +532,7 @@ def prepare_server_args(argv: Optional[list[str]] = None) -> ServerArgs: disable_cuda_graphs=parsed.disable_cuda_graphs, cuda_graph_max_bucket_size=parsed.cuda_graph_max_bucket_size, cuda_graph_num_buckets=parsed.cuda_graph_num_buckets, + enable_prefix_cache=parsed.enable_prefix_cache, host_kv_chunk_size=parsed.host_kv_chunk_size, host_kv_eviction_watermark=parsed.host_kv_eviction_watermark, enable_host_kv_eviction=parsed.enable_host_kv_eviction, diff --git a/batchgen/server/worker_manager.py b/batchgen/server/worker_manager.py index 08d5f5f1b..4cbd62919 100644 --- a/batchgen/server/worker_manager.py +++ b/batchgen/server/worker_manager.py @@ -196,6 +196,7 @@ def start(self) -> None: try: self.host_kv_manager = self.allocate_host_kv_cache( self.args.host_kv_cache_size, self.args.model, + enable_prefix_cache=self.args.enable_prefix_cache, enable_memfd=self.args.fast_init, ) except Exception as exc: @@ -499,6 +500,7 @@ def _spawn_workers(self) -> None: disable_cuda_graphs=self.args.disable_cuda_graphs, cuda_graph_max_bucket_size=self.args.cuda_graph_max_bucket_size, cuda_graph_num_buckets=self.args.cuda_graph_num_buckets, + enable_prefix_cache=self.args.enable_prefix_cache, host_kv_chunk_size=self.args.host_kv_chunk_size, enable_host_kv_eviction=self.args.enable_host_kv_eviction, host_kv_eviction_watermark=self.args.host_kv_eviction_watermark, @@ -841,11 +843,13 @@ def _configure_host_kv_cache_budget(self) -> None: @staticmethod def allocate_host_kv_cache( host_kv_cache_size_gb: int, model_name: str, + enable_prefix_cache: bool = True, enable_memfd: bool = False, ) -> Any: config = build_host_kv_config( host_kv_cache_size=host_kv_cache_size_gb * (1024**3), model_name=model_name, + enable_prefix_reuse=enable_prefix_cache, ) if enable_memfd: config.enable_memfd = True From 680762a82911998e4562018bc7b1ab6581fcaa4b Mon Sep 17 00:00:00 2001 From: luzhan <513964121@qq.com> Date: Fri, 6 Mar 2026 13:02:52 +0000 Subject: [PATCH 09/17] fix: use venv pip in install script --- scripts/install_deps.sh | 30 +++++++++++++++++++----------- 1 file changed, 19 insertions(+), 11 deletions(-) diff --git a/scripts/install_deps.sh b/scripts/install_deps.sh index 20076bac9..e730bcc24 100755 --- a/scripts/install_deps.sh +++ b/scripts/install_deps.sh @@ -42,6 +42,10 @@ print_error() { echo -e "${RED}[ERROR]${NC} $1" } +run_pip() { + python -m pip "$@" +} + check_prerequisites() { print_step "Checking prerequisites..." @@ -52,7 +56,11 @@ check_prerequisites() { fi PYTHON_VERSION=$(python -c 'import sys; print(f"{sys.version_info.major}.{sys.version_info.minor}")') - if [[ $(echo "$PYTHON_VERSION < 3.11" | bc -l) -eq 1 ]]; then + if ! python - <<'PY' +import sys +sys.exit(0 if sys.version_info >= (3, 11) else 1) +PY + then print_error "Python 3.11+ required. Found: $PYTHON_VERSION" exit 1 fi @@ -77,7 +85,7 @@ check_prerequisites() { # Check ninja (for fast builds) if ! command -v ninja &> /dev/null; then print_step "Installing ninja for faster builds..." - pip install ninja + run_pip install ninja fi print_success "ninja found" } @@ -128,7 +136,7 @@ install_torch() { fi else print_step "Installing PyTorch with CUDA 12.8 support..." - pip install torch==2.9.0+cu128 --index-url https://download.pytorch.org/whl/cu128 + run_pip install torch==2.9.0+cu128 --index-url https://download.pytorch.org/whl/cu128 print_success "PyTorch installed" fi } @@ -159,7 +167,7 @@ install_flash_attention() { print_step "Building flash-attention 3 (this may take 10-20 minutes)..." cd hopper - FLASH_ATTENTION_FORCE_BUILD=TRUE pip install . --no-build-isolation + FLASH_ATTENTION_FORCE_BUILD=TRUE run_pip install . --no-build-isolation print_success "flash-attention 3 installed" } @@ -174,7 +182,7 @@ install_flashmla() { fi print_step "Installing FlashMLA from git (this may take 5-10 minutes)..." - FLASH_MLA_DISABLE_SM100=1 pip install "git+https://github.com/deepseek-ai/FlashMLA.git@${FLASHMLA_COMMIT}" --no-build-isolation + FLASH_MLA_DISABLE_SM100=1 run_pip install "git+https://github.com/deepseek-ai/FlashMLA.git@${FLASHMLA_COMMIT}" --no-build-isolation print_success "FlashMLA installed" } @@ -206,7 +214,7 @@ install_deepgemm() { fi print_step "Building DeepGEMM (this may take 5-10 minutes)..." - pip install . --no-build-isolation + run_pip install . --no-build-isolation print_success "DeepGEMM installed" } @@ -219,7 +227,7 @@ install_batchgen_kernels() { if [[ -f "$BATCHGEN_DIR/batchgen_kernels/setup.py" ]]; then cd "$BATCHGEN_DIR/batchgen_kernels" - pip install . --no-build-isolation + run_pip install . --no-build-isolation print_success "batchgen_kernels installed" else print_warning "batchgen_kernels/setup.py not found, skipping kernel compilation" @@ -235,7 +243,7 @@ install_batchgen() { if [[ -f "$BATCHGEN_DIR/setup.py" ]]; then cd "$BATCHGEN_DIR" - pip install . + run_pip install . print_success "BatchGen installed" else print_error "Could not find BatchGen setup.py at $BATCHGEN_DIR" @@ -368,9 +376,9 @@ main() { if [[ $IS_HOPPER -eq 1 ]]; then if [[ -n "$WHEEL_DIR" && -d "$WHEEL_DIR" ]]; then print_step "Installing Hopper dependencies from pre-built wheels: $WHEEL_DIR" - pip install --find-links "$WHEEL_DIR" --no-index \ + run_pip install --find-links "$WHEEL_DIR" --no-index \ flash-attn-hopper flash-mla deep-gemm 2>/dev/null || \ - pip install "$WHEEL_DIR"/*.whl + run_pip install "$WHEEL_DIR"/*.whl print_success "Hopper dependencies installed from wheels" else install_flash_attention @@ -378,7 +386,7 @@ main() { install_deepgemm # Reinstall PyTorch — building deps from source may downgrade torch or triton print_step "Reinstalling PyTorch to ensure correct version after dependency builds..." - pip install torch==2.9.0+cu128 --index-url https://download.pytorch.org/whl/cu128 + run_pip install torch==2.9.0+cu128 --index-url https://download.pytorch.org/whl/cu128 print_success "PyTorch reinstalled" fi else From a50ac5cf061a5babec6448fb8bc4fa885eebdd4b Mon Sep 17 00:00:00 2001 From: luzhan <513964121@qq.com> Date: Fri, 6 Mar 2026 13:06:06 +0000 Subject: [PATCH 10/17] revert: restore install script --- scripts/install_deps.sh | 30 +++++++++++------------------- 1 file changed, 11 insertions(+), 19 deletions(-) diff --git a/scripts/install_deps.sh b/scripts/install_deps.sh index e730bcc24..20076bac9 100755 --- a/scripts/install_deps.sh +++ b/scripts/install_deps.sh @@ -42,10 +42,6 @@ print_error() { echo -e "${RED}[ERROR]${NC} $1" } -run_pip() { - python -m pip "$@" -} - check_prerequisites() { print_step "Checking prerequisites..." @@ -56,11 +52,7 @@ check_prerequisites() { fi PYTHON_VERSION=$(python -c 'import sys; print(f"{sys.version_info.major}.{sys.version_info.minor}")') - if ! python - <<'PY' -import sys -sys.exit(0 if sys.version_info >= (3, 11) else 1) -PY - then + if [[ $(echo "$PYTHON_VERSION < 3.11" | bc -l) -eq 1 ]]; then print_error "Python 3.11+ required. Found: $PYTHON_VERSION" exit 1 fi @@ -85,7 +77,7 @@ PY # Check ninja (for fast builds) if ! command -v ninja &> /dev/null; then print_step "Installing ninja for faster builds..." - run_pip install ninja + pip install ninja fi print_success "ninja found" } @@ -136,7 +128,7 @@ install_torch() { fi else print_step "Installing PyTorch with CUDA 12.8 support..." - run_pip install torch==2.9.0+cu128 --index-url https://download.pytorch.org/whl/cu128 + pip install torch==2.9.0+cu128 --index-url https://download.pytorch.org/whl/cu128 print_success "PyTorch installed" fi } @@ -167,7 +159,7 @@ install_flash_attention() { print_step "Building flash-attention 3 (this may take 10-20 minutes)..." cd hopper - FLASH_ATTENTION_FORCE_BUILD=TRUE run_pip install . --no-build-isolation + FLASH_ATTENTION_FORCE_BUILD=TRUE pip install . --no-build-isolation print_success "flash-attention 3 installed" } @@ -182,7 +174,7 @@ install_flashmla() { fi print_step "Installing FlashMLA from git (this may take 5-10 minutes)..." - FLASH_MLA_DISABLE_SM100=1 run_pip install "git+https://github.com/deepseek-ai/FlashMLA.git@${FLASHMLA_COMMIT}" --no-build-isolation + FLASH_MLA_DISABLE_SM100=1 pip install "git+https://github.com/deepseek-ai/FlashMLA.git@${FLASHMLA_COMMIT}" --no-build-isolation print_success "FlashMLA installed" } @@ -214,7 +206,7 @@ install_deepgemm() { fi print_step "Building DeepGEMM (this may take 5-10 minutes)..." - run_pip install . --no-build-isolation + pip install . --no-build-isolation print_success "DeepGEMM installed" } @@ -227,7 +219,7 @@ install_batchgen_kernels() { if [[ -f "$BATCHGEN_DIR/batchgen_kernels/setup.py" ]]; then cd "$BATCHGEN_DIR/batchgen_kernels" - run_pip install . --no-build-isolation + pip install . --no-build-isolation print_success "batchgen_kernels installed" else print_warning "batchgen_kernels/setup.py not found, skipping kernel compilation" @@ -243,7 +235,7 @@ install_batchgen() { if [[ -f "$BATCHGEN_DIR/setup.py" ]]; then cd "$BATCHGEN_DIR" - run_pip install . + pip install . print_success "BatchGen installed" else print_error "Could not find BatchGen setup.py at $BATCHGEN_DIR" @@ -376,9 +368,9 @@ main() { if [[ $IS_HOPPER -eq 1 ]]; then if [[ -n "$WHEEL_DIR" && -d "$WHEEL_DIR" ]]; then print_step "Installing Hopper dependencies from pre-built wheels: $WHEEL_DIR" - run_pip install --find-links "$WHEEL_DIR" --no-index \ + pip install --find-links "$WHEEL_DIR" --no-index \ flash-attn-hopper flash-mla deep-gemm 2>/dev/null || \ - run_pip install "$WHEEL_DIR"/*.whl + pip install "$WHEEL_DIR"/*.whl print_success "Hopper dependencies installed from wheels" else install_flash_attention @@ -386,7 +378,7 @@ main() { install_deepgemm # Reinstall PyTorch — building deps from source may downgrade torch or triton print_step "Reinstalling PyTorch to ensure correct version after dependency builds..." - run_pip install torch==2.9.0+cu128 --index-url https://download.pytorch.org/whl/cu128 + pip install torch==2.9.0+cu128 --index-url https://download.pytorch.org/whl/cu128 print_success "PyTorch reinstalled" fi else From 42c997e8c7c4340a18e27d4a9c9cfe31907dc0ec Mon Sep 17 00:00:00 2001 From: luzhan <513964121@qq.com> Date: Fri, 6 Mar 2026 16:13:56 +0000 Subject: [PATCH 11/17] fix: lazily import server worker entrypoint --- batchgen/server/worker_manager.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/batchgen/server/worker_manager.py b/batchgen/server/worker_manager.py index 4cbd62919..7742a475a 100644 --- a/batchgen/server/worker_manager.py +++ b/batchgen/server/worker_manager.py @@ -27,7 +27,6 @@ get_model_byte_size, ) from batchgen.server.server_args import ServerArgs -from batchgen.server_worker_main_loop import server_worker_main from batchgen.utils import config_torch_module_initializer logger = logging.getLogger(__name__) @@ -35,6 +34,14 @@ PARAMETER_SERVER_ENDPOINT_ENV = "BATCHGEN_PARAMETER_SERVER_ENDPOINT" +def _load_server_worker_main(): + # Delay this import to avoid a package-init cycle when worker subprocesses + # import `batchgen.server.process_utils` via `batchgen.server_worker_main_loop`. + from batchgen.server_worker_main_loop import server_worker_main + + return server_worker_main + + def _validate_shmem_enabled() -> None: """Check that THP shmem is enabled for --fast-init. Raises RuntimeError if not.""" import re @@ -516,7 +523,7 @@ def _spawn_workers(self) -> None: weights_memfd_fd=self._get_weights_memfd_fd(), ) self.worker_process = mp.spawn( - server_worker_main, + _load_server_worker_main(), args=( self.request_queue, self.response_queue, From fde3312ccfe75397e3fa2597c7935a6d34686914 Mon Sep 17 00:00:00 2001 From: luzhan <513964121@qq.com> Date: Tue, 10 Mar 2026 11:52:46 +0000 Subject: [PATCH 12/17] fix: initialize empty gpu page table on idle ranks --- batchgen/batchgen_worker.py | 95 +++++++++++++++++++++++++++++++++++++ 1 file changed, 95 insertions(+) diff --git a/batchgen/batchgen_worker.py b/batchgen/batchgen_worker.py index 93baec799..5ab5bf625 100644 --- a/batchgen/batchgen_worker.py +++ b/batchgen/batchgen_worker.py @@ -510,6 +510,7 @@ def __init__(self, args: BatchGenWorkerArgs): self._core_initialized = False self._batch_completed = False self._nvshmem_initialized_this_run = False + self._prefix_cache_stats_at_batch_start = None # 10. Distributed Communication Info self.dist_init_addr = args.dist_init_addr @@ -1819,6 +1820,93 @@ def _get_host_kv_free_pages(self) -> int: stats = self.host_paged_kv_worker_view.get_stats() return stats.num_free_pages + def _snapshot_prefix_cache_stats(self) -> Optional[Dict[str, float]]: + """Capture shared PrefixCache counters for batch-level reporting.""" + worker_view = getattr(self, "host_paged_kv_worker_view", None) + if worker_view is None: + return None + + try: + stats = worker_view.get_stats() + except Exception as exc: + logging.warning( + f"Rank {self.rank}: Failed to read PrefixCache stats: {exc}" + ) + return None + + hits = int(stats.num_prefix_hits) + misses = int(stats.num_prefix_misses) + lookups = hits + misses + return { + "num_total_pages": int(stats.num_total_pages), + "num_free_pages": int(stats.num_free_pages), + "num_used_pages": int(stats.num_used_pages), + "num_prefix_entries": int(stats.num_prefix_entries), + "num_prefix_hits": hits, + "num_prefix_misses": misses, + "num_prefix_evictions": int(stats.num_prefix_evictions), + "num_cache_entry_pages": int(stats.num_cache_entry_pages), + "num_shared_pages": int(stats.num_shared_pages), + "hit_rate": (hits / lookups) if lookups > 0 else 0.0, + } + + def _log_prefix_cache_batch_report(self) -> None: + """Log PrefixCache cumulative and per-batch deltas after batch completion.""" + if self.rank != 0: + return + + after = self._snapshot_prefix_cache_stats() + before = self._prefix_cache_stats_at_batch_start + if after is None: + logging.info("[PrefixCache] Stats unavailable after batch completion") + return + + if before is None: + logging.info( + "[PrefixCache] cumulative: entries=%d hits=%d misses=%d " + "hit_rate=%.2f%% evictions=%d cache_entry_pages=%d shared_pages=%d " + "host_used_pages=%d/%d", + after["num_prefix_entries"], + after["num_prefix_hits"], + after["num_prefix_misses"], + after["hit_rate"] * 100.0, + after["num_prefix_evictions"], + after["num_cache_entry_pages"], + after["num_shared_pages"], + after["num_used_pages"], + after["num_total_pages"], + ) + return + + batch_hits = after["num_prefix_hits"] - before["num_prefix_hits"] + batch_misses = after["num_prefix_misses"] - before["num_prefix_misses"] + batch_evictions = after["num_prefix_evictions"] - before["num_prefix_evictions"] + batch_lookup_total = batch_hits + batch_misses + batch_hit_rate = (batch_hits / batch_lookup_total) if batch_lookup_total > 0 else 0.0 + + logging.info( + "[PrefixCache] batch: hits=%d misses=%d hit_rate=%.2f%% " + "entries_delta=%d cache_pages_delta=%d shared_pages_delta=%d evictions_delta=%d; " + "cumulative: entries=%d hits=%d misses=%d hit_rate=%.2f%% " + "evictions=%d cache_entry_pages=%d shared_pages=%d host_used_pages=%d/%d", + batch_hits, + batch_misses, + batch_hit_rate * 100.0, + after["num_prefix_entries"] - before["num_prefix_entries"], + after["num_cache_entry_pages"] - before["num_cache_entry_pages"], + after["num_shared_pages"] - before["num_shared_pages"], + batch_evictions, + after["num_prefix_entries"], + after["num_prefix_hits"], + after["num_prefix_misses"], + after["hit_rate"] * 100.0, + after["num_prefix_evictions"], + after["num_cache_entry_pages"], + after["num_shared_pages"], + after["num_used_pages"], + after["num_total_pages"], + ) + def _get_or_create_gloo_group(self): """Get or create a Gloo process group for CPU tensor migrations. @@ -2876,6 +2964,7 @@ def process_new_batch( per_sequence_max_tokens: Optional per-sequence max output token limits. Falls back to self.max_decoding_length if None or if individual entry is None. """ + self._prefix_cache_stats_at_batch_start = self._snapshot_prefix_cache_stats() logging.info( f"Rank {self.rank}: Processing global batch of {len(global_prompts)} sequences" ) @@ -4728,6 +4817,7 @@ def generate(self): # Compute and log batch statistics self._log_batch_statistics() + self._log_prefix_cache_batch_report() # ============ Gather Results in Original Order ============ # Detokenize locally on each rank to avoid gathering large token tensors. @@ -5257,6 +5347,11 @@ def _config_decoding_for_batch( # Mark initial reservation done seq.mark_initial_gpu_reservation_done() self._sequences_with_gpu_kv.add(uuid) + else: + # Idle ranks still participate in CUDA graph warmup/capture. Keep an + # explicit empty page table so graph segments can access KV caches + # without tripping over an uninitialized gpu_table. + self.gpu_paged_kv_cache_manager.clear_page_table() if self.rank == 0: logging.info(f"[DECODE] Config completed: {(time.perf_counter() - start_time)*1000:.1f}ms, {len(decode_uuids)} sequences") From 927d8924ad468aa7508134d56d93d05af3876ad6 Mon Sep 17 00:00:00 2001 From: luzhan <513964121@qq.com> Date: Tue, 10 Mar 2026 12:04:50 +0000 Subject: [PATCH 13/17] fix: keep prefix cache runtime config in sync --- batchgen/batchgen_worker.py | 21 ++++++++++++++++----- 1 file changed, 16 insertions(+), 5 deletions(-) diff --git a/batchgen/batchgen_worker.py b/batchgen/batchgen_worker.py index 5ab5bf625..c52a898f1 100644 --- a/batchgen/batchgen_worker.py +++ b/batchgen/batchgen_worker.py @@ -1396,6 +1396,14 @@ def _initialize_core_components(self, num_queries: int) -> None: ) self.core_engine.host_paged_kv_worker_view = self.host_paged_kv_worker_view + host_kv_runtime_cfg = getattr(self.engine_config, "Host_Paged_KV_Config", None) + if host_kv_runtime_cfg is not None: + host_kv_runtime_cfg.enable_prefix_reuse = bool( + getattr(worker_kv_config, "enable_prefix_reuse", False) + ) + host_kv_runtime_cfg.page_size = int( + getattr(worker_kv_config, "page_size_tokens", host_kv_runtime_cfg.page_size) + ) self.engine_config.Basic_Config.num_queries = num_queries # Set CUDA graph config from command-line args @@ -5096,10 +5104,14 @@ def _config_prefill_for_batch(self, prefill_uuids: List[str]) -> None: self.core_engine.host_paged_kv_worker_view.register_sequences(global_sequence_ids) host_cfg = getattr(self.engine_config, "Host_Paged_KV_Config", None) - enable_prefix_reuse = ( - host_cfg is not None - and getattr(host_cfg, "enable_prefix_reuse", False) - ) + enable_prefix_reuse = bool(getattr(self.args, "enable_prefix_cache", False)) + page_size = seq.PAGE_SIZE + if host_cfg is not None: + enable_prefix_reuse = ( + enable_prefix_reuse + or bool(getattr(host_cfg, "enable_prefix_reuse", False)) + ) + page_size = max(1, int(getattr(host_cfg, "page_size", page_size))) if enable_prefix_reuse: _, reused_prefix_tokens = self.core_engine.host_paged_kv_worker_view.allocate_pages_for_sequences_with_prefix( list(zip(global_sequence_ids, sequence_tokens)), @@ -5108,7 +5120,6 @@ def _config_prefill_for_batch(self, prefill_uuids: List[str]) -> None: ) reused_sequences = sum(1 for tokens in reused_prefix_tokens if tokens > 0) reused_tokens_total = sum(int(tokens) for tokens in reused_prefix_tokens) - page_size = max(1, int(getattr(host_cfg, "page_size_tokens", 1))) if reused_sequences > 0: logging.info( f"Rank {self.rank}: PrefixCache allocation summary " From 2d58cd8f5414714079f431d6e7191be5dd352f52 Mon Sep 17 00:00:00 2001 From: luzhan <513964121@qq.com> Date: Tue, 10 Mar 2026 12:14:31 +0000 Subject: [PATCH 14/17] fix: persist worker host kv config for runtime sync --- batchgen/batchgen_worker.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/batchgen/batchgen_worker.py b/batchgen/batchgen_worker.py index c52a898f1..b1bd4bf25 100644 --- a/batchgen/batchgen_worker.py +++ b/batchgen/batchgen_worker.py @@ -451,6 +451,7 @@ def __init__(self, args: BatchGenWorkerArgs): host_kv_cache_size=args.global_host_kv_cache_size_gb * (1024**3), enable_prefix_reuse=args.enable_prefix_cache, ) + self._worker_host_kv_config = worker_kv_config if args.fast_init: worker_kv_config.enable_memfd = True worker_kv_config.memfd_creator_pid = args.kv_memfd_pid @@ -1397,6 +1398,7 @@ def _initialize_core_components(self, num_queries: int) -> None: self.core_engine.host_paged_kv_worker_view = self.host_paged_kv_worker_view host_kv_runtime_cfg = getattr(self.engine_config, "Host_Paged_KV_Config", None) + worker_kv_config = getattr(self, "_worker_host_kv_config", None) if host_kv_runtime_cfg is not None: host_kv_runtime_cfg.enable_prefix_reuse = bool( getattr(worker_kv_config, "enable_prefix_reuse", False) From 1a5193997b4067575439d4e7de16233be0dd66a2 Mon Sep 17 00:00:00 2001 From: luzhan <513964121@qq.com> Date: Tue, 10 Mar 2026 15:39:31 +0000 Subject: [PATCH 15/17] test: expand prefix cache coverage --- test/paged_kv/test_host_prefix_reuse.py | 149 +++++++++++++++++++++ test/paged_kv/test_prefix_cache_binding.py | 39 ++++++ 2 files changed, 188 insertions(+) diff --git a/test/paged_kv/test_host_prefix_reuse.py b/test/paged_kv/test_host_prefix_reuse.py index 2532f98d4..efee63a96 100644 --- a/test/paged_kv/test_host_prefix_reuse.py +++ b/test/paged_kv/test_host_prefix_reuse.py @@ -218,6 +218,155 @@ def test_prefix_reuse_includes_decode_tokens() -> None: _shm_unlink(shm_name) +def test_prefix_reuse_waits_for_full_decode_page_before_commit() -> None: + if not torch.cuda.is_available(): + raise SkipTest("CUDA is not available") + + shm_name = _random_shm_name() + cfg = _make_mla_config( + shm_name, enable_prefix_reuse=True, page_size_tokens=8 + ) + worker = bg.MLAHostPagedKVWorkerView(cfg) + + try: + worker.initialize(device_index=0, create_region=True) + + prompt_tokens = [500 + i for i in range(cfg.page_size_tokens)] + decode_tokens = [600 + i for i in range(cfg.page_size_tokens - 1)] + full_tokens = prompt_tokens + decode_tokens + + worker.register_sequences([351]) + pages_1, reused_1 = worker.allocate_pages_for_sequences_with_prefix( + [(351, len(full_tokens))], prompt_tokens, [0, len(prompt_tokens)] + ) + assert reused_1 == [0] + assert len(pages_1[0]) == 2 + + prefill_k = torch.full( + (1, len(prompt_tokens), cfg.num_k_heads, cfg.k_head_dim), + fill_value=1.0, + dtype=torch.bfloat16, + device="cuda:0", + ) + worker.async_offload_layer_kv_to_host( + layer_idx=0, + sequence_ids=[351], + k_tensor=prefill_k, + v_tensor=None, + sequence_lengths=[len(prompt_tokens)], + ).wait() + + for step, token_id in enumerate(decode_tokens): + decode_k = torch.full( + (1, 1, cfg.num_k_heads, cfg.k_head_dim), + fill_value=2.0 + step, + dtype=torch.bfloat16, + device="cuda:0", + ) + worker.async_append_decode_kv_to_host( + layer_idx=0, + sequence_ids=[351], + k_tensor=decode_k, + v_tensor=None, + sequence_lengths=[len(prompt_tokens) + step], + decode_token_ids=[token_id], + ).wait() + + worker.release_sequence_pages([351]) + + worker.register_sequences([352]) + pages_2, reused_2 = worker.allocate_pages_for_sequences_with_prefix( + [(352, len(full_tokens))], full_tokens, [0, len(full_tokens)] + ) + + assert len(pages_2[0]) == 2 + assert reused_2 == [len(prompt_tokens)] + + stats = worker.get_stats() + assert stats.num_prefix_hits >= 1 + assert stats.num_shared_pages >= 1 + + worker.release_sequence_pages([352]) + finally: + worker.shutdown() + del worker + _shm_unlink(shm_name) + + +def test_prefix_reuse_skips_decode_extension_when_token_ids_missing() -> None: + if not torch.cuda.is_available(): + raise SkipTest("CUDA is not available") + + shm_name = _random_shm_name() + cfg = _make_mla_config( + shm_name, enable_prefix_reuse=True, page_size_tokens=8 + ) + worker = bg.MLAHostPagedKVWorkerView(cfg) + + try: + worker.initialize(device_index=0, create_region=True) + + prompt_tokens = [700 + i for i in range(cfg.page_size_tokens)] + decode_tokens = [800 + i for i in range(cfg.page_size_tokens)] + full_tokens = prompt_tokens + decode_tokens + + worker.register_sequences([361]) + pages_1, reused_1 = worker.allocate_pages_for_sequences_with_prefix( + [(361, len(full_tokens))], prompt_tokens, [0, len(prompt_tokens)] + ) + assert reused_1 == [0] + assert len(pages_1[0]) == 2 + + prefill_k = torch.full( + (1, len(prompt_tokens), cfg.num_k_heads, cfg.k_head_dim), + fill_value=1.0, + dtype=torch.bfloat16, + device="cuda:0", + ) + worker.async_offload_layer_kv_to_host( + layer_idx=0, + sequence_ids=[361], + k_tensor=prefill_k, + v_tensor=None, + sequence_lengths=[len(prompt_tokens)], + ).wait() + + for step in range(len(decode_tokens)): + decode_k = torch.full( + (1, 1, cfg.num_k_heads, cfg.k_head_dim), + fill_value=3.0 + step, + dtype=torch.bfloat16, + device="cuda:0", + ) + worker.async_append_decode_kv_to_host( + layer_idx=0, + sequence_ids=[361], + k_tensor=decode_k, + v_tensor=None, + sequence_lengths=[len(prompt_tokens) + step], + ).wait() + + worker.release_sequence_pages([361]) + + worker.register_sequences([362]) + pages_2, reused_2 = worker.allocate_pages_for_sequences_with_prefix( + [(362, len(full_tokens))], full_tokens, [0, len(full_tokens)] + ) + + assert len(pages_2[0]) == 2 + assert reused_2 == [len(prompt_tokens)] + + stats = worker.get_stats() + assert stats.num_prefix_hits >= 1 + assert stats.num_shared_pages >= 1 + + worker.release_sequence_pages([362]) + finally: + worker.shutdown() + del worker + _shm_unlink(shm_name) + + def test_prefix_batch_allocation_failure_rolls_back() -> None: if not torch.cuda.is_available(): raise SkipTest("CUDA is not available") diff --git a/test/paged_kv/test_prefix_cache_binding.py b/test/paged_kv/test_prefix_cache_binding.py index 3733d6c79..5677790c3 100644 --- a/test/paged_kv/test_prefix_cache_binding.py +++ b/test/paged_kv/test_prefix_cache_binding.py @@ -135,3 +135,42 @@ def test_prefix_cache_harness_duplicate_commit_avoids_spurious_evict() -> None: assert stats_after.prefix_evict_count == 0 assert cache.page_refcount(0) == 1 assert cache.page_refcount(1) == 1 + + +def test_prefix_cache_harness_lookup_refreshes_lru_recency() -> None: + cache = _make_prefix_cache(num_pages=12, prefix_page_budget=4) + tokens_a = [6000 + i for i in range(128)] + tokens_b = [7000 + i for i in range(128)] + tokens_c = [8000 + i for i in range(128)] + + assert cache.commit(tokens_a, [0, 1]) is True + assert cache.commit(tokens_b, [2, 3]) is True + + pages, reused = cache.lookup(tokens_a, 2) + assert pages == [0, 1] + assert reused == 2 + + assert cache.commit(tokens_c, [4, 5]) is True + stats = cache.get_stats() + assert stats.prefix_evict_count >= 1 + assert stats.prefix_entry_count == 2 + assert stats.prefix_used_pages == 4 + + pages, reused = cache.lookup(tokens_a, 2) + assert pages == [0, 1] + assert reused == 2 + + pages, reused = cache.lookup(tokens_b, 2) + assert pages == [] + assert reused == 0 + + pages, reused = cache.lookup(tokens_c, 2) + assert pages == [4, 5] + assert reused == 2 + + assert cache.page_refcount(0) == 1 + assert cache.page_refcount(1) == 1 + assert cache.page_refcount(2) == 0 + assert cache.page_refcount(3) == 0 + assert cache.page_refcount(4) == 1 + assert cache.page_refcount(5) == 1 From 7adf83488989c7b2f702a1836c0186a0f97ffd41 Mon Sep 17 00:00:00 2001 From: luzhan <513964121@qq.com> Date: Tue, 10 Mar 2026 15:46:22 +0000 Subject: [PATCH 16/17] test: add multiprocess prefix reuse coverage --- test/paged_kv/test_host_prefix_reuse.py | 339 ++++++++++++++++++++++++ 1 file changed, 339 insertions(+) diff --git a/test/paged_kv/test_host_prefix_reuse.py b/test/paged_kv/test_host_prefix_reuse.py index efee63a96..67dc38320 100644 --- a/test/paged_kv/test_host_prefix_reuse.py +++ b/test/paged_kv/test_host_prefix_reuse.py @@ -1,7 +1,10 @@ import ctypes import errno +import multiprocessing as mp +import queue import random import string +import traceback from unittest import SkipTest import torch @@ -55,6 +58,289 @@ def _make_mla_config( return cfg +def _put_mp_success(result_queue, role: str, **payload) -> None: + result_queue.put({"role": role, "ok": True, **payload}) + + +def _put_mp_error(result_queue, role: str) -> None: + result_queue.put({"role": role, "ok": False, "error": traceback.format_exc()}) + + +def _prefix_prompt_creator_proc( + shm_name: str, + prompt_tokens: list[int], + result_queue, + ready_event, + done_event, +) -> None: + worker = None + try: + torch.cuda.set_device(0) + cfg = _make_mla_config(shm_name, enable_prefix_reuse=True) + worker = bg.MLAHostPagedKVWorkerView(cfg) + worker.initialize(device_index=0, create_region=True) + + worker.register_sequences([501]) + pages, reused = worker.allocate_pages_for_sequences_with_prefix( + [(501, len(prompt_tokens))], prompt_tokens, [0, len(prompt_tokens)] + ) + assert reused == [0] + assert len(pages[0]) == 1 + + prefill_k = torch.full( + (1, len(prompt_tokens), cfg.num_k_heads, cfg.k_head_dim), + fill_value=1.0, + dtype=torch.bfloat16, + device="cuda:0", + ) + worker.async_offload_layer_kv_to_host( + layer_idx=0, + sequence_ids=[501], + k_tensor=prefill_k, + v_tensor=None, + sequence_lengths=[len(prompt_tokens)], + ).wait() + worker.release_sequence_pages([501]) + + stats = worker.get_stats() + _put_mp_success( + result_queue, + "creator", + prefix_entries=stats.num_prefix_entries, + shared_pages=stats.num_shared_pages, + ) + ready_event.set() + + if not done_event.wait(timeout=30): + raise TimeoutError("Timed out waiting for attacher process") + except Exception: + ready_event.set() + _put_mp_error(result_queue, "creator") + raise + finally: + if worker is not None: + worker.shutdown() + del worker + + +def _prefix_prompt_attacher_proc( + shm_name: str, + prompt_tokens: list[int], + result_queue, + ready_event, + done_event, +) -> None: + worker = None + try: + if not ready_event.wait(timeout=30): + raise TimeoutError("Timed out waiting for creator process") + + torch.cuda.set_device(0) + cfg = _make_mla_config(shm_name, enable_prefix_reuse=True) + worker = bg.MLAHostPagedKVWorkerView(cfg) + worker.initialize(device_index=0, create_region=False) + + worker.register_sequences([502]) + pages, reused = worker.allocate_pages_for_sequences_with_prefix( + [(502, len(prompt_tokens))], prompt_tokens, [0, len(prompt_tokens)] + ) + stats = worker.get_stats() + worker.release_sequence_pages([502]) + + _put_mp_success( + result_queue, + "attacher", + reused_tokens=reused, + allocated_pages=len(pages[0]), + prefix_hits=stats.num_prefix_hits, + shared_pages=stats.num_shared_pages, + ) + done_event.set() + except Exception: + done_event.set() + _put_mp_error(result_queue, "attacher") + raise + finally: + if worker is not None: + worker.shutdown() + del worker + + +def _prefix_decode_creator_proc( + shm_name: str, + page_size_tokens: int, + prompt_tokens: list[int], + decode_tokens: list[int], + result_queue, + ready_event, + done_event, +) -> None: + worker = None + try: + if len(prompt_tokens) != page_size_tokens or len(decode_tokens) != page_size_tokens: + raise AssertionError("decode multiprocess helper expects exact full-page prompt/decode") + + torch.cuda.set_device(0) + cfg = _make_mla_config( + shm_name, enable_prefix_reuse=True, page_size_tokens=page_size_tokens + ) + worker = bg.MLAHostPagedKVWorkerView(cfg) + worker.initialize(device_index=0, create_region=True) + + full_tokens = prompt_tokens + decode_tokens + worker.register_sequences([601]) + pages, reused = worker.allocate_pages_for_sequences_with_prefix( + [(601, len(full_tokens))], prompt_tokens, [0, len(prompt_tokens)] + ) + assert reused == [0] + assert len(pages[0]) == 2 + + prefill_k = torch.full( + (1, len(prompt_tokens), cfg.num_k_heads, cfg.k_head_dim), + fill_value=1.0, + dtype=torch.bfloat16, + device="cuda:0", + ) + worker.async_offload_layer_kv_to_host( + layer_idx=0, + sequence_ids=[601], + k_tensor=prefill_k, + v_tensor=None, + sequence_lengths=[len(prompt_tokens)], + ).wait() + + for step, token_id in enumerate(decode_tokens): + decode_k = torch.full( + (1, 1, cfg.num_k_heads, cfg.k_head_dim), + fill_value=2.0 + step, + dtype=torch.bfloat16, + device="cuda:0", + ) + worker.async_append_decode_kv_to_host( + layer_idx=0, + sequence_ids=[601], + k_tensor=decode_k, + v_tensor=None, + sequence_lengths=[len(prompt_tokens) + step], + decode_token_ids=[token_id], + ).wait() + + worker.release_sequence_pages([601]) + stats = worker.get_stats() + _put_mp_success( + result_queue, + "creator", + prefix_entries=stats.num_prefix_entries, + shared_pages=stats.num_shared_pages, + ) + ready_event.set() + + if not done_event.wait(timeout=30): + raise TimeoutError("Timed out waiting for attacher process") + except Exception: + ready_event.set() + _put_mp_error(result_queue, "creator") + raise + finally: + if worker is not None: + worker.shutdown() + del worker + + +def _prefix_decode_attacher_proc( + shm_name: str, + page_size_tokens: int, + full_tokens: list[int], + result_queue, + ready_event, + done_event, +) -> None: + worker = None + try: + if not ready_event.wait(timeout=30): + raise TimeoutError("Timed out waiting for creator process") + + torch.cuda.set_device(0) + cfg = _make_mla_config( + shm_name, enable_prefix_reuse=True, page_size_tokens=page_size_tokens + ) + worker = bg.MLAHostPagedKVWorkerView(cfg) + worker.initialize(device_index=0, create_region=False) + + worker.register_sequences([602]) + pages, reused = worker.allocate_pages_for_sequences_with_prefix( + [(602, len(full_tokens))], full_tokens, [0, len(full_tokens)] + ) + stats = worker.get_stats() + worker.release_sequence_pages([602]) + + _put_mp_success( + result_queue, + "attacher", + reused_tokens=reused, + allocated_pages=len(pages[0]), + prefix_hits=stats.num_prefix_hits, + shared_pages=stats.num_shared_pages, + ) + done_event.set() + except Exception: + done_event.set() + _put_mp_error(result_queue, "attacher") + raise + finally: + if worker is not None: + worker.shutdown() + del worker + + +def _run_mp_pair(creator_target, creator_args, attacher_target, attacher_args): + ctx = mp.get_context("spawn") + result_queue = ctx.Queue() + ready_event = ctx.Event() + done_event = ctx.Event() + + creator = ctx.Process( + target=creator_target, + args=(*creator_args, result_queue, ready_event, done_event), + ) + attacher = ctx.Process( + target=attacher_target, + args=(*attacher_args, result_queue, ready_event, done_event), + ) + + creator.start() + attacher.start() + + creator.join(timeout=60) + attacher.join(timeout=60) + + if creator.is_alive(): + creator.terminate() + creator.join(timeout=5) + raise AssertionError("creator process timed out") + if attacher.is_alive(): + attacher.terminate() + attacher.join(timeout=5) + raise AssertionError("attacher process timed out") + + results = {} + for _ in range(2): + try: + item = result_queue.get(timeout=5) + except queue.Empty: + break + results[item["role"]] = item + + if creator.exitcode != 0: + raise AssertionError(results.get("creator", {}).get("error", "creator process failed")) + if attacher.exitcode != 0: + raise AssertionError(results.get("attacher", {}).get("error", "attacher process failed")) + + assert results["creator"]["ok"] is True + assert results["attacher"]["ok"] is True + return results + + def test_prefix_reuse_disabled_regression() -> None: if not torch.cuda.is_available(): raise SkipTest("CUDA is not available") @@ -367,6 +653,59 @@ def test_prefix_reuse_skips_decode_extension_when_token_ids_missing() -> None: _shm_unlink(shm_name) +def test_prefix_reuse_across_process_attach() -> None: + if not torch.cuda.is_available(): + raise SkipTest("CUDA is not available") + + shm_name = _random_shm_name() + prompt_tokens = [900 + i for i in range(64)] + + try: + results = _run_mp_pair( + _prefix_prompt_creator_proc, + (shm_name, prompt_tokens), + _prefix_prompt_attacher_proc, + (shm_name, prompt_tokens), + ) + + assert results["creator"]["prefix_entries"] >= 1 + assert results["creator"]["shared_pages"] >= 1 + assert results["attacher"]["reused_tokens"] == [64] + assert results["attacher"]["allocated_pages"] == 1 + assert results["attacher"]["prefix_hits"] >= 1 + assert results["attacher"]["shared_pages"] >= 1 + finally: + _shm_unlink(shm_name) + + +def test_prefix_reuse_decode_extension_across_process_attach() -> None: + if not torch.cuda.is_available(): + raise SkipTest("CUDA is not available") + + shm_name = _random_shm_name() + page_size_tokens = 8 + prompt_tokens = [1000 + i for i in range(page_size_tokens)] + decode_tokens = [1100 + i for i in range(page_size_tokens)] + full_tokens = prompt_tokens + decode_tokens + + try: + results = _run_mp_pair( + _prefix_decode_creator_proc, + (shm_name, page_size_tokens, prompt_tokens, decode_tokens), + _prefix_decode_attacher_proc, + (shm_name, page_size_tokens, full_tokens), + ) + + assert results["creator"]["prefix_entries"] >= 1 + assert results["creator"]["shared_pages"] >= 1 + assert results["attacher"]["reused_tokens"] == [len(full_tokens)] + assert results["attacher"]["allocated_pages"] == 2 + assert results["attacher"]["prefix_hits"] >= 1 + assert results["attacher"]["shared_pages"] >= 1 + finally: + _shm_unlink(shm_name) + + def test_prefix_batch_allocation_failure_rolls_back() -> None: if not torch.cuda.is_available(): raise SkipTest("CUDA is not available") From 551d05a59a0effec00531e5c46d1b62474bfe683 Mon Sep 17 00:00:00 2001 From: luzhan <513964121@qq.com> Date: Tue, 10 Mar 2026 15:48:12 +0000 Subject: [PATCH 17/17] test: fix multiprocess prefix stats assertions --- test/paged_kv/test_host_prefix_reuse.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/test/paged_kv/test_host_prefix_reuse.py b/test/paged_kv/test_host_prefix_reuse.py index 67dc38320..e96745b11 100644 --- a/test/paged_kv/test_host_prefix_reuse.py +++ b/test/paged_kv/test_host_prefix_reuse.py @@ -107,6 +107,7 @@ def _prefix_prompt_creator_proc( result_queue, "creator", prefix_entries=stats.num_prefix_entries, + cache_entry_pages=stats.num_cache_entry_pages, shared_pages=stats.num_shared_pages, ) ready_event.set() @@ -231,6 +232,7 @@ def _prefix_decode_creator_proc( result_queue, "creator", prefix_entries=stats.num_prefix_entries, + cache_entry_pages=stats.num_cache_entry_pages, shared_pages=stats.num_shared_pages, ) ready_event.set() @@ -669,7 +671,7 @@ def test_prefix_reuse_across_process_attach() -> None: ) assert results["creator"]["prefix_entries"] >= 1 - assert results["creator"]["shared_pages"] >= 1 + assert results["creator"]["cache_entry_pages"] >= 1 assert results["attacher"]["reused_tokens"] == [64] assert results["attacher"]["allocated_pages"] == 1 assert results["attacher"]["prefix_hits"] >= 1 @@ -697,7 +699,7 @@ def test_prefix_reuse_decode_extension_across_process_attach() -> None: ) assert results["creator"]["prefix_entries"] >= 1 - assert results["creator"]["shared_pages"] >= 1 + assert results["creator"]["cache_entry_pages"] >= 2 assert results["attacher"]["reused_tokens"] == [len(full_tokens)] assert results["attacher"]["allocated_pages"] == 2 assert results["attacher"]["prefix_hits"] >= 1