diff --git a/.gitignore b/.gitignore index 35ab91011..ad224bcdb 100644 --- a/.gitignore +++ b/.gitignore @@ -4,6 +4,7 @@ log_*.log log *.log +tmp/ *.pt runner.sh diff --git a/batchgen/batchgen_worker.py b/batchgen/batchgen_worker.py index f9df07409..75588e499 100644 --- a/batchgen/batchgen_worker.py +++ b/batchgen/batchgen_worker.py @@ -359,6 +359,7 @@ class BatchGenWorkerArgs: disable_cuda_graphs: bool = True # Disable CUDA graph capture for decode attention (default: off due to 128K+ crash) 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 detokenization_include_special_tokens: bool = False # When True, include special tokens in detokenized output # Dynamic host KV reservation host_kv_chunk_size: int = 8192 # Initial host KV chunk size in tokens @@ -420,6 +421,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}" @@ -526,6 +528,7 @@ def __init__(self, args: BatchGenWorkerArgs): model_name=args.model_name, host_kv_cache_size=host_budget_bytes, core_engine_module=core_engine, + enable_prefix_reuse=args.enable_prefix_cache, enable_memfd=args.fast_init, memfd_creator_pid=args.kv_memfd_pid if args.fast_init else -1, memfd_fd=args.kv_memfd_fd if args.fast_init else -1, @@ -540,7 +543,9 @@ def __init__(self, args: BatchGenWorkerArgs): worker_kv_config = build_host_kv_config( model_name=args.model_name, host_kv_cache_size=host_budget_bytes, + 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 @@ -601,6 +606,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 @@ -1846,7 +1852,12 @@ def _flush_deferred_kv_to_host(self) -> None: self._deferred_kv_entries_aux = [] return - sequence_ids, sequence_lengths = batch_info if batch_info is not None else (None, None) + sequence_ids = sequence_lengths = decode_token_ids = None + if batch_info is not None: + if len(batch_info) == 3: + sequence_ids, sequence_lengths, decode_token_ids = batch_info + else: + sequence_ids, sequence_lengths = batch_info # ONE sync for ALL layers across BOTH caches — the key optimization if not hasattr(self, '_kv_offload_event'): @@ -1882,6 +1893,7 @@ def _flush_deferred_kv_to_host(self) -> None: entries=_prepared_entries, sequence_ids=sequence_ids, sequence_lengths=sequence_lengths, + decode_token_ids=decode_token_ids, ) if task is not None: self._pending_kv_append_tasks.append(task) @@ -1916,6 +1928,7 @@ def _flush_deferred_kv_to_host(self) -> None: k_tensor=k_tensor, v_tensor=v_tensor, sequence_lengths=sequence_lengths, + decode_token_ids=decode_token_ids, ) self._pending_kv_append_tensors.append(k_tensor) @@ -1984,6 +1997,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 = [] @@ -1995,6 +2010,31 @@ 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) + 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] + ): + 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: @@ -2004,6 +2044,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: @@ -2038,17 +2083,16 @@ def _append_decode_kv_to_host_async( f"Rank {self.rank}: NaN detected in k_tensor BEFORE host append (layer={layer_idx}) - affected_seqs={nan_seq_info}" ) - # Launch async D2H append — no CPU-side sync needed here. - # The C++ side runs on a background thread with its own D2H stream. - # Tensor references are kept alive in _pending_kv_append_tensors to - # prevent GC/memory reuse. All tasks are waited at decision boundary - # via _wait_pending_kv_append_tasks(). + # Launch async D2H append. The C++ side waits for the producer + # stream and keeps tensor references alive through the task lifetime. + 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, ) if not hasattr(self, '_pending_kv_append_tensors'): @@ -2209,6 +2253,13 @@ def _initialize_core_components(self, num_queries: int) -> None: self.host_paged_kv_worker_view_aux = self.host_paged_kv_worker_view.auxiliary else: 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(self.args, "enable_prefix_cache", False)) + host_view_config = getattr(self.core_engine.host_paged_kv_worker_view, "config", None) + host_kv_runtime_cfg.page_size = int( + getattr(host_view_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 @@ -2622,6 +2673,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. @@ -3485,6 +3623,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" ) @@ -5770,6 +5909,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. @@ -6064,6 +6204,8 @@ 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] chunk_size = self._get_effective_chunk_size() for uuid in my_prefill_uuids: @@ -6084,6 +6226,10 @@ def _config_prefill_for_batch(self, prefill_uuids: List[str]) -> None: seq.host_token_capacity = initial_capacity seq.host_pages_allocated = math.ceil(initial_capacity / seq.PAGE_SIZE) + 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)) + # Safety assertion: log if selection over-admitted. This should not # happen after the EVICTED-length fix in _prepare_prefill_batch — # if it fires, there's another selection bug to investigate. @@ -6113,9 +6259,38 @@ def _config_prefill_for_batch(self, prefill_uuids: List[str]) -> None: ) 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 = 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)), + 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) + 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)) + ) # DSA: mirror registration on auxiliary host KV aux_view = getattr(self, "host_paged_kv_worker_view_aux", None) if aux_view is not None: @@ -6364,6 +6539,11 @@ def _config_decoding_for_batch( f"sequences. Clearing local_decode_indices to avoid inconsistent state." ) local_decode_indices.clear() + 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") @@ -9102,12 +9282,37 @@ def decoding_continuous( if _kv_worker_view is not None: _kv_seq_ids = [] _kv_seq_lengths = [] + _kv_decode_token_ids = [] + _kv_has_decode_tokens = True for local_idx in current_batch: uuid = self._local_to_uuid_map[local_idx] seq = self.global_batch.get_sequence(uuid) _kv_seq_ids.append(seq.global_idx) - _kv_seq_lengths.append(seq.current_context_length - 1) - self._deferred_kv_batch = (_kv_seq_ids, _kv_seq_lengths) + write_pos = seq.current_context_length - 1 + _kv_seq_lengths.append(write_pos) + decode_pos = write_pos - seq.prompt_length + if decode_pos >= 0: + query_entry = self.query_book.get(local_idx) if self.query_book is not None else None + 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] + ): + _kv_has_decode_tokens = False + else: + _kv_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] + ): + _kv_has_decode_tokens = False + else: + _kv_decode_token_ids.append(int(seq.input_ids[0, write_pos].item())) + _kv_decode_token_ids_arg = _kv_decode_token_ids if _kv_has_decode_tokens else None + self._deferred_kv_batch = (_kv_seq_ids, _kv_seq_lengths, _kv_decode_token_ids_arg) self._deferred_kv_entries = [] self._deferred_kv_entries_aux = [] self._deferred_kv_worker_view = _kv_worker_view @@ -9117,6 +9322,7 @@ def decoding_continuous( # SYNC MODE: Immediately write each layer's KV to host (no deferral) _sync_kv_seq_ids = _kv_seq_ids _sync_kv_seq_lengths = _kv_seq_lengths + _sync_kv_decode_token_ids = _kv_decode_token_ids_arg _sync_kv_worker_view = _kv_worker_view def kv_append_callback(layer_idx: int, k_tensor: torch.Tensor, v_tensor: torch.Tensor = None): if k_tensor.dim() == 3: @@ -9130,6 +9336,7 @@ def kv_append_callback(layer_idx: int, k_tensor: torch.Tensor, v_tensor: torch.T k_tensor=k_tensor, v_tensor=v_tensor, sequence_lengths=_sync_kv_seq_lengths, + decode_token_ids=_sync_kv_decode_token_ids, ) if task is not None: task.wait() @@ -9530,12 +9737,45 @@ 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) + 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] + ): + 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) @@ -9566,15 +9806,16 @@ def _append_decode_kv_to_host_async( f"affected_seqs={nan_seq_info}" ) - # Launch async D2H append — no CPU-side sync needed. - # Tensor references kept alive in _pending_kv_append_tensors. - # All tasks waited at decision boundary via _wait_pending_kv_append_tasks(). + # Launch async D2H append. The C++ side waits for the producer stream; + # tensor references are kept alive until pending tasks are flushed. + 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, ) # Store tensor references alongside task to prevent GC/memory reuse diff --git a/batchgen/config/config.py b/batchgen/config/config.py index 3e14001ef..13c164707 100644 --- a/batchgen/config/config.py +++ b/batchgen/config/config.py @@ -176,6 +176,45 @@ 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 + # 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 + prefix_entry_capacity: int = 0 + 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: 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/dual_host_kv_coordinator.py b/batchgen/kv_cache/dual_host_kv_coordinator.py index 73ae6ba18..20e98122d 100644 --- a/batchgen/kv_cache/dual_host_kv_coordinator.py +++ b/batchgen/kv_cache/dual_host_kv_coordinator.py @@ -34,7 +34,9 @@ def _try_set_logger_name(config, name: str) -> bool: return False -def _build_host_config_from_profile(profile, shm_name: str, num_pages: int) -> Any: +def _build_host_config_from_profile( + profile, shm_name: str, num_pages: int, enable_prefix_reuse: bool = True +) -> Any: """Build a bg_lib.HostPagedKVConfig from a _HostKVModelProfile.""" from batchgen.kv_cache.host_kv_mananger_config import _dtype_size_bytes @@ -55,6 +57,9 @@ def _build_host_config_from_profile(profile, shm_name: str, num_pages: int) -> A profile.sequence_table_capacity or config.num_pages ) config.alignment_bytes = profile.alignment_bytes + config.enable_prefix_reuse = bool(enable_prefix_reuse) + config.prefix_min_reuse_pages = 1 + config.prefix_min_store_pages = 2 return config @@ -107,6 +112,7 @@ def from_budget( model_name: str, host_kv_cache_size: int, core_engine_module, + enable_prefix_reuse: bool = True, enable_memfd: bool = False, memfd_creator_pid: int = -1, memfd_fd: int = -1, @@ -127,10 +133,10 @@ def from_budget( ) primary_config = _build_host_config_from_profile( - primary_profile, HOST_KV_SHM_NAME, num_pages, + primary_profile, HOST_KV_SHM_NAME, num_pages, enable_prefix_reuse, ) aux_config = _build_host_config_from_profile( - aux_profile, HOST_KV_AUX_SHM_NAME, num_pages, + aux_profile, HOST_KV_AUX_SHM_NAME, num_pages, enable_prefix_reuse, ) # Set distinct logger names to avoid C++ logger name collision @@ -177,6 +183,7 @@ def create_managers( cls, model_name: str, host_kv_cache_size: int, + enable_prefix_reuse: bool = True, enable_memfd: bool = False, ) -> Optional[Tuple[Any, Any]]: """Server-side factory: create and initialize both host KV managers. @@ -194,10 +201,10 @@ def create_managers( ) primary_config = _build_host_config_from_profile( - primary_profile, HOST_KV_SHM_NAME, num_pages, + primary_profile, HOST_KV_SHM_NAME, num_pages, enable_prefix_reuse, ) aux_config = _build_host_config_from_profile( - aux_profile, HOST_KV_AUX_SHM_NAME, num_pages, + aux_profile, HOST_KV_AUX_SHM_NAME, num_pages, enable_prefix_reuse, ) # Set distinct logger names to avoid C++ logger name collision diff --git a/batchgen/kv_cache/host_kv_mananger_config.py b/batchgen/kv_cache/host_kv_mananger_config.py index 15d1290ab..99eebe505 100644 --- a/batchgen/kv_cache/host_kv_mananger_config.py +++ b/batchgen/kv_cache/host_kv_mananger_config.py @@ -227,7 +227,37 @@ def _resolve_profile(model_name: str) -> _HostKVModelProfile: return _PROFILE_REGISTRY[_PROFILE_ALIASES[alias]] -def build_host_kv_config(model_name: str, host_kv_cache_size: int) -> Any: +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, + enable_prefix_reuse: bool = True, +) -> Any: """Builds a core HostPagedKVConfig for the given model and host budget.""" if host_kv_cache_size is None: @@ -271,6 +301,10 @@ 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 = bool(enable_prefix_reuse) + config.prefix_min_reuse_pages = 1 + config.prefix_min_store_pages = 2 + _apply_capacity_defaults(config) return config @@ -356,7 +390,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) @@ -386,6 +424,10 @@ 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 = bool(enable_prefix_reuse) + config.prefix_min_reuse_pages = 1 + config.prefix_min_store_pages = 2 + _apply_capacity_defaults(config) return config diff --git a/batchgen/server/server_args.py b/batchgen/server/server_args.py index ff7d81e9d..8ff5f84bb 100644 --- a/batchgen/server/server_args.py +++ b/batchgen/server/server_args.py @@ -92,6 +92,7 @@ class ServerArgs: disable_cuda_graphs: bool = True # Disable CUDA graph capture for decode attention (128K+ crash: corrupted num_tokens_per_rank) 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 detokenization_include_special_tokens: bool = False # When True, include special tokens in detokenized output # Dynamic host KV reservation settings host_kv_chunk_size: int = 8192 # Initial host KV chunk size in tokens (default: 8K) @@ -349,6 +350,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", + ) parser.add_argument( "--detokenization-include-special-tokens", action="store_true", @@ -546,6 +560,7 @@ def prepare_server_args(argv: Optional[list[str]] = None) -> ServerArgs: cuda_graph_max_bucket_size=parsed.cuda_graph_max_bucket_size, cuda_graph_num_buckets=parsed.cuda_graph_num_buckets, detokenization_include_special_tokens=parsed.detokenization_include_special_tokens, + 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 e49f674bd..ca49b5f7b 100644 --- a/batchgen/server/worker_manager.py +++ b/batchgen/server/worker_manager.py @@ -35,6 +35,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 @@ -204,6 +212,7 @@ def _diag(msg): _diag(">>> allocate_host_kv_cache") result = 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, ) _diag("<<< allocate_host_kv_cache") @@ -629,6 +638,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, detokenization_include_special_tokens=self.args.detokenization_include_special_tokens, host_kv_chunk_size=self.args.host_kv_chunk_size, enable_host_kv_eviction=self.args.enable_host_kv_eviction, @@ -648,7 +658,7 @@ def _spawn_workers(self) -> None: ) from batchgen.server_worker_main_loop import server_worker_main self.worker_process = mp.spawn( - server_worker_main, + _load_server_worker_main(), args=( self.request_queue, self.response_queue, @@ -1005,6 +1015,7 @@ 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: from batchgen.kv_cache.dual_host_kv_coordinator import DualHostKVCoordinator @@ -1013,6 +1024,7 @@ def allocate_host_kv_cache( dual = DualHostKVCoordinator.create_managers( model_name=model_name, host_kv_cache_size=int(host_kv_cache_size_gb * (1024**3)), + enable_prefix_reuse=enable_prefix_cache, enable_memfd=enable_memfd, ) if dual is not None: @@ -1025,6 +1037,7 @@ def allocate_host_kv_cache( 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 diff --git a/core/KV_Storage/host_paged_kv_backend.cpp b/core/KV_Storage/host_paged_kv_backend.cpp index 51e6c0d52..44f7c315f 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 @@ -14,7 +15,6 @@ #include #include #include -#include #include #include #include @@ -30,7 +30,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; @@ -84,7 +84,7 @@ class ScopedMutexLock { ~ScopedMutexLock() { const int rc = pthread_mutex_unlock(mu_); if (rc != 0) { - std::terminate(); // Unlock failure is irrecoverable here. + std::terminate(); } } @@ -92,13 +92,54 @@ 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{ static_cast(InitState::kUninitialized)}; @@ -112,10 +153,35 @@ struct SharedHeader { std::uint64_t num_layers = 0; std::uint64_t page_size_tokens = 0; std::uint32_t has_v_cache = 0; + std::uint32_t enable_prefix_reuse = 0; + + 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; + 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}; + std::atomic active_sequences{0}; - pthread_mutex_t allocation_mutex{}; - pthread_mutex_t sequence_mutex{}; + 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 = kInvalidIndex; + std::int32_t lru_tail = kInvalidIndex; + + pthread_mutex_t metadata_mutex{}; }; std::size_t SafeHardwareConcurrency() { @@ -147,43 +213,6 @@ void TouchPagesMultiThreaded(void* ptr, std::size_t size, std::size_t stride) { } } -// 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; - } - - uintptr_t raw_addr = reinterpret_cast(addr); - uintptr_t aligned_addr = (raw_addr + alignment - 1) & ~(alignment - 1); - void* final_addr = reinterpret_cast(aligned_addr); - - // 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); - } - - return ret; -} - -} // namespace - struct HostPagedKVBackend::SharedState { explicit SharedState(const HostPagedKVConfig& cfg, std::size_t data_bytes, std::uint64_t fingerprint, bool has_v) @@ -194,15 +223,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; } @@ -216,11 +258,24 @@ 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; @@ -229,23 +284,61 @@ struct HostPagedKVBackend::SharedState { 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; - 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; + SequenceEntry* FindSequenceEntryLocked(std::int64_t sequence_id) const; + SequenceEntry* FindOrInsertSequenceEntryLocked(std::int64_t sequence_id, + bool* is_new, + std::int64_t* previous_marker); + + 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() { @@ -259,23 +352,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; @@ -287,13 +411,63 @@ 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() { // Always zero the entire mapping — historically we skipped the data // region under --fast-init because memfd pages are kernel-zeroed on @@ -305,6 +479,7 @@ void HostPagedKVBackend::SharedState::ConstructSharedState() { // allocation is the belt to the gather-level suspenders. std::memset(mapping, 0, total_bytes); MapPointers(); + header->magic = kSharedMemoryMagic; header->layout_fingerprint = layout_fingerprint; header->config_hash = HashHostKVConfig(config); @@ -314,20 +489,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(), @@ -347,18 +568,11 @@ void HostPagedKVBackend::SharedState::ConstructSharedState() { "pthread_mutexattr_setrobust failed"); } - if (const int rc = pthread_mutex_init(&header->allocation_mutex, &attr); + if (const int rc = pthread_mutex_init(&header->metadata_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); - 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); @@ -373,10 +587,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)); } } @@ -421,6 +631,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( @@ -451,7 +675,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) { @@ -460,15 +684,19 @@ 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 (target->sequence_id != sequence_id) { - *target = SequenceEntry(); - target->sequence_id = sequence_id; + if (previous_marker != nullptr) { + *previous_marker = target->sequence_id; } + *target = SequenceEntry(); + target->sequence_id = sequence_id; if (is_new != nullptr) { *is_new = true; } @@ -484,6 +712,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); @@ -647,9 +991,8 @@ void HostPagedKVBackend::SharedState::Initialize(bool create_region) { } } - 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); @@ -661,6 +1004,7 @@ void HostPagedKVBackend::SharedState::Initialize(bool create_region) { mapping = static_cast(mapped); MapPointers(); + BindPrefixCache(); if (created_region) { header->init_state.store( @@ -678,130 +1022,311 @@ 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) { - 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) + ")"); - } - 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]; - } - header->free_stack_top.store(new_top, std::memory_order_relaxed); + + 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)) + + ")"); } - { - ScopedMutexLock lock(&header->sequence_mutex); - bool is_new = false; - SequenceEntry* entry = - FindOrInsertSequenceEntryLocked(sequence_id, &is_new); - if (is_new) { - header->active_sequences.fetch_add(1, 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, nullptr); + 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); + + 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; + }; + + ScopedMutexLock lock(&header->metadata_mutex); + std::vector rollback_records; + rollback_records.reserve(sequence_ids.size()); + + 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"); + } + + 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 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"); + } + + 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); + } + + std::vector pages; + pages.reserve(required_pages); + + 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 < 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; } - 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; + } catch (...) { + for (auto it = rollback_records.rbegin(); it != rollback_records.rend(); + ++it) { + SequenceRollbackRecord& record = *it; + SequenceEntry* entry = record.entry; + if (entry == nullptr) { + continue; + } + + 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); } - ++entry->num_pages; } + throw; } - 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_cache_entry_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; } @@ -844,76 +1369,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); } } @@ -922,6 +1411,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 b8bff88c1..6e6ba356e 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_cache_entry_pages = 0; + std::size_t num_shared_pages = 0; }; struct HostPagedKVConfig { @@ -34,12 +41,26 @@ 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; bool enable_memfd = false; int memfd_creator_pid = -1; int memfd_fd = -1; std::string logger_name; // Custom logger name (empty = use default) }; +struct PrefixAllocationBatchResult { + std::vector> allocated_pages; + std::vector reused_prefix_tokens; +}; + inline std::uint64_t HashCombine(std::uint64_t seed, std::uint64_t value) { seed ^= value + 0x9e3779b97f4a7c15ULL + (seed << 6) + (seed >> 2); return seed; @@ -95,6 +116,37 @@ 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(128, config.num_pages / 2); + } + 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; } @@ -124,7 +176,20 @@ 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 + << ", enable_memfd=" << config.enable_memfd + << ", memfd_creator_pid=" << config.memfd_creator_pid + << ", memfd_fd=" << config.memfd_fd << ")"; return oss.str(); } @@ -135,7 +200,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 + << ", cache_entry_pages=" << stats.num_cache_entry_pages + << ", shared_pages=" << stats.num_shared_pages << ")"; return oss.str(); } @@ -153,6 +224,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); seed = HashCombine(seed, static_cast(sanitized.enable_memfd)); return seed; } @@ -176,6 +256,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); @@ -183,6 +269,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..aace747f5 --- /dev/null +++ b/core/KV_Storage/host_paged_kv_prefix_cache.cpp @@ -0,0 +1,641 @@ +#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::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) { + 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::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 + + 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..7306d8caa --- /dev/null +++ b/core/KV_Storage/host_paged_kv_prefix_cache.h @@ -0,0 +1,177 @@ +#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 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); + 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 5b4bab2c5..937a708fe 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 @@ -265,6 +266,70 @@ 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); + + 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]; + 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), + 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), + std::move(alloc_result.reused_prefix_tokens)}; + } + std::vector GrowSequencePages( std::int64_t sequence_id, std::size_t num_pages) { if (num_pages == 0) { @@ -309,6 +374,7 @@ class HostPagedKVWorkerView { ResetCopyStreams(); UnregisterPinnedMemory(); page_table_.Clear(); + ClearPrefixStates(); } KVAsyncTask AsyncLoadLayerKVToDevice( @@ -838,16 +904,17 @@ class HostPagedKVWorkerView { void UnregisterSequence(std::int64_t sequence_id) { page_table_.Remove(sequence_id); + RemovePrefixState(sequence_id); } void UnregisterSequences(const std::vector& sequence_ids) { 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) { @@ -899,11 +966,28 @@ class HostPagedKVWorkerView { const auto producer_cuda_stream = at::cuda::getCurrentCUDAStream(device_index_).stream(); + 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, - sequence_ids = std::move(sequence_ids), - sequence_lengths = std::move(sequence_lengths), - prepared_k, prepared_v, tokens_per_sequence, - producer_cuda_stream]() { + sequence_ids = std::move(sequence_ids), + sequence_lengths = std::move(sequence_lengths), + prepared_k, prepared_v, tokens_per_sequence, + prefix_skip_tokens, producer_cuda_stream]() { c10::cuda::OptionalCUDAGuard device_guard(device_index_); const auto cuda_stream = CopyStream(CopyDirection::kDeviceToHost); this->WaitForProducerStream(cuda_stream, producer_cuda_stream); @@ -939,13 +1023,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) { @@ -953,7 +1043,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); }); @@ -962,7 +1054,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, @@ -973,7 +1065,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); @@ -983,6 +1076,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); }); @@ -991,7 +1085,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(); @@ -1002,6 +1098,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()) { @@ -1024,10 +1126,11 @@ class HostPagedKVWorkerView { at::cuda::getCurrentCUDAStream(device_index_).stream(); return LaunchAsyncTask([this, layer_idx, - sequence_ids = std::move(sequence_ids), - sequence_lengths = std::move(sequence_lengths), - prepared_k, prepared_v, - producer_cuda_stream]() mutable { + sequence_ids = std::move(sequence_ids), + sequence_lengths = std::move(sequence_lengths), + prepared_k, prepared_v, + decode_token_ids = std::move(decode_token_ids), + producer_cuda_stream]() mutable { c10::cuda::OptionalCUDAGuard device_guard(device_index_); const auto cuda_stream = CopyStream(CopyDirection::kDeviceToHost); this->WaitForProducerStream(cuda_stream, producer_cuda_stream); @@ -1088,6 +1191,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); }); } @@ -1117,18 +1225,29 @@ class HostPagedKVWorkerView { KVAsyncTask AsyncAppendDecodeKVToHostBatchedKernel( std::vector entries, std::vector sequence_ids, - SequenceLengths sequence_lengths) { + SequenceLengths sequence_lengths, + std::optional> decode_token_ids = + std::nullopt) { if (entries.empty() || sequence_ids.empty()) { return LaunchAsyncTask([] {}); } EnsureDeviceReady(); const std::size_t batch = sequence_ids.size(); + if (decode_token_ids.has_value() && + decode_token_ids->size() != batch) { + throw std::invalid_argument( + "AsyncAppendDecodeKVToHostBatchedKernel: decode_token_ids " + "size must match sequence_ids size"); + } // Validate each entry and detect V presence bool has_v = false; + bool has_last_layer = false; for (const auto& e : entries) { geometry_.EnsureLayerBounds(e.layer_idx, "AsyncAppendDecodeKVToHostBatchedKernel"); + has_last_layer = + has_last_layer || (e.layer_idx + 1 == config_.num_layers); const std::size_t tokens_per_seq = ValidateKTensorShape(e.k_tensor, batch); if (tokens_per_seq != 1) { @@ -1227,7 +1346,10 @@ class HostPagedKVWorkerView { v_src_host = std::move(v_src_host), v_dst_host = std::move(v_dst_host), k_token_bytes, v_token_bytes, k_total, - v_total, producer_cuda_stream]() mutable { + v_total, producer_cuda_stream, + sequence_ids = std::move(sequence_ids), + decode_token_ids = std::move(decode_token_ids), + has_last_layer]() mutable { c10::cuda::OptionalCUDAGuard device_guard(device_index_); const auto cuda_stream = CopyStream(CopyDirection::kDeviceToHost); this->WaitForProducerStream(cuda_stream, producer_cuda_stream); @@ -1279,6 +1401,12 @@ class HostPagedKVWorkerView { } this->SynchronizeWithEvent(cuda_stream); + if (has_last_layer) { + if (decode_token_ids.has_value()) { + this->AppendDecodeTokens(sequence_ids, *decode_token_ids); + } + this->MaybeCommitPrefix(config_.num_layers - 1, sequence_ids); + } }); } @@ -1566,6 +1694,128 @@ class HostPagedKVWorkerView { } } + struct PrefixSequenceState { + std::vector tokens; + std::size_t reused_prefix_tokens = 0; + std::size_t committed_token_count = 0; + }; + + void UpdatePrefixState(std::int64_t sequence_id, + std::vector prompt_tokens, + 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.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, + 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 { + 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 CacheCommitItem { + std::int64_t sequence_id = 0; + std::vector tokens; + std::size_t token_count = 0; + }; + + 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) { + auto it = prefix_states_.find(sequence_id); + if (it == prefix_states_.end()) { + continue; + } + 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.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, + item.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()) { @@ -2302,6 +2552,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_; // Scratch device buffers for AsyncAppendDecodeKVToHostBatchedKernel — // pointer arrays (src + dst) uploaded once per batched call. Sized diff --git a/core/batchgen_Binding.cpp b/core/batchgen_Binding.cpp index fd536985a..83aa7c884 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) @@ -151,12 +383,14 @@ 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_append_decode_kv_to_host_batched_kernel", [](WorkerView& self, py::list entries_py, std::vector sequence_ids, - batchgen::kv::SequenceLengths sequence_lengths) { + batchgen::kv::SequenceLengths sequence_lengths, + std::optional> decode_token_ids) { std::vector entries; entries.reserve(entries_py.size()); for (auto item : entries_py) { @@ -176,10 +410,11 @@ void BindHostPagedWorkerView(py::module& m, const char* name) { } return self.AsyncAppendDecodeKVToHostBatchedKernel( std::move(entries), std::move(sequence_ids), - std::move(sequence_lengths)); + std::move(sequence_lengths), std::move(decode_token_ids)); }, py::arg("entries"), py::arg("sequence_ids"), py::arg("sequence_lengths"), + py::arg("decode_token_ids") = py::none(), "Batched variant: all (layer × seq) host-KV writes issued by " "one UVA kernel launch on the DtoH stream. Replaces the " "per-layer async_append_decode_kv_to_host loop of 78×bsz " @@ -229,6 +464,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) { @@ -278,7 +536,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 @@ -326,24 +584,21 @@ 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"), @@ -353,8 +608,8 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { py::arg("memfd_creator_pid") = -1, py::arg("memfd_fd") = -1) .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) @@ -374,6 +629,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_readwrite("enable_memfd", &kv::HostPagedKVConfig::enable_memfd) .def_readwrite("memfd_creator_pid", &kv::HostPagedKVConfig::memfd_creator_pid) @@ -394,11 +667,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_cache_entry_pages", + &kv::HostPagedKVStats::num_cache_entry_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 56f44bf72..67734dea0 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; } @@ -376,10 +420,16 @@ ModelConfig parse_model_config(const py::object& model_config) { std::shared_ptr init_logger(const std::string& log_level, const std::string& logger_name) { auto logger = spdlog::get(logger_name); - if (logger) { - return logger; + if (!logger) { + try { + logger = spdlog::stdout_color_mt(logger_name); + } catch (const spdlog::spdlog_ex&) { + logger = spdlog::get(logger_name); + if (!logger) { + throw; + } + } } - logger = spdlog::stdout_color_mt(logger_name); // Set colors for all five standard levels auto console_sink = dynamic_cast( diff --git a/op_builder/core_engine.py b/op_builder/core_engine.py index 1a03f1c2f..f69a6e31f 100644 --- a/op_builder/core_engine.py +++ b/op_builder/core_engine.py @@ -32,10 +32,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", ] @@ -44,9 +48,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 @@ -58,8 +62,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 @@ -106,4 +110,4 @@ def extra_ldflags(self): return flags 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/tests/integration/paged_kv/test_host_prefix_reuse.py b/tests/integration/paged_kv/test_host_prefix_reuse.py new file mode 100644 index 000000000..e96745b11 --- /dev/null +++ b/tests/integration/paged_kv/test_host_prefix_reuse.py @@ -0,0 +1,743 @@ +import ctypes +import errno +import multiprocessing as mp +import queue +import random +import string +import traceback +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, 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 = page_size_tokens + 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 _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, + cache_entry_pages=stats.num_cache_entry_pages, + 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, + cache_entry_pages=stats.num_cache_entry_pages, + 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") + + 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) + + +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_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_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"]["cache_entry_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"]["cache_entry_pages"] >= 2 + 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") + + 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/tests/integration/paged_kv/test_prefix_cache_binding.py b/tests/integration/paged_kv/test_prefix_cache_binding.py new file mode 100644 index 000000000..5677790c3 --- /dev/null +++ b/tests/integration/paged_kv/test_prefix_cache_binding.py @@ -0,0 +1,176 @@ +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 + + +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 + + +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