diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/config.h b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/config.h index 6dbedcd0445f..b88ab1645d2b 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/config.h +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/config.h @@ -148,6 +148,9 @@ struct AttentionLayerConfig // nullopt or 0 = no sink tokens. std::optional numSinkTokens; + // Layers in different ownership domains must not share lifecycle or storage pools. + std::string cacheDomain = "target"; + [[nodiscard]] std::optional windowSize() const noexcept { return slidingWindowSize; diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/lifeCycleRegistry.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/lifeCycleRegistry.cpp index c26e77f61b2f..1fcaeeceb9a5 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/lifeCycleRegistry.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/lifeCycleRegistry.cpp @@ -37,7 +37,7 @@ LifeCycle makeLifeCycle(LayerConfig const& layer, int tokensPerBlock) if constexpr (std::is_same_v) return SsmLifeCycle{}; else - return AttnLifeCycle::make(cfg.slidingWindowSize, cfg.numSinkTokens, tokensPerBlock); + return AttnLifeCycle::make(cfg.slidingWindowSize, cfg.numSinkTokens, tokensPerBlock, cfg.cacheDomain); }, layer); } diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/lifeCycleRegistry.h b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/lifeCycleRegistry.h index f655fa7ce092..317c137244be 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/lifeCycleRegistry.h +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/lifeCycleRegistry.h @@ -25,6 +25,8 @@ #include #include #include +#include +#include #include #include #include @@ -40,6 +42,7 @@ struct AttnLifeCycle { std::optional windowSize; // nullopt = no sliding window int numSinkBlocks = 0; // divUp(numSinkTokens, tokensPerBlock) + std::string cacheDomain = "target"; HalfOpenRange getStaleRange(int historyLength, int tokensPerBlock) const { @@ -56,24 +59,24 @@ struct AttnLifeCycle bool operator==(AttnLifeCycle const& o) const noexcept { - return windowSize == o.windowSize && numSinkBlocks == o.numSinkBlocks; + return windowSize == o.windowSize && numSinkBlocks == o.numSinkBlocks && cacheDomain == o.cacheDomain; } bool operator<(AttnLifeCycle const& o) const noexcept { - if (windowSize != o.windowSize) - return windowSize < o.windowSize; - return numSinkBlocks < o.numSinkBlocks; + return std::tie(windowSize, numSinkBlocks, cacheDomain) + < std::tie(o.windowSize, o.numSinkBlocks, o.cacheDomain); } - static AttnLifeCycle make(std::optional ws, std::optional numSinkTokens, int tokensPerBlock) + static AttnLifeCycle make( + std::optional ws, std::optional numSinkTokens, int tokensPerBlock, std::string cacheDomain = "target") { TLLM_CHECK_DEBUG(tokensPerBlock > 0); TLLM_CHECK_DEBUG(!ws.has_value() || *ws > 0); TLLM_CHECK_DEBUG(!numSinkTokens.has_value() || *numSinkTokens >= 0); TLLM_CHECK_DEBUG((!numSinkTokens.has_value() || *numSinkTokens == 0) || ws.has_value()); int sinkBlocks = divUp(numSinkTokens.value_or(0), tokensPerBlock); - return AttnLifeCycle{ws, sinkBlocks}; + return AttnLifeCycle{ws, sinkBlocks, std::move(cacheDomain)}; } }; diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storage/config.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storage/config.cpp index 7411577d03b1..d85f96db1cae 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storage/config.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storage/config.cpp @@ -204,20 +204,21 @@ StorageConfig createStorageConfig(KVCacheManagerConfig const& config) slotGroups.push_back(std::move(var)); } - // Merge SlotDescVariants that share the same slotSizeList. - // Key: tuple of sizes (sorted desc). - std::map, std::vector> poolGroupsBySizes; + // Merge equal storage sizes only within the same cache domain. + std::map>, std::vector> poolGroupsByLayout; for (auto& sg : slotGroups) { + auto const* attn = std::get_if(®istry[sg.lifeCycleId]); + std::string const cacheDomain = attn ? attn->cacheDomain : "target"; auto sizes = sg.slotSizeList(); - poolGroupsBySizes[sizes.raw()].push_back(std::move(sg)); + poolGroupsByLayout[{cacheDomain, sizes.raw()}].push_back(std::move(sg)); } StorageConfig out; out.cacheTiers = TypedVec{config.cacheTiers}; out.expansion = expansionMap; - for (auto& [sizes, variants] : poolGroupsBySizes) + for (auto& [layout, variants] : poolGroupsByLayout) { SlotDesc sd; sd.variants = std::move(variants); diff --git a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storageManager.cpp b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storageManager.cpp index 592684575178..13b56c03dc10 100644 --- a/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storageManager.cpp +++ b/cpp/tensorrt_llm/batch_manager/kv_cache_manager_v2/storageManager.cpp @@ -396,12 +396,14 @@ StorageManager::StorageManager(LifeCycleRegistry const& lifeCycles, StorageConfi TypedVec coldGrouping(numLifeCycles()); TypedVec coldSlotDescList; - std::map coldGroupByPageBytes; + std::map, PoolGroupIndex> coldGroupByLayout; for (LifeCycleId lifeCycle{0}; lifeCycle < numLifeCycles(); ++lifeCycle) { + auto const* attn = std::get_if(&lifeCycles[lifeCycle]); + std::string const cacheDomain = attn ? attn->cacheDomain : "target"; size_t const coldPageBytes = coldPageBytesByLifeCycle[lifeCycle]; - auto [it, inserted] = coldGroupByPageBytes.emplace( - coldPageBytes, PoolGroupIndex{static_cast(coldSlotDescList.size().value())}); + auto [it, inserted] = coldGroupByLayout.emplace( + std::pair{cacheDomain, coldPageBytes}, PoolGroupIndex{static_cast(coldSlotDescList.size().value())}); PoolGroupIndex const coldPgIdx = it->second; if (inserted) { diff --git a/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp b/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp index f45743b79290..081f391cbbc5 100644 --- a/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp +++ b/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp @@ -1341,13 +1341,15 @@ void KvCacheManagerV2Bindings::initBindings(nb::module_& m) .def("__bool__", [](kv::ScratchDesc const& self) { return static_cast(self); }); nb::class_(m, "AttnLifeCycle") - .def(nb::init, int>(), nb::arg("window_size").none(), nb::arg("num_sink_blocks")) + .def(nb::init, int, std::string>(), nb::arg("window_size").none(), + nb::arg("num_sink_blocks"), nb::arg("cache_domain") = "target") // Sink tokens round up to whole blocks. Bound rather than repeated in Python so the // connector's view of a life cycle is built by the same code as the allocator's. .def_static("make", &kv::AttnLifeCycle::make, nb::arg("window_size").none(), nb::arg("num_sink_tokens").none(), - nb::arg("tokens_per_block")) + nb::arg("tokens_per_block"), nb::arg("cache_domain") = "target") .def_prop_ro("window_size", [](kv::AttnLifeCycle const& self) { return self.windowSize; }) .def_ro("num_sink_blocks", &kv::AttnLifeCycle::numSinkBlocks) + .def_ro("cache_domain", &kv::AttnLifeCycle::cacheDomain) .def("get_stale_range", &kv::AttnLifeCycle::getStaleRange, nb::arg("history_length"), nb::arg("tokens_per_block")) .def("__eq__", &kv::AttnLifeCycle::operator==); @@ -1545,13 +1547,15 @@ void KvCacheManagerV2Bindings::initBindings(nb::module_& m) .def_rw("tokens_per_block_override", &kv::BufferConfig::tokensPerBlockOverride) DEF_COPY(kv::BufferConfig); nb::class_(m, "AttentionLayerConfig") - .def(nb::init, std::optional, std::optional>(), + .def( + nb::init, std::optional, std::optional, std::string>(), nb::arg("layer_id"), nb::arg("buffers"), nb::arg("sliding_window_size") = std::nullopt, - nb::arg("num_sink_tokens") = std::nullopt) + nb::arg("num_sink_tokens") = std::nullopt, nb::arg("cache_domain") = "target") .def_rw("layer_id", &kv::AttentionLayerConfig::layerId) .def_rw("buffers", &kv::AttentionLayerConfig::buffers) .def_rw("sliding_window_size", &kv::AttentionLayerConfig::slidingWindowSize) .def_rw("num_sink_tokens", &kv::AttentionLayerConfig::numSinkTokens) + .def_rw("cache_domain", &kv::AttentionLayerConfig::cacheDomain) .def_prop_ro("window_size", &kv::AttentionLayerConfig::windowSize) DEF_COPY(kv::AttentionLayerConfig); nb::enum_(m, "LayerType") diff --git a/tensorrt_llm/_torch/attention/backends/sparse/deepseek_v4/cache_manager.py b/tensorrt_llm/_torch/attention/backends/sparse/deepseek_v4/cache_manager.py index f7268e60b007..35b73cbbf85c 100644 --- a/tensorrt_llm/_torch/attention/backends/sparse/deepseek_v4/cache_manager.py +++ b/tensorrt_llm/_torch/attention/backends/sparse/deepseek_v4/cache_manager.py @@ -26,8 +26,11 @@ _RESERVED_REQUEST_IDS, GPU_LEVEL, KVCacheManagerV2, + _estimate_cache_size_components, _fill_kv_pages, + _get_generation_kv_capacity, ) +from tensorrt_llm._torch.pyexecutor.kv_cache.standalone_draft_cache import StandaloneDraftLayout from tensorrt_llm._utils import ( TensorWrapper, convert_to_torch_tensor, @@ -68,6 +71,24 @@ NVFP4_VECTOR_SIZE = 16 +def _draft_cache_size_components( + layout: StandaloneDraftLayout, + tokens_per_block: int, + generation_capacity_headroom: int, + target_context_bytes: int, +) -> tuple[int, int, int]: + """The draft contribution to static profiling and runtime byte quotas.""" + context, generation, per_request = _estimate_cache_size_components( + [layout.bytes_per_layer_token] * layout.num_layers, + [layout.retention_window_size] * layout.num_layers, + tokens_per_block, + scratch=False, + generation_capacity_headroom=generation_capacity_headroom, + standalone_draft_reserve=layout.extra_tokens, + ) + return context, generation, per_request + layout.extra_tokens * target_context_bytes + + def get_attn_dim( head_dim: int, index_head_dim: int, compress_ratio: int, attn_type: DeepseekV4AttentionType ) -> int: @@ -923,22 +944,14 @@ def _get_extra_quota_padding(self) -> int: """Ensure each attention type has minimal space when max_tokens is small.""" return len(DeepseekV4AttentionType) * (2 << 20) - def _get_quota_from_max_tokens(self, max_tokens: int) -> int: + def _get_cache_cost_components(self) -> tuple[int, int, int, int]: + """Return context/generation bytes per token, then fixed bytes per request.""" compress_ratios = [self._compress_ratios[layer] for layer in self.pp_layers] has_fp8_kv_cache = self.dtype == DataType.FP8 - non_sliding_attn_size_per_token = _estimate_non_sliding_attn_size_per_token( - self.head_dim, - self.index_head_dim, - compress_ratios, - has_fp8_kv_cache, - indexer_k_dtype=self._indexer_k_dtype, - use_fp8_ds_mla=self.use_fp8_ds_mla, - has_nvfp4_compress=self._use_nvfp4_compress, - nvfp4_residual_dim=self.nvfp4_residual_dim, - ) + non_sliding_attn_size_per_token = self.get_cache_bytes_per_token() ( context_swa_size_per_token, - _, + context_swa_size_per_request, ) = _estimate_swa_cache_size( self.head_dim, self.index_head_dim, @@ -966,65 +979,56 @@ def _get_quota_from_max_tokens(self, max_tokens: int) -> int: indexer_k_dtype=self._indexer_k_dtype, use_fp8_ds_mla=self.use_fp8_ds_mla, ) + context_size_per_token = non_sliding_attn_size_per_token + context_swa_size_per_token + generation_size_per_token = non_sliding_attn_size_per_token + generation_swa_size_per_token + if self.draft_layout is not None: + draft_context, draft_generation, draft_per_request = _draft_cache_size_components( + self.draft_layout, + self.tokens_per_block, + self._generation_kv_capacity_headroom, + context_size_per_token, + ) + context_size_per_token += draft_context + generation_size_per_token += draft_generation + context_swa_size_per_request += draft_per_request + generation_swa_size_per_request += draft_per_request + return ( + context_size_per_token, + generation_size_per_token, + context_swa_size_per_request, + generation_swa_size_per_request, + ) + + def _get_quota_from_max_tokens(self, max_tokens: int) -> int: + ( + context_size_per_token, + generation_size_per_token, + _, + generation_size_per_request, + ) = self._get_cache_cost_components() max_context_tokens = ( self._max_num_tokens if self._max_num_tokens is not None else max_tokens ) context_tokens = min(max_tokens, max_context_tokens) generation_tokens = max_tokens - context_tokens - generation_quota = ( - max_tokens * non_sliding_attn_size_per_token - + generation_tokens * generation_swa_size_per_token - + self.max_batch_size * generation_swa_size_per_request + return int( + context_tokens * context_size_per_token + + generation_tokens * generation_size_per_token + + self.max_batch_size * generation_size_per_request + + self._get_extra_quota_padding() ) - context_extra_quota = context_tokens * context_swa_size_per_token - padding = self._get_extra_quota_padding() - return int(generation_quota + context_extra_quota + padding) def _get_max_tokens_from_quota(self, quota: int) -> float: - compress_ratios = [self._compress_ratios[layer] for layer in self.pp_layers] - has_fp8_kv_cache = self.dtype == DataType.FP8 - non_sliding_attn_size_per_token = _estimate_non_sliding_attn_size_per_token( - self.head_dim, - self.index_head_dim, - compress_ratios, - has_fp8_kv_cache, - indexer_k_dtype=self._indexer_k_dtype, - use_fp8_ds_mla=self.use_fp8_ds_mla, - has_nvfp4_compress=self._use_nvfp4_compress, - nvfp4_residual_dim=self.nvfp4_residual_dim, - ) - context_swa_size_per_token, _ = _estimate_swa_cache_size( - self.head_dim, - self.index_head_dim, - compress_ratios, - has_fp8_kv_cache, - self.tokens_per_block, - self._swa_window_size, - context=True, - scratch=self.enable_swa_scratch_reuse, - indexer_k_dtype=self._indexer_k_dtype, - use_fp8_ds_mla=self.use_fp8_ds_mla, - ) ( - generation_swa_size_per_token, - generation_swa_size_per_request, - ) = _estimate_swa_cache_size( - self.head_dim, - self.index_head_dim, - compress_ratios, - has_fp8_kv_cache, - self.tokens_per_block, - self._swa_window_size, - context=False, - scratch=False, - indexer_k_dtype=self._indexer_k_dtype, - use_fp8_ds_mla=self.use_fp8_ds_mla, - ) + context_size_per_token, + generation_size_per_token, + _, + generation_size_per_request, + ) = self._get_cache_cost_components() padding = self._get_extra_quota_padding() - size_per_batch = self.max_batch_size * generation_swa_size_per_request + padding + size_per_batch = self.max_batch_size * generation_size_per_request + padding if quota < size_per_batch: return 0 - context_size_per_token = non_sliding_attn_size_per_token + context_swa_size_per_token if self._max_num_tokens is None: return (quota - size_per_batch) / context_size_per_token @@ -1032,7 +1036,6 @@ def _get_max_tokens_from_quota(self, quota: int) -> float: if quota <= context_limit_quota: return (quota - size_per_batch) / context_size_per_token - generation_size_per_token = non_sliding_attn_size_per_token + generation_swa_size_per_token if generation_size_per_token <= 0: return float("inf") return self._max_num_tokens + (quota - context_limit_quota) / generation_size_per_token @@ -1187,6 +1190,12 @@ def _add_layer( layers=layers, ) + def _append_standalone_draft_layers( + self, config: KVCacheManagerConfigPy + ) -> KVCacheManagerConfigPy: + # Preserve target virtual-layer indices when registering draft layers. + return super()._append_standalone_draft_layers(config, register_model_layers=False) + def _init_indexer_dtype(self, sparse_attn_config: DeepSeekV4SparseAttentionConfig) -> None: # Indexer compressor cache layout. Two modes are supported: # - "fp8" (FP8 blockwise): 1 byte per value + 1 fp32 scale per 128 @@ -1371,34 +1380,15 @@ def _get_generation_bytes(self, request: llm_request.LlmRequest) -> int: return self._get_cache_bytes_for_tokens(total_tokens, context=False) def _get_cache_bytes_for_tokens(self, total_tokens: int, *, context: bool) -> int: - has_fp8_kv_cache = self.dtype == DataType.FP8 - compress_ratios = [self._compress_ratios[layer] for layer in self.pp_layers] - non_sliding_attn_size_per_token = _estimate_non_sliding_attn_size_per_token( - self.head_dim, - self.index_head_dim, - compress_ratios, - has_fp8_kv_cache, - indexer_k_dtype=self._indexer_k_dtype, - use_fp8_ds_mla=self.use_fp8_ds_mla, - has_nvfp4_compress=self._use_nvfp4_compress, - nvfp4_residual_dim=self.nvfp4_residual_dim, - ) - swa_size_per_token, swa_size_per_request = _estimate_swa_cache_size( - self.head_dim, - self.index_head_dim, - compress_ratios, - has_fp8_kv_cache, - self.tokens_per_block, - self._swa_window_size, - context=context, - scratch=self.enable_swa_scratch_reuse, - indexer_k_dtype=self._indexer_k_dtype, - use_fp8_ds_mla=self.use_fp8_ds_mla, - ) - return int( - total_tokens * (non_sliding_attn_size_per_token + swa_size_per_token) - + swa_size_per_request - ) + ( + context_size_per_token, + generation_size_per_token, + context_size_per_request, + generation_size_per_request, + ) = self._get_cache_cost_components() + if context: + return int(total_tokens * context_size_per_token + context_size_per_request) + return int(total_tokens * generation_size_per_token + generation_size_per_request) def get_needed_resource_to_completion(self, request: llm_request.LlmRequest) -> int: if self._is_generation_request(request): @@ -1412,6 +1402,8 @@ def get_layer_bytes_per_token( local_layer_idx: int, data_role: DataRole, ) -> int: + if self._is_standalone_draft_layer(local_layer_idx): + return super().get_layer_bytes_per_token(local_layer_idx, data_role) # The generic layers in the base config are replaced by # _build_cache_config, so their buffer sizes are only placeholders. return 1 @@ -1686,6 +1678,29 @@ def get_cache_size_per_token(model_config: ModelConfig, mapping: Mapping, **kwar use_fp8_ds_mla=use_fp8_ds_mla, ) max_batch_size = int(kwargs.get("max_batch_size") or 0) + draft_layout = kwargs.get("draft_layout") + if draft_layout is not None: + _, headroom = _get_generation_kv_capacity(kwargs.get("spec_config"), is_draft=False) + target_context_swa_bytes, _ = _estimate_swa_cache_size( + head_dim, + index_head_dim, + compress_ratios, + has_fp8_kv_cache, + kwargs["tokens_per_block"], + model_config.sparse_attention_config.window_size, + context=True, + scratch=False, + indexer_k_dtype=indexer_k_dtype, + use_fp8_ds_mla=use_fp8_ds_mla, + ) + _, draft_generation, draft_per_request = _draft_cache_size_components( + draft_layout, + kwargs["tokens_per_block"], + headroom + draft_layout.extra_tokens, + non_sliding_attn_size_per_token + target_context_swa_bytes, + ) + swa_size_per_token += draft_generation + swa_size_per_request += draft_per_request return ( non_sliding_attn_size_per_token + swa_size_per_token, swa_size_per_request * max_batch_size, diff --git a/tensorrt_llm/_torch/disaggregation/native/auxiliary.py b/tensorrt_llm/_torch/disaggregation/native/auxiliary.py index 3620b9adb8bf..6b847d2a0336 100644 --- a/tensorrt_llm/_torch/disaggregation/native/auxiliary.py +++ b/tensorrt_llm/_torch/disaggregation/native/auxiliary.py @@ -1,3 +1,6 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + from abc import ABC, abstractmethod from collections import deque, namedtuple from dataclasses import dataclass, field @@ -111,6 +114,63 @@ def build_aux_transfer_layout( AuxSlot = namedtuple("AuxSlot", ["id", "buffer"]) +_DRAFT_HISTORY_VERSION = 2 +_DRAFT_HISTORY_FIELDS = 10 +_DRAFT_DTYPE_CODES = {"torch.float16": 1, "torch.bfloat16": 2} +_DRAFT_BACKEND_CODES = {"VANILLA": 1, "TRTLLM": 2, "DSv4": 3} + + +def _encode_draft_history(history: dict[str, Any]) -> list[int]: + """Serialize the manager's validated history and rank-local storage identity.""" + if history is None: + raise ValueError("Standalone draft transfer requires draft history metadata") + layout = history["layout"] + return [ + _DRAFT_HISTORY_VERSION, + history["valid_length"], + history["position"], + layout["num_layers"], + layout["num_kv_heads"], + layout["head_dim"], + _DRAFT_DTYPE_CODES[layout["dtype"]], + _DRAFT_BACKEND_CODES[layout["attention_backend"]], + layout["kv_factor"], + layout["window_size"] or 0, + ] + + +def _decode_draft_history(values: list[int]) -> dict[str, Any]: + if values[0] != _DRAFT_HISTORY_VERSION: + raise ValueError("Missing or unsupported standalone draft history metadata version") + ( + _, + valid_length, + position, + num_layers, + num_kv_heads, + head_dim, + dtype_code, + backend_code, + kv_factor, + window_size, + ) = values + dtypes = {code: name for name, code in _DRAFT_DTYPE_CODES.items()} + backends = {code: name for name, code in _DRAFT_BACKEND_CODES.items()} + layout = { + "num_layers": num_layers, + "num_kv_heads": num_kv_heads, + "head_dim": head_dim, + "dtype": dtypes.get(dtype_code), + "attention_backend": backends.get(backend_code), + "kv_factor": kv_factor, + "window_size": window_size or None, + } + return { + "valid_length": valid_length, + "position": position, + "layout": layout, + } + class AuxBufferBase(ABC): """ @@ -166,8 +226,15 @@ def get_slot_data(self, slot: int) -> tuple[list[int], list[int], tuple[int, int class AuxBuffer(AuxBufferBase): - def __init__(self, max_slot_num: int, beam_width: int, max_draft_len: int, device: str = "cpu"): - # public constructor args remain the same, internals are private + def __init__( + self, + max_slot_num: int, + beam_width: int, + max_draft_len: int, + device: str = "cpu", + *, + draft_history: bool = False, + ) -> None: self._max_slot_num = int(max_slot_num) self._beam_width = int(beam_width) self._max_draft_len = int(max_draft_len) @@ -196,35 +263,32 @@ def __init__(self, max_slot_num: int, beam_width: int, max_draft_len: int, devic self._prompt_token_counts_buffer = torch.zeros( self._max_slot_num, 2, dtype=data_type, device=self._device ) + # Version zero denotes an unfilled or recycled slot. + self._draft_history_buffer = ( + torch.zeros( + self._max_slot_num, _DRAFT_HISTORY_FIELDS, dtype=torch.int64, device=self._device + ) + if draft_history + else None + ) + + buffers = [ + self._first_tokens_buffer, + self._draft_tokens_buffer, + self._token_counts_buffer, + self._prompt_token_counts_buffer, + ] + if self._draft_history_buffer is not None: + buffers.append(self._draft_history_buffer) self._meta = AuxBufferMeta( - ptrs=np.array( - [ - self._first_tokens_buffer.data_ptr(), - self._draft_tokens_buffer.data_ptr(), - self._token_counts_buffer.data_ptr(), - self._prompt_token_counts_buffer.data_ptr(), - ], - dtype=np.int64, - ), + ptrs=np.array([buffer.data_ptr() for buffer in buffers], dtype=np.int64), size=np.array( - [ - self._first_tokens_buffer.numel() * self._first_tokens_buffer.element_size(), - self._draft_tokens_buffer.numel() * self._draft_tokens_buffer.element_size(), - self._token_counts_buffer.numel() * self._token_counts_buffer.element_size(), - self._prompt_token_counts_buffer.numel() - * self._prompt_token_counts_buffer.element_size(), - ], + [buffer.numel() * buffer.element_size() for buffer in buffers], dtype=np.int64, ), item_sizes=np.array( - [ - self._first_tokens_buffer[0].numel() * self._first_tokens_buffer.element_size(), - self._draft_tokens_buffer[0].numel() * self._draft_tokens_buffer.element_size(), - self._token_counts_buffer[0].numel() * self._token_counts_buffer.element_size(), - self._prompt_token_counts_buffer[0].numel() - * self._prompt_token_counts_buffer.element_size(), - ], + [buffer[0].numel() * buffer.element_size() for buffer in buffers], dtype=np.int64, ), device=self._device, @@ -245,6 +309,8 @@ def alloc_slot(self) -> AuxSlot: ) self._occupied_slots.add(slot_id) self._slot_token_counts[slot_id] = (0, 0) + if self._draft_history_buffer is not None: + self._draft_history_buffer[slot_id].zero_() return AuxSlot(slot_id, self) def free_slot(self, slot: int) -> None: @@ -265,6 +331,11 @@ def free_slot(self, slot: int) -> None: def meta(self) -> AuxBufferMeta: return self._meta + @property + def has_draft_history(self) -> bool: + """Whether this buffer transfers standalone drafter history metadata.""" + return self._draft_history_buffer is not None + def fill_slot(self, slot: int, request: LlmRequest) -> None: if slot not in self._occupied_slots: raise ValueError( @@ -301,6 +372,11 @@ def fill_slot(self, slot: int, request: LlmRequest) -> None: self._prompt_token_counts_buffer[slot].copy_( torch.tensor([prompt_tokens, cached_tokens], dtype=torch.int32, device=self._device) ) + if self._draft_history_buffer is not None: + values = _encode_draft_history(request.py_draft_transfer_history) + self._draft_history_buffer[slot].copy_( + torch.tensor(values, dtype=torch.int64, device=self._device) + ) @staticmethod def _resolve_prompt_token_counts(request: LlmRequest) -> tuple[int, int]: @@ -331,3 +407,11 @@ def get_slot_data(self, slot: int) -> tuple[list[int], list[int], tuple[int, int first_gen_tokens, draft_tokens = self.get_slot_tokens(slot) prompt_tokens, cached_tokens = self._prompt_token_counts_buffer[slot].tolist() return first_gen_tokens, draft_tokens, (int(prompt_tokens), int(cached_tokens)) + + def get_slot_draft_history(self, slot: int) -> dict[str, Any]: + """Decode transferred history for validation against the receiving request/cache.""" + if slot not in self._occupied_slots: + raise ValueError(f"Cannot read slot {slot}: slot is not currently allocated.") + if self._draft_history_buffer is None: + raise ValueError("Standalone draft history transfer is not enabled for this buffer") + return _decode_draft_history(self._draft_history_buffer[slot].tolist()) diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index aac8cd167081..161cc1cbed23 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -1640,7 +1640,7 @@ def _handle_cancel_session(self, message: list[bytes]): @nvtx_range("_respond_with_kv") def _respond_with_kv(self, _send_id: bytes, message: list[bytes]): # _sessions_lock prevents a race between session lookup and req_info save. - # session.lock atomically saves peer info and snapshots tasks against send(). + # session.lock saves peer info and snapshots tasks against send() and send_aux(). info: RecvReqInfo = RecvReqInfo.from_bytes(message[1]) with self._sessions_lock: session = self._get_session(info.unique_rid) @@ -1657,6 +1657,8 @@ def _respond_with_kv(self, _send_id: bytes, message: list[bytes]): self._save_peer_req_info(info) tasks = list(session.kv_tasks) terminal = session.has_failed() + if not terminal and session.aux_task is not None: + tasks.append(session.aux_task) include_aux = terminal and bool( session._claim_unsubmitted_aux_failures_locked((info,)) ) @@ -1907,7 +1909,9 @@ def __init__( self._timeout_s = timeout_s self._overall_timeout_s = overall_timeout_s self._deadline_monotonic_s: Optional[float] = None - self._need_aux = params.schedule_style == DisaggScheduleStyle.GENERATION_FIRST + self._need_aux = params.schedule_style == DisaggScheduleStyle.GENERATION_FIRST or ( + aux_buffer is not None and aux_buffer.has_draft_history + ) self._enforce_physical_ownership = getattr(sender, "_enforce_physical_ownership", False) self._sender: Sender # narrow base class type for Pylance self.request_id = request_id @@ -3018,7 +3022,9 @@ def __init__( ): super().__init__(receiver, SessionArgsBase(params, prompt_len=prompt_len)) self._timeout_s = timeout_s - self._need_aux = params.schedule_style == DisaggScheduleStyle.GENERATION_FIRST + self._need_aux = params.schedule_style == DisaggScheduleStyle.GENERATION_FIRST or ( + aux_buffer is not None and aux_buffer.has_draft_history + ) self._enforce_physical_ownership = getattr(receiver, "_enforce_physical_ownership", False) self._receiver: Receiver # narrow base class type for Pylance self.request_id = request_id @@ -3482,6 +3488,8 @@ def unpack_aux(self, request: LlmRequest) -> None: """Read token data from the aux buffer slot into the given request.""" assert self._aux_buffer is not None, "No aux_buffer set for this session" assert self.aux_slot is not None, "No aux_slot set for this session" + if self._aux_buffer.has_draft_history: + self.unpack_draft_history(request) first_gen_tokens, draft_tokens, (prompt_tokens, cached_tokens) = ( self._aux_buffer.get_slot_data(self.aux_slot) ) @@ -3497,6 +3505,12 @@ def unpack_aux(self, request: LlmRequest) -> None: }, } + def unpack_draft_history(self, request: LlmRequest) -> None: + """Read standalone history without changing context-first token fields.""" + if self._aux_buffer is None or self.aux_slot is None: + raise ValueError("Standalone draft transfer requires an auxiliary buffer slot") + request.py_draft_transfer_history = self._aux_buffer.get_slot_draft_history(self.aux_slot) + def is_completed(self) -> bool: """Non-blocking check: has the transfer completed successfully? @@ -3751,7 +3765,10 @@ def _create_nixl_agent( def _make_aux_buffer( kvm: KVCacheManager, max_slots: int, max_draft_len: Optional[int] = None ) -> Optional[AuxBuffer]: + draft_history = getattr(kvm, "draft_layout", None) is not None if max_slots <= 0: + if draft_history: + raise ValueError("Standalone draft transfer requires auxiliary buffer slots") return None if max_draft_len is None: max_draft_len = max(0, int(getattr(kvm, "max_draft_len", 0))) @@ -3760,6 +3777,7 @@ def _make_aux_buffer( beam_width=max(1, int(getattr(kvm, "max_beam_width", 1))), max_draft_len=max_draft_len, device="cpu", + draft_history=draft_history, ) diff --git a/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py b/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py index a5fb8f56c061..12c660c7f698 100644 --- a/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py +++ b/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py @@ -481,6 +481,8 @@ def _compute_global_layer_ids(manager, lg_idx: int) -> List[int]: inverse = {} for (model_layer, attn_type), layer_id in manager._layer_attn_to_layer_id.items(): inverse[layer_id] = (model_layer, attn_type.value) + for draft_layer in getattr(manager, "draft_layer_ids", ()): + inverse[manager.layer_offsets[draft_layer]] = (draft_layer, 0) # Use the full enum range for consistent encoding across all PP ranks. # Different PP ranks may have different subsets of attention types (e.g., @@ -740,7 +742,9 @@ def _window_size_for_layer(internal_layer_id: int): # may exceed the length of num_kv_heads_per_layer. Use index 0 as # all layers within a pool group share the same kv_heads count. first_local_layer = all_internal_layer_ids[0] - if first_local_layer < len(manager.num_kv_heads_per_layer): + if manager._is_standalone_draft_layer(first_local_layer): + num_kv_heads = manager.draft_layout.num_kv_heads + elif first_local_layer < len(manager.num_kv_heads_per_layer): num_kv_heads = manager.num_kv_heads_per_layer[first_local_layer] else: num_kv_heads = manager.num_kv_heads_per_layer[0] diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index 701166e618cb..39bdac726f96 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -427,6 +427,15 @@ def _describe_local(self, req: LlmRequest) -> Chunk: # prompt positions through beam 0's block table), so the # positional path applies to every beam width. ordinals = adapter.get_block_ordinals(req, idx, lg) + if ( + isinstance(self._kv_cache_manager, KVCacheManagerV2) + and self._kv_cache_manager.draft_layout is not None + and lg.sliding_window_size is not None + ): + stale_end = max(0, (req.prompt_len + 1 - lg.sliding_window_size) // tpb) + prompt_pages = ordinals[stale_end:prompt_blocks] + if prompt_pages.size != prompt_blocks - stale_end or np.any(prompt_pages < 0): + raise ValueError("Missing allocated prompt pages for windowed KV transfer") group = self._positional_window(ordinals, prompt_blocks, cached_per_lg[idx] // tpb) groups.append(group) @@ -486,10 +495,75 @@ def _chunk_num_bytes(self, chunk: Chunk) -> int: total += n * pool.slot_bytes return total - @staticmethod - def _need_aux_transfer(req: LlmRequest) -> bool: + def _need_aux_transfer(self, req: LlmRequest) -> bool: + params = req.py_disaggregated_params + return getattr(self._kv_cache_manager, "draft_layout", None) is not None or ( + params is not None and params.schedule_style == DisaggScheduleStyle.GENERATION_FIRST + ) + + def _validate_draft_transfer(self, req: LlmRequest) -> None: + manager = getattr(self, "_kv_cache_manager", None) + if getattr(manager, "draft_layout", None) is None: + return params = req.py_disaggregated_params - return params is not None and params.schedule_style == DisaggScheduleStyle.GENERATION_FIRST + if params is not None and params.schedule_style == DisaggScheduleStyle.GENERATION_FIRST: + raise ValueError( + "DSpark draft-state transfer requires context_first scheduling; " + "generation_first is not yet supported for draft history." + ) + if self.pipeline_transfer_enabled: + raise ValueError("DSpark draft-state transfer does not support pipelined transfer.") + + @staticmethod + def _validate_draft_history_range(req: LlmRequest, history: dict) -> None: + # Full-attention drafters retain the whole prompt; rolling drafters + # retain its complete live suffix. Neither includes speculative scratch. + window_size = history["layout"]["window_size"] + expected_length = req.prompt_len + if window_size is not None: + if type(window_size) is not int or window_size <= 0: + raise ValueError("Invalid draft history window size") + expected_length = min(expected_length, window_size) + if history["valid_length"] != expected_length or history["position"] != req.prompt_len: + raise ValueError( + "DSpark transfer requires valid draft history and sequence position " + f"covering the complete prompt ({req.prompt_len} tokens, " + f"{expected_length} retained)." + ) + + def _pack_draft_history(self, req: LlmRequest) -> None: + manager = getattr(self, "_kv_cache_manager", None) + if getattr(manager, "draft_layout", None) is None: + return + self._validate_draft_transfer(req) + history = self._kv_cache_manager.export_draft_history(req.py_request_id) + self._validate_draft_history_range(req, history) + req.py_draft_transfer_history = history + + def _restore_draft_history(self, req: LlmRequest) -> None: + history = getattr(req, "py_draft_transfer_history", None) + manager = getattr(self, "_kv_cache_manager", None) + has_draft = getattr(manager, "draft_layout", None) is not None + if history is None: + if has_draft: + raise ValueError( + "Standalone DSpark generation requires draft history from a prefill worker " + "with matching speculative configuration; draft history metadata is missing." + ) + return + if not has_draft: + raise ValueError( + "Received standalone DSpark draft history without a manager-owned draft cache." + ) + self._validate_draft_history_range(req, history) + # K/V already occupies receiver-local pages; restore only validity and position. + self._kv_cache_manager.restore_draft_history(req.py_request_id, history) + + def _prepare_received_history(self, session: RxSessionBase, req: LlmRequest) -> None: + if self._need_aux_transfer(req): + self._apply_aux(session, req) + self._assert_disagg_history_declared(req) + self._restore_draft_history(req) def _validate_bridge_req(self, req: LlmRequest, synchronous: bool = False) -> bool: if not getattr(self, "_fp4_mla_bridge_enabled", False): @@ -774,6 +848,11 @@ def _close_session_or_raise(self, session: object, rid: int, outcome: str) -> No def _apply_aux(self, session, req: LlmRequest): """Unpack aux tokens from session into request's context_phase_params.""" + params = req.py_disaggregated_params + if params is not None and params.schedule_style != DisaggScheduleStyle.GENERATION_FIRST: + # Context-first tokens and usage already arrived in the context response. + session.unpack_draft_history(req) + return session.unpack_aux(req) first_gen_tokens = req.py_first_gen_tokens # type: ignore[attr-defined] draft_tokens = req.py_draft_tokens @@ -897,6 +976,7 @@ def respond_and_send_async(self, req: LlmRequest) -> None: if not self._validate_bridge_req(req): return + self._pack_draft_history(req) self._ever_had_send_session = True # Keep the latest slice's transfer-start timestamp. req.set_kv_cache_transfer_start(tensorrt_llm.bindings.global_steady_clock_now()) @@ -945,6 +1025,7 @@ def respond_and_send_async(self, req: LlmRequest) -> None: def request_and_receive_sync(self, req: LlmRequest) -> None: if not self._validate_bridge_req(req, synchronous=True): return + self._validate_draft_transfer(req) rid = get_unique_rid(req) self._ever_had_recv_session = True if rid in self._recv_sessions: @@ -971,9 +1052,7 @@ def request_and_receive_sync(self, req: LlmRequest) -> None: req.set_kv_cache_size( self._chunk_num_bytes(extent.local) * self._kv_size_rank_factor ) - if self._need_aux_transfer(req): - self._apply_aux(session, req) - self._assert_disagg_history_declared(req) + self._prepare_received_history(session, req) req.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE else: req.state = LlmRequestState.DISAGG_TRANS_ERROR @@ -1023,6 +1102,7 @@ def request_and_receive_async(self, req: LlmRequest) -> None: """ if not self._validate_bridge_req(req): return + self._validate_draft_transfer(req) self._ever_had_recv_session = True req.set_kv_cache_transfer_start(tensorrt_llm.bindings.global_steady_clock_now()) rid = get_unique_rid(req) @@ -1163,6 +1243,9 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]) -> GenT 0 if need_progress else wait_num, block_all, ) + has_draft_history = ( + getattr(getattr(self, "_kv_cache_manager", None), "draft_layout", None) is not None + ) completed, failed, cancelled = [], [], [] for rid in to_process: @@ -1178,6 +1261,18 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]) -> GenT cancelled.append(rid) elif result == WaitResult.COMPLETED: req = self._recv_reqs[rid] + if has_draft_history: + try: + # Validate local pages/history before rank consensus; any peer + # failure follows ordinary failed-request KV/history cleanup. + self._prepare_received_history(session, req) + except (ValueError, RuntimeError) as error: + logger.warning( + f"Disagg draft history validation FAILED rank={self._dist.rank} " + f"rid={rid}: {error}" + ) + failed.append(rid) + continue if session.transfer_end_time is not None: req.set_kv_cache_transfer_end(session.transfer_end_time) if session.kv_cache_size_bytes > 0: @@ -1225,9 +1320,8 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]) -> GenT req = self._recv_reqs[rid] # transfer_end already stamped at completion detection above. req.set_kv_cache_size(getattr(req, "py_kv_cache_xfer_bytes", 0)) - if self._need_aux_transfer(req): - self._apply_aux(session, req) - self._assert_disagg_history_declared(req) + if not has_draft_history: + self._prepare_received_history(session, req) self._close_session_or_raise(session, rid, "completed") req.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE del self._recv_reqs[rid] diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 608c524af6d3..8b0eab0fb677 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -63,12 +63,14 @@ from .connectors.kv_cache_connector import KvCacheConnectorManager from .dwdp import DwdpManager from .guided_decoder import GuidedDecoder -from .kv_cache.kv_cache_manager_v2 import KVCacheManagerV2 +from .kv_cache.kv_cache_manager_v2 import (KVCacheManagerV2, + get_draft_cache_unsupported_reason) from .kv_cache.mamba_cache_manager import (BaseMambaCacheManager, CppMambaHybridCacheManager, MambaHybridCacheManagerV2, MixedMambaHybridCacheManager, use_py_mamba_cache_manager) +from .kv_cache.standalone_draft_cache import StandaloneDraftLayout from .llm_request import ExecutorResponse, LlmRequestState from .model_engine import PyTorchModelEngine from .py_executor import PyExecutor @@ -862,6 +864,7 @@ def __init__( self._disable_overlap_scheduler = llm_args.disable_overlap_scheduler self._draft_config = draft_config self._skip_est = skip_est + self._validate_standalone_draft_cache() # Admission cap (tokens of summed context attended-KV) that the fp8 context-MLA workspace reservation # covers, computed in configure_kv_cache_capacity and carried to the KV manager so the scheduler # reads it directly instead of re-deriving it from pool layout. None until reserved (or w == 0). @@ -1018,11 +1021,15 @@ def _get_kv_size_per_token(self, model_config = self._model_engine.model.model_config use_separate_draft_kv_cache = ( self._should_create_separate_draft_kv_cache()) + draft_kwargs = {} + if self._uses_unified_standalone_draft_cache(): + draft_kwargs["draft_layout"] = self._get_standalone_draft_layout() total = self._per_manager_cache_cost( self._kv_cache_manager_cls, model_config, kv_cache_config, - use_separate_draft_kv_cache=use_separate_draft_kv_cache) + use_separate_draft_kv_cache=use_separate_draft_kv_cache, + **draft_kwargs) if self._is_encoder_decoder(): total += CacheCost.from_raw(self._get_cross_kv_size_per_token()) draft_cost = self._get_draft_cache_cost( @@ -1338,6 +1345,9 @@ def _get_token_num_for_estimation(self) -> int: if spec_cfg is not None: num_extra_tokens_per_seq += spec_cfg.tokens_per_gen_step - 1 num_extra_tokens_per_seq += get_num_extra_kv_tokens(spec_cfg) + if self._uses_unified_standalone_draft_cache(): + draft_layout = self._get_standalone_draft_layout() + num_extra_tokens_per_seq += draft_layout.extra_tokens if self._dummy_reqs is None: self._dummy_reqs = self._create_dummy_context_requests( @@ -1382,6 +1392,8 @@ def _get_token_num_for_estimation(self) -> int: self._model_engine.max_seq_len, self._kv_cache_config.max_attention_window, ) + if self._uses_unified_standalone_draft_cache(): + num_pool_groups += 1 num_cache_blocks *= num_pool_groups # Dummy context requests use the configured maximum beam width. Scale @@ -1768,12 +1780,16 @@ def _create_kv_cache_manager( kv_cache_manager_cls = self._get_model_kv_cache_manager_cls( model_engine, kv_cache_config) - # When using separate draft KV cache in one-model speculative decoding, - # use layer_mask to include only target layers. The draft layers should - # only be in the separate draft KV cache manager. + # Keep the target layer layout separate from standalone draft layouts. + # Legacy modes construct a separate manager; unified standalone DSpark + # passes an explicit draft layout to the target's owner instead. # We still pass spec_config so that num_extra_kv_tokens is calculated. spec_dec_layer_mask = None - if self._should_create_separate_draft_kv_cache(): + standalone_draft_layout = (self._get_standalone_draft_layout() if + self._uses_unified_standalone_draft_cache() + else None) + if (self._should_create_separate_draft_kv_cache() + or standalone_draft_layout is not None): num_target_layers = model_engine.model.model_config.pretrained_config.num_hidden_layers spec_dec_layer_mask = [True] * num_target_layers @@ -1802,6 +1818,7 @@ def _create_kv_cache_manager( self._llm_args.kv_cache_config.kv_events_config, cold_page_codec_provider=cold_page_codec_provider, joint_kv_cache_reuse=self._joint_kv_cache_reuse, + standalone_draft_layout=standalone_draft_layout, ) if not self._skip_est: @@ -1829,6 +1846,115 @@ def _create_kv_cache_manager( return kv_cache_manager + def _is_standalone_dspark(self) -> bool: + spec_config = self._speculative_config + return (spec_config is not None + and spec_config.spec_dec_mode.is_dspark() + and not spec_config.draft_is_embedded_in_target + and not spec_config._use_shared_kv_cache) + + def _is_embedded_dspark(self) -> bool: + spec_config = self._speculative_config + return (spec_config is not None + and spec_config.spec_dec_mode.is_dspark() + and spec_config.draft_is_embedded_in_target) + + def _uses_unified_standalone_draft_cache(self) -> bool: + if (not (self._is_standalone_dspark() or self._is_embedded_dspark()) + or not self._is_kv_cache_manager_v2): + return False + # Disaggregation validates unified support; aggregate may use legacy state. + return (self._is_disagg + or self._unified_draft_cache_unsupported_reason() is None) + + def _validate_standalone_draft_cache(self) -> None: + """Reject unsupported DSpark state ownership before profiling.""" + if not (self._is_standalone_dspark() or self._is_embedded_dspark()): + return + if not self._is_kv_cache_manager_v2: + if self._kv_cache_config.use_kv_cache_manager_v2 is True: + raise ValueError( + "DSpark requested KVCacheManagerV2 but its " + "configuration resolved to V1. Remove unsupported V2 " + "features, including beam search, instead of falling " + "back to private draft state.") + if self._is_disagg: + raise ValueError("DSpark disaggregation requires " + "kv_cache_config.use_kv_cache_manager_v2=True " + "on both workers.") + return + if self._is_disagg: + reason = self._unified_draft_cache_unsupported_reason() + if reason is not None: + raise ValueError(reason) + + def _unified_draft_cache_unsupported_reason(self) -> Optional[str]: + """Shared admission requirements for unified aggregate and disagg KV.""" + reason = get_draft_cache_unsupported_reason(self._kv_cache_config) + if reason is not None: + return reason + if (self._is_standalone_dspark() + and is_mla(self._draft_config.pretrained_config)): + return "Unified DSpark KV cache does not yet support MLA drafters." + if (self._speculative_config.draft_len_schedule is not None + or self._speculative_config.max_concurrency is not None): + return ( + "Unified DSpark KV cache does not yet support " + "draft_len_schedule or max_concurrency: skipped drafting " + "would lose accepted-token history before speculation resumes.") + if self._mapping.pp_size != 1 or self._mapping.cp_size != 1: + return "Unified DSpark KV cache requires PP=1 and CP=1." + if self._kv_connector_manager is not None: + return ("Unified DSpark KV cache does not yet support " + "KV cache connectors.") + transceiver_config = self._cache_transceiver_config + if self._is_disagg and (transceiver_config is None + or transceiver_config.transceiver_runtime + != "PYTHON" + or transceiver_config.backend != "NIXL"): + return ("DSpark draft-state transfer requires the PYTHON " + "NIXL transceiver on both workers.") + return None + + def _get_standalone_draft_layout(self) -> StandaloneDraftLayout: + """Describe distinct draft layers without borrowing target shapes.""" + if self._is_embedded_dspark(): + draft = self._model_engine.model.draft_model + return StandaloneDraftLayout( + num_layers=draft.num_stages, + num_kv_heads=1, + head_dim=int(draft._attn_params["head_dim"]), + dtype=torch.bfloat16, + extra_tokens=0, + attention_backend="DSv4", + kv_factor=1, + window_size=int(draft._attn_params["window_size"]), + ) + config = self._draft_config.pretrained_config + num_heads = config.num_attention_heads + num_kv_heads = getattr(config, "num_key_value_heads", num_heads) + head_dim = getattr(config, "head_dim", None) + if head_dim is None: + head_dim = config.hidden_size // num_heads + attention_tp_size = (1 if self._mapping.enable_attention_dp else + self._mapping.tp_size) + if (num_kv_heads % attention_tp_size != 0 + and attention_tp_size % num_kv_heads != 0): + raise ValueError( + "Standalone DSpark KV heads must divide attention TP or be " + "divisible by it.") + attention_backend = self._speculative_config.attention_backend + if attention_backend == "AUTO": + attention_backend = self._model_engine.model.draft_model.dflash_attention_backend + return StandaloneDraftLayout( + num_layers=config.num_hidden_layers, + num_kv_heads=max(1, num_kv_heads // attention_tp_size), + head_dim=head_dim, + dtype=torch.bfloat16, + extra_tokens=self._speculative_config.max_draft_len + 1, + attention_backend=attention_backend, + ) + def _should_create_separate_draft_kv_cache(self) -> bool: """ Check if we need a separate draft KV cache manager for one-model mode. @@ -1841,6 +1967,8 @@ def _should_create_separate_draft_kv_cache(self) -> bool: if self._speculative_config is None: # No drafter at all, so there is nothing to give a manager to. return False + if self._uses_unified_standalone_draft_cache(): + return False # Narrower than is_external_drafter(): PARD and DRAFT_TARGET_ONE_MODEL # never reach the arena this carve-out exists for. spec_dec_mode = self._speculative_config.spec_dec_mode @@ -2793,12 +2921,16 @@ def _create_kv_cache_manager( disable_overlap_scheduler: bool = False, cold_page_codec_provider: Optional[object] = None, kv_events_config: Optional[KVEventsConfig] = None, + standalone_draft_layout: Optional[StandaloneDraftLayout] = None, joint_kv_cache_reuse: bool = False, max_cuda_graph_batch_size: Optional[int] = None) -> KVCacheManager: """ Returns: A KVCacheManager instance for the given model engine or model config """ + if standalone_draft_layout is not None and not issubclass( + kv_cache_manager_cls, KVCacheManagerV2): + raise ValueError("Standalone draft layouts require KVCacheManagerV2.") if cold_page_codec_provider is not None and not issubclass( kv_cache_manager_cls, KVCacheManagerV2): raise ValueError( @@ -2966,6 +3098,9 @@ def _create_kv_cache_manager( "cold_page_codec_provider"] = cold_page_codec_provider manager_extra_kwargs["kv_events_config"] = kv_events_config manager_extra_kwargs["joint_kv_cache_reuse"] = joint_kv_cache_reuse + if standalone_draft_layout is not None: + manager_extra_kwargs[ + "standalone_draft_layout"] = standalone_draft_layout manager_extra_kwargs[ "disable_overlap_scheduler"] = disable_overlap_scheduler # Vocab size also enables multimodal event decoding and its per-block diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py index 4205de376ff5..33d8f7ffb5cc 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py @@ -119,6 +119,7 @@ request_context, ) from ..scheduler import ScheduledRequests +from .standalone_draft_cache import StandaloneDraftHistory, StandaloneDraftLayout if TYPE_CHECKING: from tensorrt_llm._torch.attention.backends.interface import AttentionMetadata @@ -525,11 +526,14 @@ def _estimate_cache_size_components( *, scratch: bool, generation_capacity_headroom: int, + standalone_draft_reserve: int = 0, ) -> tuple[int, int, int]: """Return context/generation bytes per token and generation bytes per request. Static profiling and runtime quota conversion must charge the same SWA - retention pages and context scratch space. Resume-watermark normalization + retention pages and context scratch space. The standalone draft reserve + is a shared allocation envelope, including windowed attention groups; + it does not represent committed history. Resume-watermark normalization is separate from these usable-capacity costs. """ full_attn_size = _estimate_full_attn_size_per_token(layer_sizes, attention_windows) @@ -547,7 +551,7 @@ def _estimate_cache_size_components( return ( full_attn_size + context_swa_size, full_attn_size + generation_swa_size, - generation_swa_per_request, + generation_swa_per_request + standalone_draft_reserve * sum(layer_sizes), ) @@ -965,6 +969,8 @@ def _update_kv_cache_draft_token_location( run_kv_cache_relocation = True if not run_kv_cache_relocation: return + if getattr(cache_manager, "draft_layout", None) is not None: + raise ValueError("Unified standalone draft KV does not support tree-token relocation") requests = scheduled_batch.all_requests() ( accepted_draft_token_offsets, @@ -1105,7 +1111,23 @@ def _settle_context_cursor(req: LlmRequest, reuse: int, tokens_per_block: int) - req.context_chunk_size = req.context_remaining_length +def get_draft_cache_unsupported_reason(kv_cache_config: KvCacheConfig) -> Optional[str]: + """Restrictions shared by creator admission and direct manager construction.""" + if kv_cache_config.enable_swa_scratch_reuse: + return ( + "Unified DSpark draft KV cannot use SWA scratch reuse; " + "draft prefill requires ordinary pages. " + "set kv_cache_config.enable_swa_scratch_reuse=False" + ) + if kv_cache_config.pool_ratio is not None: + return "Unified DSpark draft KV does not yet support explicit pool_ratio" + return None + + class KVCacheManagerV2(BaseResourceManager): + draft_layout: Optional[StandaloneDraftLayout] = None + draft_layer_ids: tuple[int, ...] = () + _standalone_draft_reserve: int = 0 # Filled lazily by _cold_pool_group_membership(); the grouping is fixed after construction. # Declared on the class so it is present even when an instance is built without running __init__. _cold_pool_group_membership_cache: Optional[tuple[tuple[int, frozenset[int]], ...]] = None @@ -1145,9 +1167,21 @@ def __init__( is_estimating_kv_cache: bool = False, cold_page_codec_provider: Optional[object] = None, joint_kv_cache_reuse: bool = False, + standalone_draft_layout: Optional[StandaloneDraftLayout] = None, max_cuda_graph_batch_size: Optional[int] = None, **kwargs, ) -> None: + self.draft_layout = standalone_draft_layout + self.draft_layer_ids: tuple[int, ...] = () + self.draft_history: dict[int, StandaloneDraftHistory] = {} + self._draft_dummy_request_ids: set[int] = set() + self._standalone_draft_reserve = ( + standalone_draft_layout.extra_tokens if standalone_draft_layout is not None else 0 + ) + if standalone_draft_layout is not None: + reason = get_draft_cache_unsupported_reason(kv_cache_config) + if reason is not None: + raise ValueError(reason) self.mapping = mapping self.dtype = dtype self._validate_speculative_config(spec_config) @@ -1246,6 +1280,7 @@ def __init__( self._kv_reserve_draft_tokens, self._generation_kv_capacity_headroom = ( _get_generation_kv_capacity(spec_config, is_draft=self.is_draft) ) + self._generation_kv_capacity_headroom += self._standalone_draft_reserve self.event_buffer_max_size = kv_cache_config.event_buffer_max_size self.enable_stats = enable_stats @@ -1509,6 +1544,7 @@ def append_to_kv_heads_per_layer( cache_tiers=cache_tiers, ) config = self._build_cache_config(config) + config = self._append_standalone_draft_layers(config) config = self._remove_zero_size_buffers(config) has_host_cache_tier = any( isinstance(tier, HostCacheTierConfig) for tier in config.cache_tiers @@ -1997,7 +2033,8 @@ def _fill_fresh_kv_pages(self, request_id: int) -> None: # Same conversion the attention backends index the buffer # with: layers in one pool can carry different scales. scale = self.get_layer_page_index_scale(layer_idx) - pages = [page * scale // self.kv_factor for page in fresh] + factor = self.get_layer_kv_factor(layer_idx) + pages = [page * scale // factor for page in fresh] pages = [page for page in pages if 0 <= page < buffer.shape[0]] if pages: if _fill_kv_pages(buffer, pages, self._fresh_page_fill): @@ -2034,6 +2071,10 @@ def _get_pool_roles(self, pool_id: int) -> Tuple[DataRole, Optional[DataRole]]: When present, role B must be addressable from role A using a constant page-index offset. """ + if self.draft_layout is not None: + layer_id = int(self.impl.layer_grouping[pool_id][0]) + if self._is_standalone_draft_layer(layer_id): + return Role.KEY, Role.VALUE if self.draft_layout.kv_factor == 2 else None role_b = None if self.kv_cache_type == CacheTypeCpp.SELFKONLY else Role.VALUE return Role.KEY, role_b @@ -2043,6 +2084,11 @@ def _get_block_scale_role(self, role_a: DataRole, layer_id: int) -> Optional[Dat return None return Role.KEY_BLOCK_SCALE + def _get_layer_block_scale_role(self, layer_id: int, role_a: DataRole) -> Optional[DataRole]: + if self._is_standalone_draft_layer(layer_id): + return None + return self._get_block_scale_role(role_a, layer_id) + def _build_pool_mapping_tensors(self): """Build the (kv_cache_pool_pointers, kv_cache_pool_mapping) tensors. @@ -2065,7 +2111,7 @@ def _build_pool_mapping_tensors(self): ] ) if self.dtype == DataType.NVFP4: - block_scale_role = self._get_block_scale_role(role_a, layer_id) + block_scale_role = self._get_layer_block_scale_role(layer_id, role_a) block_scale_pool_pointers_list.append( [ self.impl.get_mem_pool_base_address( @@ -2097,7 +2143,7 @@ def _build_pool_mapping_tensors(self): # shift lands the origin on the pool's slot-0 scale address. # This keeps block_scale_offset == offset without depending on # the non-contractual layer_grouping order. - block_scale_role = self._get_block_scale_role(role_a, layer_id) + block_scale_role = self._get_layer_block_scale_role(layer_id, role_a) if block_scale_role is not None: rep_offset = self._kv_pool_mapping_offset(layer_id, pool_id, key_base_addr) scale_stride = ( @@ -2137,7 +2183,7 @@ def _build_pool_mapping_tensors(self): if self.dtype != DataType.NVFP4 or role_a != Role.KEY: block_scale_offset = None else: - block_scale_role = self._get_block_scale_role(role_a, layer_id) + block_scale_role = self._get_layer_block_scale_role(layer_id, role_a) if block_scale_role is None: block_scale_offset = None else: @@ -2262,7 +2308,7 @@ def _kv_pool_mapping_offset( return exact_div( addr_offset, self.get_layer_bytes_per_token(layer_id, Role.KEY) - * self.kv_factor + * self.get_layer_kv_factor(self.pp_layers[layer_id]) * self.tokens_per_block, ) @@ -2274,6 +2320,15 @@ def _get_runtime_cache_size_layer_components(self) -> tuple[List[int], List[Opti self.get_layer_bytes_per_token(local_layer_idx=local_layer_idx, data_role=Role.ALL) ) attention_windows.append(self.max_attention_window_vec[local_layer_idx]) + # Quota sizing precedes specialized target-layer construction. Include + # draft storage then; after append it is already in the per-layer arrays. + if self.draft_layout is not None and not self.draft_layer_ids: + layer_sizes.extend( + [self.draft_layout.bytes_per_layer_token] * self.draft_layout.num_layers + ) + attention_windows.extend( + [self.draft_layout.retention_window_size] * self.draft_layout.num_layers + ) return layer_sizes, attention_windows def _get_generation_request_capacity(self) -> int: @@ -2296,15 +2351,16 @@ def _get_max_tokens_from_quota_impl(self, quota: int) -> float: ( context_size_per_token, generation_size_per_token, - generation_swa_size_per_request, + generation_size_per_request, ) = _estimate_cache_size_components( layer_sizes, attention_windows, self.tokens_per_block, scratch=self.enable_swa_scratch_reuse, generation_capacity_headroom=self._generation_kv_capacity_headroom, + standalone_draft_reserve=self._standalone_draft_reserve, ) - size_per_batch = self._get_generation_request_capacity() * generation_swa_size_per_request + size_per_batch = self._get_generation_request_capacity() * generation_size_per_request if quota < size_per_batch: return 0 context_limit_quota = self.max_num_tokens * context_size_per_token + size_per_batch @@ -2332,20 +2388,21 @@ def _get_quota_from_max_tokens_impl(self, max_tokens: int) -> int: ( context_size_per_token, generation_size_per_token, - generation_swa_size_per_request, + generation_size_per_request, ) = _estimate_cache_size_components( layer_sizes, attention_windows, self.tokens_per_block, scratch=self.enable_swa_scratch_reuse, generation_capacity_headroom=self._generation_kv_capacity_headroom, + standalone_draft_reserve=self._standalone_draft_reserve, ) context_tokens = min(max_tokens, self.max_num_tokens) generation_tokens = max_tokens - context_tokens return int( context_tokens * context_size_per_token + generation_tokens * generation_size_per_token - + self._get_generation_request_capacity() * generation_swa_size_per_request + + self._get_generation_request_capacity() * generation_size_per_request ) def _get_event_num_blocks_per_cache_level( @@ -2874,6 +2931,66 @@ def _build_cache_config(self, config: KVCacheManagerConfigPy) -> KVCacheManagerC """Customize the general cache config for a specialized cache manager.""" return config + def _append_standalone_draft_layers( + self, config: KVCacheManagerConfigPy, *, register_model_layers: bool = True + ) -> KVCacheManagerConfigPy: + layout = self.draft_layout + if layout is None: + return config + first_global_id = max(self.pp_layers, default=-1) + 1 + first_local_id = len(config.layers) + self.draft_layer_ids = tuple(range(first_global_id, first_global_id + layout.num_layers)) + layers = list(config.layers) + for offset, global_id in enumerate(self.draft_layer_ids): + local_id = first_local_id + offset + layers.append( + AttentionLayerConfig( + layer_id=LayerId(local_id), + buffers=[ + BufferConfig( + role=role, + size=layout.bytes_per_layer_token + // layout.kv_factor + * self.tokens_per_block, + ) + for role in (Role.KEY, Role.VALUE)[: layout.kv_factor] + ], + sliding_window_size=layout.retention_window_size, + cache_domain="standalone_draft", + ) + ) + self.layer_offsets[global_id] = local_id + if register_model_layers: + self.pp_layers.append(global_id) + self.num_kv_heads_per_layer.append(layout.num_kv_heads) + self.total_num_kv_heads_per_layer.append(layout.num_kv_heads) + self.head_dim_per_layer.append(layout.head_dim) + self.max_attention_window_vec.append(layout.retention_window_size) + if register_model_layers: + self.num_local_layers = len(self.pp_layers) + self.num_layers += layout.num_layers + + def reserve_scratch(batch: BatchDesc) -> BatchDesc: + return BatchDesc( + [ + KVCacheDesc( + capacity=desc.capacity + (layout.extra_tokens if desc.capacity else 0), + history_length=desc.history_length, + ) + for desc in batch.kv_caches + ], + system_prompt_length=batch.system_prompt_length, + ) + + return replace( + config, + layers=layers, + constraints=[reserve_scratch(batch) for batch in config.constraints], + typical_step=reserve_scratch(config.typical_step) + if config.typical_step is not None + else None, + ) + def _remove_zero_size_buffers(self, config: KVCacheManagerConfigPy) -> KVCacheManagerConfigPy: """Exclude empty buffers before creating the runtime storage pools.""" if config.layers and all( @@ -2898,6 +3015,7 @@ def _remove_zero_size_buffers(self, config: KVCacheManagerConfigPy) -> KVCacheMa buffers=buffers, sliding_window_size=layer.sliding_window_size, num_sink_tokens=layer.num_sink_tokens, + cache_domain=layer.cache_domain, ) ) else: @@ -3010,50 +3128,173 @@ def is_attention_layer(self, layer_idx: int) -> bool: def get_buffers(self, layer_idx: int, kv_layout: str = "NHD") -> Optional[torch.Tensor]: layer_offset = self.layer_offsets[layer_idx] - addr_key = self.impl.get_mem_pool_base_address(layer_offset, Role.KEY, PageIndexMode.SHARED) - if self.kv_cache_type != CacheTypeCpp.SELFKONLY: + if self._is_standalone_draft_layer(layer_offset): + return self.get_draft_buffers(self.draft_layer_ids.index(layer_idx), kv_layout) + return self._get_kv_buffer_view( + layer_offset, + dtype=self.dtype, + kv_factor=self.kv_factor, + num_kv_heads=self.num_kv_heads_per_layer[layer_offset], + head_dim=self.head_dim_per_layer[layer_offset], + kv_layout=kv_layout, + ) + + def _get_kv_buffer_view( + self, + layer_id: int, + *, + dtype: DataType | torch.dtype, + kv_factor: int, + num_kv_heads: int, + head_dim: int, + kv_layout: str, + ) -> torch.Tensor: + """View shared pool pages for a manager-local layer and its KV geometry.""" + if kv_layout not in ("NHD", "HND"): + raise ValueError(f"Unsupported KV layout: {kv_layout}") + addr_key = self.impl.get_mem_pool_base_address(layer_id, Role.KEY, PageIndexMode.SHARED) + if kv_factor == 2: addr_value = self.impl.get_mem_pool_base_address( - layer_offset, Role.VALUE, PageIndexMode.SHARED + layer_id, Role.VALUE, PageIndexMode.SHARED ) - page_size_key = self.impl.get_page_stride(layer_offset, Role.KEY) - page_size_value = self.impl.get_page_stride(layer_offset, Role.VALUE) - - assert addr_key + page_size_value == addr_value and page_size_key == page_size_value - - assert kv_layout in ["NHD", "HND"], f"Unsupported kv_layout: {kv_layout}" + page_size_key = self.impl.get_page_stride(layer_id, Role.KEY) + page_size_value = self.impl.get_page_stride(layer_id, Role.VALUE) + if addr_key + page_size_key != addr_value or page_size_key != page_size_value: + raise ValueError("K/V buffers must have adjacent equal-sized pages") element_per_container = 1 - dtype = self.dtype if dtype == DataType.NVFP4: element_per_container = 2 dtype = torch.int8 - layer_head_dim = self.head_dim_per_layer[layer_offset] - if kv_layout == "NHD": - shape = [ - self.impl.get_page_index_upper_bound(layer_offset, Role.KEY) // self.kv_factor, - self.kv_factor, - self.tokens_per_block, - self.num_kv_heads_per_layer[layer_offset], - layer_head_dim // element_per_container, - ] - else: - shape = [ - self.impl.get_page_index_upper_bound(layer_offset, Role.KEY) // self.kv_factor, - self.kv_factor, - self.num_kv_heads_per_layer[layer_offset], - self.tokens_per_block, - layer_head_dim // element_per_container, - ] - + dimensions = ( + [self.tokens_per_block, num_kv_heads] + if kv_layout == "NHD" + else [num_kv_heads, self.tokens_per_block] + ) return convert_to_torch_tensor( TensorWrapper( addr_key, dtype, - shape, + [ + self.impl.get_page_index_upper_bound(layer_id, Role.KEY) // kv_factor, + kv_factor, + *dimensions, + head_dim // element_per_container, + ], ) ) + def _is_standalone_draft_layer(self, local_layer_idx: int) -> bool: + return self.draft_layout is not None and any( + self.layer_offsets[layer_id] == local_layer_idx for layer_id in self.draft_layer_ids + ) + + def get_layer_cache_dtype(self, layer_idx: int) -> DataType: + if layer_idx in self.draft_layer_ids: + return DataType.BF16 if self.draft_layout.dtype == torch.bfloat16 else DataType.HALF + return self.dtype + + def get_layer_kv_factor(self, layer_idx: int) -> int: + return self.draft_layout.kv_factor if layer_idx in self.draft_layer_ids else self.kv_factor + + def get_draft_buffers(self, local_layer_idx: int, kv_layout: str = "HND") -> torch.Tensor: + """View authoritative draft pages with their independent geometry.""" + layout = self.draft_layout + if layout is None or not 0 <= local_layer_idx < layout.num_layers: + raise ValueError("No standalone draft cache layer at this index") + layer_idx = self.draft_layer_ids[local_layer_idx] + return self._get_kv_buffer_view( + self.layer_offsets[layer_idx], + dtype=self.get_layer_cache_dtype(layer_idx), + kv_factor=layout.kv_factor, + num_kv_heads=layout.num_kv_heads, + head_dim=layout.head_dim, + kv_layout=kv_layout, + ) + + def get_draft_block_table( + self, + request_ids: List[int], + histories: Optional[Sequence[StandaloneDraftHistory]] = None, + ) -> torch.Tensor: + """Current rank-local page mappings; unused tail entries point to page zero.""" + for request_id in request_ids: + cache = self.kv_cache_map.get(request_id) + if cache is None or not cache.is_active: + raise ValueError(f"Standalone draft request {request_id} has no active cache") + batch_indices = self.get_batch_cache_indices(request_ids, self.draft_layer_ids[0]) + table = torch.zeros((len(request_ids), self.max_blocks_per_seq), dtype=torch.int32) + for row, indices in enumerate(batch_indices): + cache = self.kv_cache_map[request_ids[row]] + if len(indices) != cache.num_blocks or len(indices) > self.max_blocks_per_seq: + raise ValueError("Standalone draft cache has an incomplete or oversized page table") + first_required_block = 0 + if self.draft_layout.window_size is not None: + history = ( + histories[row] + if histories is not None + else self.get_draft_history(request_ids[row]) + ) + history_start = ( + history.position - history.valid_length + if history is not None + else max(0, cache.history_length - self.draft_layout.window_size) + ) + first_required_block = history_start // self.tokens_per_block + if any(index == BAD_PAGE_INDEX for index in indices[first_required_block:]): + raise ValueError( + "Draft cache contains missing pages in its required history or scratch" + ) + table[row, : len(indices)] = torch.tensor(indices, dtype=torch.int32) + return table + + def get_draft_history(self, request_id: int) -> Optional[StandaloneDraftHistory]: + return self.draft_history.get(request_id) + + def set_draft_history(self, request_id: int, valid_length: int, position: int) -> None: + cache = self.kv_cache_map.get(request_id) + if cache is None or not cache.is_active: + raise ValueError( + f"Standalone draft request {request_id} has no active cache allocation" + ) + history = StandaloneDraftHistory(valid_length, position) + if history.position > cache.capacity: + raise ValueError("Standalone draft history exceeds allocated capacity") + if ( + self.draft_layout.window_size is not None + and history.valid_length > self.draft_layout.window_size + ): + raise ValueError("Draft history exceeds its retention window") + self.draft_history[request_id] = history + + def export_draft_history(self, request_id: int) -> Optional[dict]: + if self.draft_layout is None: + return None + history = self.get_draft_history(request_id) + if history is None: + raise ValueError( + f"Standalone draft request {request_id} has no valid history to transfer" + ) + return { + "valid_length": history.valid_length, + "position": history.position, + "layout": self.draft_layout.transfer_identity(), + } + + def restore_draft_history(self, request_id: int, metadata: dict) -> None: + if ( + self.draft_layout is None + or metadata.get("layout") != self.draft_layout.transfer_identity() + ): + raise ValueError("Standalone draft transfer layout does not match the receiving worker") + valid_length = metadata.get("valid_length") + position = metadata.get("position") + history = StandaloneDraftHistory(valid_length, position) + # Validate receiver-local allocation before publishing history. + self.get_draft_block_table([request_id], [history]) + self.set_draft_history(request_id, valid_length, position) + def get_index_k_buffer( self, layer_idx: int, @@ -3188,7 +3429,9 @@ def get_num_available_tokens( replicates all tokens on every rank, so both bounds constrain the same request-length variable. """ - extra_tokens = self.num_extra_kv_tokens + max_num_draft_tokens + extra_tokens = ( + self.num_extra_kv_tokens + max_num_draft_tokens + self._standalone_draft_reserve + ) # Token num upper bound is the maximum number of tokens that can be allocated in the kv cache manager. # We need to add extra tokens to the token num upper bound to account for the extra tokens. clamped = ( @@ -3215,9 +3458,10 @@ def get_num_free_blocks(self) -> int: assert not set(self.kv_cache_map) - reserved, ( "get_num_free_blocks is only used when the kv cache manager is empty" ) - max_num_pages = max( + max_num_blocks = max( [ self.impl.get_page_index_upper_bound(layer_id, Role.KEY) + // self.get_layer_kv_factor(self.pp_layers[layer_id]) for layer_id in typed_range(LayerId(self.num_local_layers)) ] ) @@ -3226,7 +3470,7 @@ def get_num_free_blocks(self) -> int: # page (the guard is today). Summing ``num_blocks`` keeps this a page # count if a reservation ever spans more than one. reserved_pages = sum(int(self.kv_cache_map[req_id].num_blocks) for req_id in reserved) - return max_num_pages // self.kv_factor - reserved_pages + return max_num_blocks - reserved_pages def commit_scheduled_kv_cache_stats(self, scheduled_batch: ScheduledRequests) -> None: if self.is_draft or (not self.enable_stats and not self._request_stats_enabled_ids): @@ -3575,6 +3819,15 @@ def prepare_context_cache(self, req: LlmRequest, reuse_limit: int | None = None) kv_cache.enable_swa_scratch_reuse = False if not self._resume_and_restore(req.py_request_id, kv_cache): return None + if self.draft_layout is not None and req.py_request_id not in self.draft_history: + # The reuse match covers every cache domain. Initialize once, + # after resume succeeds, so overlap/chunk retries cannot rewind + # draft history that has already advanced on the device. + reused = kv_cache.num_committed_tokens + valid_length = reused + if self.draft_layout.window_size is not None: + valid_length = min(valid_length, self.draft_layout.window_size) + self.set_draft_history(req.py_request_id, valid_length, reused) return kv_cache.num_committed_tokens # Subsequent chunk: cache must exist from first chunk. It may be @@ -3629,7 +3882,12 @@ def resize_context(self, req: LlmRequest, num_tokens: int) -> bool: if kv_cache is None: return False - target = req.context_current_position + num_tokens + self.num_extra_kv_tokens + target = ( + req.context_current_position + + num_tokens + + self.num_extra_kv_tokens + + self._standalone_draft_reserve + ) capacity = max(kv_cache.capacity, target) pre_cap = kv_cache.capacity @@ -3684,7 +3942,12 @@ def prepare_disagg_gen_init(self, req: LlmRequest) -> bool: # Helix requests carry the rank-local strided slice in prompt_len; # the global ledger sizes off the full prompt instead. prompt_len = req.total_input_len_cp if self._has_cp_helix else req.prompt_len - target = prompt_len + get_draft_token_length(req) + self.num_extra_kv_tokens + target = ( + prompt_len + + get_draft_token_length(req) + + self.num_extra_kv_tokens + + self._standalone_draft_reserve + ) capacity = max(kv_cache.capacity, target) pre_cap = kv_cache.capacity @@ -4985,7 +5248,9 @@ def release_resources( release_resources(req) return None kv_cache.stop_committing() - dummy_capacity = token_num + self.num_extra_kv_tokens + dummy_capacity = ( + token_num + self.num_extra_kv_tokens + self._standalone_draft_reserve + ) if is_gen and not materialize_history: kv_cache.enable_swa_scratch_reuse = False # Need to hint the committed history to activate stale-block @@ -5114,6 +5379,9 @@ def free_resources(self, request: LlmRequest, pin_on_release: bool = False): if self.conversation_manager is not None: self.conversation_manager.finish_request(request) self._allocated_draft_lens.pop(request.py_request_id, None) + if self.draft_layout is not None: + self.draft_history.pop(request.py_request_id, None) + self._draft_dummy_request_ids.discard(request.py_request_id) self._request_stats_enabled_ids.discard(request.py_request_id) # The next owner of these pages fills them again; keeping the set would # both leak and let a recycled page skip its fill. @@ -5162,6 +5430,7 @@ def get_batch_cache_indices( is_kv_aggregate=not raw_indices, num_blocks_per_seq=num_blocks_per_seq, index_scale=index_scale, + kv_factor=self.get_layer_kv_factor(layer_idx) if layer_idx is not None else None, ) def _get_batch_cache_indices_by_pool_id( @@ -5172,11 +5441,12 @@ def _get_batch_cache_indices_by_pool_id( is_kv_aggregate: bool = True, num_blocks_per_seq: Optional[Sequence[int]] = None, index_scale: Optional[int] = None, + kv_factor: Optional[int] = None, ) -> List[List[int]]: if is_kv_aggregate: # Div by kv_factor to index kv cache with size # [num_blocks, kv_factor, tokens_per_block, num_kv_heads, head_dim] - div_factor = self.kv_factor + div_factor = self.kv_factor if kv_factor is None else kv_factor else: div_factor = 1 @@ -5232,7 +5502,9 @@ def get_batch_cache_indices_flat( # scale so this flat block table matches get_batch_cache_indices() # and never feeds out-of-range page ids to FlashInfer. scale = self.get_layer_page_index_scale(layer_idx) - div_factor = self.kv_factor + div_factor = ( + self.get_layer_kv_factor(layer_idx) if layer_idx is not None else self.kv_factor + ) out_tensor = torch.empty(sum(num_blocks), dtype=torch.int32, pin_memory=prefer_pinned()) out = out_tensor.numpy() @@ -5253,6 +5525,8 @@ def get_batch_cache_indices_flat( return out_tensor def get_cache_bytes_per_token(self) -> int: + if self.draft_layout is not None: + return sum(self._get_runtime_cache_size_layer_components()[0]) data_roles = [Role.KEY] if self.kv_cache_type != CacheTypeCpp.SELFKONLY: data_roles.append(Role.VALUE) @@ -5268,6 +5542,14 @@ def get_cache_bytes_per_token(self) -> int: ) def get_layer_bytes_per_token(self, local_layer_idx: int, data_role: Role): + if self._is_standalone_draft_layer(local_layer_idx): + if data_role == Role.ALL: + return self.draft_layout.bytes_per_layer_token + if data_role in (Role.KEY, Role.VALUE): + if data_role == Role.VALUE and self.draft_layout.kv_factor == 1: + return 0 + return self.draft_layout.bytes_per_layer_token // self.draft_layout.kv_factor + return 0 if self.dtype not in ( DataType.FP8, DataType.HALF, @@ -5378,6 +5660,9 @@ def shutdown(self): for kv_cache in self.kv_cache_map.values(): kv_cache.close() self.kv_cache_map.clear() + if self.draft_layout is not None: + self.draft_history.clear() + self._draft_dummy_request_ids.clear() self._disagg_receive_ready.clear() self._request_stats_enabled_ids.clear() self._fresh_pages_filled.clear() @@ -5423,6 +5708,7 @@ def get_cache_size_per_token( max_num_tokens: int = 0, spec_config=None, is_draft: bool = False, + draft_layout: Optional[StandaloneDraftLayout] = None, **kwargs, ): layer_sizes, attention_windows = _get_static_cache_size_layer_components( @@ -5455,10 +5741,16 @@ def get_cache_size_per_token( _, generation_capacity_headroom = _get_generation_kv_capacity( spec_config, is_draft=is_draft ) + draft_reserve = 0 + if draft_layout is not None: + layer_sizes.extend([draft_layout.bytes_per_layer_token] * draft_layout.num_layers) + attention_windows.extend([draft_layout.retention_window_size] * draft_layout.num_layers) + draft_reserve = draft_layout.extra_tokens + generation_capacity_headroom += draft_reserve ( context_size_per_token, cache_size_per_token, - swa_size_per_request, + generation_size_per_request, ) = _estimate_cache_size_components( layer_sizes, attention_windows, @@ -5469,11 +5761,12 @@ def get_cache_size_per_token( and not is_draft ), generation_capacity_headroom=generation_capacity_headroom, + standalone_draft_reserve=draft_reserve, ) # The affine slope covers all tokens; context additionally retains SWA # pages for the current token batch beyond the generation windows. fixed_cost = ( - swa_size_per_request * max_batch_size + generation_size_per_request * max_batch_size + (context_size_per_token - cache_size_per_token) * max_num_tokens ) bytes_per_slot = _get_single_swa_pool_slot_bytes( @@ -5584,6 +5877,14 @@ def update_resources( if req.state in (LlmRequestState.GENERATION_COMPLETE, LlmRequestState.CONTEXT_INIT) else kv_cache.capacity - rewind_len ) + if self.draft_layout is not None and new_capacity is not None: + draft_history = self.get_draft_history(req.py_request_id) + if draft_history is not None: + # Rejected target verification slots do not revoke draft + # feature history. Retain it and its next-forward scratch. + new_capacity = max( + new_capacity, draft_history.position + self._standalone_draft_reserve + ) history_length = ( None # Reuse (history's consumer) is disabled under helix, and @@ -5684,6 +5985,8 @@ def _create_kv_cache( if is_dummy: self.impl.mark_stats_excluded(request_id) kv_cache.discard_pending_stats() + if self.draft_layout is not None: + self._draft_dummy_request_ids.add(request_id) index = self.index_mapper.add_new_sequence(request_id) for i in range(self.max_beam_width): for pool_idx in range(self.num_pools): diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache/mamba_cache_manager.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/mamba_cache_manager.py index ad417618950b..4741be8cace5 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache/mamba_cache_manager.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/mamba_cache_manager.py @@ -33,7 +33,10 @@ from tensorrt_llm._torch.disaggregation.resource.page import (MapperKind, RoleLayout) from tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2 import ( - _RESERVED_REQUEST_IDS, BlockReusePolicy, KVCacheManagerV2, Role) + _RESERVED_REQUEST_IDS, BlockReusePolicy, KVCacheManagerV2, Role, + _estimate_cache_size_components) +from tensorrt_llm._torch.pyexecutor.kv_cache.standalone_draft_cache import \ + StandaloneDraftLayout from tensorrt_llm._torch.pyexecutor.kv_cache_stats import \ KVCacheV2IterationStatsReport from tensorrt_llm._torch.pyexecutor.llm_request import ( @@ -2141,6 +2144,7 @@ def _estimate_mamba_hybrid_cache_cost( cap_partial_attention_snapshots: bool, is_draft: bool = False, use_separate_draft_kv_cache: bool = False, + draft_layout: Optional[StandaloneDraftLayout] = None, **kwargs, ) -> Tuple[int, int]: spec_config = kwargs.get("spec_config") @@ -2203,6 +2207,19 @@ def _estimate_mamba_hybrid_cache_cost( if (has_unaligned_periodic_snapshot and not cap_partial_attention_snapshots): regular_slope += math.ceil(attention_block_bytes / interval) + if draft_layout is not None: + # Hybrid target attention is full attention. Reserve capture/noise + # capacity only in its attention pools and the distinct draft pools; + # recurrent state retains the fixed/snapshot accounting above. + _, attention_slope, draft_fixed = _estimate_cache_size_components( + [attention_slope, draft_layout.bytes_per_token], + [None, None], + tokens_per_block, + scratch=False, + generation_capacity_headroom=1, + standalone_draft_reserve=draft_layout.extra_tokens, + ) + intercept += max_batch_size * draft_fixed return attention_slope + regular_slope, intercept @@ -3786,7 +3803,9 @@ def get_cache_size_per_token(model_config, ) def _is_local_mamba_layer(self, local_layer_idx: int) -> bool: - return self._mamba_layer_mask[self.pp_layers[local_layer_idx]] + layer_id = self.pp_layers[local_layer_idx] + return layer_id < len( + self._mamba_layer_mask) and self._mamba_layer_mask[layer_id] def _get_pool_roles(self, pool_id: int) -> Tuple[DataRole, Optional[DataRole]]: @@ -4201,7 +4220,7 @@ def get_num_free_blocks(self) -> int: else: attention_pages.append( self.impl.get_page_index_upper_bound(layer_id, Role.KEY) // - self.kv_factor) + self.get_layer_kv_factor(self.pp_layers[local_layer_idx])) if attention_pages: return max(attention_pages) return max(ssm_pages) if ssm_pages else 0 diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache/standalone_draft_cache.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/standalone_draft_cache.py new file mode 100644 index 000000000000..8b2e76df2eff --- /dev/null +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/standalone_draft_cache.py @@ -0,0 +1,118 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Storage and history contracts for manager-owned DSpark/DFlash drafters.""" + +from dataclasses import dataclass +from typing import TYPE_CHECKING + +import torch + +if TYPE_CHECKING: + from .kv_cache_manager_v2 import KVCacheManagerV2 + + +@dataclass(frozen=True) +class StandaloneDraftLayout: + """Rank-local draft storage, independent of target geometry and retention.""" + + num_layers: int + num_kv_heads: int + head_dim: int + dtype: torch.dtype + extra_tokens: int + attention_backend: str + kv_factor: int = 2 + window_size: int | None = None + + def __post_init__(self) -> None: + if min(self.num_layers, self.num_kv_heads, self.head_dim) <= 0: + raise ValueError("Standalone draft cache dimensions must be positive") + if self.extra_tokens < 0: + raise ValueError("Standalone draft scratch capacity must be nonnegative") + if self.dtype not in (torch.float16, torch.bfloat16): + raise ValueError("Standalone draft KV supports FP16 and BF16 storage") + if self.attention_backend not in ("VANILLA", "TRTLLM", "DSv4"): + raise ValueError("Unsupported managed draft attention backend") + if self.kv_factor not in (1, 2): + raise ValueError("Draft KV storage requires one or two planes") + if self.window_size is not None and self.window_size <= 0: + raise ValueError("Draft history window must be positive") + + @property + def bytes_per_layer_token(self) -> int: + return self.kv_factor * self.num_kv_heads * self.head_dim * self.dtype.itemsize + + @property + def retention_window_size(self) -> int | None: + # V2 retains window_size - 1 committed rows between forwards. + return self.window_size + 1 if self.window_size is not None else None + + @property + def bytes_per_token(self) -> int: + return self.num_layers * self.bytes_per_layer_token + + def transfer_identity(self) -> dict: + return { + "num_layers": self.num_layers, + "num_kv_heads": self.num_kv_heads, + "head_dim": self.head_dim, + "dtype": str(self.dtype), + "attention_backend": self.attention_backend, + "kv_factor": self.kv_factor, + "window_size": self.window_size, + } + + +@dataclass(frozen=True) +class StandaloneDraftHistory: + """Committed length and next absolute position, independent of target/scratch state.""" + + valid_length: int + position: int + + def __post_init__(self) -> None: + if type(self.valid_length) is not int or type(self.position) is not int: + raise ValueError("Standalone draft history requires integer length and position") + if self.valid_length < 0 or self.position < self.valid_length: + raise ValueError("Invalid standalone draft history length or position") + + +@dataclass(frozen=True) +class DraftHistoryUpdate: + """One execution's draft history, published after its sampling event completes.""" + + manager: "KVCacheManagerV2" + request_ids: tuple[int, ...] + cache_instances: tuple[object, ...] + values_host: torch.Tensor + + @classmethod + def capture( + cls, + manager: "KVCacheManagerV2", + request_ids: tuple[int, ...], + values: torch.Tensor, + ) -> "DraftHistoryUpdate": + """Queue an independent length/position readback on the execution stream.""" + if values.shape != (len(request_ids), 2): + raise ValueError("Draft history update requires one length/position pair per request") + values_host = torch.empty( + values.shape, dtype=values.dtype, device="cpu", pin_memory=values.is_cuda + ) + values_host.copy_(values, non_blocking=True) + return cls( + manager, + tuple(request_ids), + tuple(manager.kv_cache_map[request_id] for request_id in request_ids), + values_host, + ) + + def publish(self) -> None: + """Publish completed writes without reviving released or replaced requests.""" + for request_id, cache, (length, position) in zip( + self.request_ids, self.cache_instances, self.values_host.tolist() + ): + current = self.manager.kv_cache_map.get(request_id) + if current is cache and current.is_active: + self.manager.set_draft_history(request_id, length, position) diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 2e5ce22d0296..b6d40eb43890 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -57,6 +57,7 @@ from ..models.checkpoints.base_checkpoint_loader import BaseCheckpointLoader from ..models.modeling_multimodal_mixin import (MultimodalModelMixin, _build_request_multimodal_input) +from ..models.modeling_speculative import SpecDecOneEngineForCausalLM from ..models.modeling_utils import DecoderModelForCausalLM from ..modules.mamba.mamba2_metadata import Mamba2Metadata from ..moe.expert_statistic import ExpertStatistic @@ -6125,6 +6126,13 @@ def _forward_scheduled(self, scheduled_requests: ScheduledRequests, can_run_graph, execution_promoted_context_ids, use_lora_graph=use_lora_graph) + spec_worker = (self.model.spec_worker if isinstance( + self.model, SpecDecOneEngineForCausalLM) else None) + if spec_worker is not None: + spec_worker.prepare_managed_draft_cache(self.model.draft_model, + spec_metadata, + attn_metadata, + resource_manager) if execution_promoted_context_ids: self.iter_states[ 'num_ctx_requests'] = scheduled_requests.num_context_requests @@ -6208,6 +6216,16 @@ def capture_postprocess_fn(inputs: Dict[str, Any]): restore_attn_metadata_after_draft_replay( attn_metadata, saved_draft) + if (spec_worker is not None and not self.is_warmup + and not self.cuda_graph_runner.is_warmup_only): + draft_history_update = spec_worker.snapshot_managed_draft_history( + ) + if draft_history_update is not None: + # Graph output dictionaries persist across replays. Each + # overlapped iteration must retain its own host readback. + outputs = dict(outputs) + outputs['draft_history_update'] = draft_history_update + if self.forward_pass_callable is not None: self.forward_pass_callable() diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 4c46b24ee23b..08b9aae49d92 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -7692,6 +7692,9 @@ def _sample_async(self, scheduled_batch, sample_state = self.sampler.sample_async( scheduled_batch, batch_outputs, num_context_logits_prefix_sum) + if sample_state is not None: + sample_state.draft_history_update = batch_outputs.get( + 'draft_history_update') self._maybe_record_hang_diagnostic_phase( "sampling_returned", scheduled_batch) return sample_state @@ -7718,10 +7721,15 @@ def _setup_sampler_step(self, requests: ScheduledRequests): @nvtx_range("_update_requests") def _update_requests(self, - sample_state: SampleState, + sample_state: SampleState | None, resource_manager: Optional[ResourceManager] = None): + if sample_state is None: + return try: self.sampler.update_requests(sample_state, resource_manager) + if sample_state.draft_history_update is not None: + sample_state.draft_history_update.publish() + sample_state.draft_history_update = None self._accumulate_spec_dec_stats(sample_state) except Exception as e: traceback.print_exc() diff --git a/tensorrt_llm/_torch/pyexecutor/sampler/sampler.py b/tensorrt_llm/_torch/pyexecutor/sampler/sampler.py index cca1cbd12a6c..8611a85d00a7 100644 --- a/tensorrt_llm/_torch/pyexecutor/sampler/sampler.py +++ b/tensorrt_llm/_torch/pyexecutor/sampler/sampler.py @@ -54,6 +54,7 @@ from tensorrt_llm.sampling_params import SamplingParams from ...utils import torch_multi_arange +from ..kv_cache.standalone_draft_cache import DraftHistoryUpdate from ..llm_request import LlmRequest, LlmRequestState, get_draft_token_length from ..resource_manager import ResourceManager from ..scheduler import ScheduledRequests @@ -139,6 +140,7 @@ class SampleState(Generic[GenericSampleStateTensorsHost, GenericSampleStateTenso host: Optional[GenericSampleStateTensorsHost] = None sampler_event: Optional[SamplerEvent] = None runtime_draft_len: Optional[int] = None + draft_history_update: DraftHistoryUpdate | None = None # Generic bounds not supported, https://github.com/python/typing/issues/548 diff --git a/tensorrt_llm/_torch/speculative/dflash.py b/tensorrt_llm/_torch/speculative/dflash.py index d8097dffbd58..8302bd8c1bbb 100644 --- a/tensorrt_llm/_torch/speculative/dflash.py +++ b/tensorrt_llm/_torch/speculative/dflash.py @@ -28,8 +28,9 @@ from ..attention.backends import AttentionMetadata from ..pyexecutor.kv_cache.mamba_cache_manager import MambaHybridCacheManager +from ..pyexecutor.kv_cache.standalone_draft_cache import DraftHistoryUpdate from ..pyexecutor.llm_request import ATTENTION_DP_DUMMY_REQUEST_ID -from ..pyexecutor.resource_manager import BaseResourceManager +from ..pyexecutor.resource_manager import BaseResourceManager, ResourceManagerType from .accept_stats import maybe_create_recorder from .dflash_attention import ( get_dflash_paged_append, @@ -44,6 +45,7 @@ if TYPE_CHECKING: from ...llmapi.llm_args import DFlashDecodingConfig + from ..pyexecutor.resource_manager import ResourceManager def compute_dflash_ctx_buffer_bytes( @@ -542,7 +544,7 @@ def prepare(self): # Update slot mapping for DFlash context buffers worker = getattr(self, "_dflash_worker", None) - if worker is not None and worker._ctx_buf_inited: + if worker is not None and worker._ctx_buf_inited and not worker._has_unified_draft_cache(): current = set(self.request_ids) evicted = {} for rid in list(worker._req_to_slot.keys()): @@ -666,12 +668,16 @@ def __init__( # graph compatible. self._ctx_buf_inited = False self._ctx_len = None + self._ctx_position_offset = None + self._managed_cache_bindings = {} + self._managed_request_ids = () # Host shadows of _ctx_len and of each request's prompt progress. self._ctx_len_host = None self._req_ctx_pos = {} # Snapshot for rolling back in-place _ctx_len updates when a forward # fails (or after warmup). See _ensure_spec_dec_state_restored. self._saved_ctx_len = None + self._saved_ctx_position_offset = None self._saved_ctx_len_host = None self._saved_req_ctx_pos = None self._ctx_len_restore_pending = False @@ -785,6 +791,116 @@ def set_draft_model(self, draft_model) -> None: "(d2t vocab mapping is not supported)." ) + def get_draft_kv_cache_manager( + self, resource_manager: Optional["ResourceManager"] + ) -> Optional[BaseResourceManager]: + if resource_manager is not None: + manager = resource_manager.get_resource_manager(ResourceManagerType.KV_CACHE_MANAGER) + if getattr(manager, "draft_layout", None) is not None: + return manager + return super().get_draft_kv_cache_manager(resource_manager) + + def _has_unified_draft_cache(self) -> bool: + return getattr(self._ctx_kv_manager, "draft_layout", None) is not None + + def _prepare_managed_history(self, request_ids: list[int]) -> list[int]: + """Restore newly resident histories without rewinding overlapping iterations.""" + updates = {self._dummy_slot: 0} + offsets = {self._dummy_slot: 0} + restore_rows = [] + real_request_ids = { + request_id + for request_id in request_ids + if request_id != ATTENTION_DP_DUMMY_REQUEST_ID + and request_id < self._graph_dummy_id_floor + and request_id not in self._ctx_kv_manager._draft_dummy_request_ids + } + for request_id in list(self._req_to_slot): + if request_id not in real_request_ids: + slot = self._req_to_slot.pop(request_id) + updates[slot] = 0 + offsets[slot] = 0 + self._req_ctx_pos.pop(request_id, None) + self._managed_cache_bindings.pop(request_id, None) + self._free_slots.append(slot) + for row, request_id in enumerate(request_ids): + if request_id in real_request_ids: + self._assign_slot(request_id) + slot = self._req_to_slot.get(request_id) + if slot is not None: + cache = self._ctx_kv_manager.kv_cache_map[request_id] + if self._managed_cache_bindings.get(request_id) is cache: + continue + history = self._ctx_kv_manager.get_draft_history(request_id) + updates[slot] = history.valid_length if history is not None else 0 + self._req_ctx_pos[request_id] = history.position if history is not None else 0 + offsets[slot] = self._req_ctx_pos[request_id] - updates[slot] + self._managed_cache_bindings[request_id] = cache + restore_rows.append(row) + self._write_ctx_len(updates) + for slot, offset in offsets.items(): + self._ctx_position_offset[slot] = offset + self._batch_to_slot[: len(request_ids)].copy_( + torch.tensor( + [self._req_to_slot.get(request_id, self._dummy_slot) for request_id in request_ids], + dtype=torch.long, + device=self._batch_to_slot.device, + ) + ) + self._managed_request_ids = tuple( + request_id for request_id in request_ids if request_id in real_request_ids + ) + return restore_rows + + def _gather_managed_context(self, request_ids: list[int], restore_rows: list[int]) -> None: + """Refresh VANILLA's dense staging; manager pages remain authoritative.""" + if self._dflash_attention_backend != "VANILLA": + return + for row in restore_rows: + request_id = request_ids[row] + slot = self._req_to_slot.get(request_id) + if slot is None: + continue + length = self._ctx_len_host[slot] + positions = torch.arange(length, device=self._ctx_k_buf.device) + pages = self._ctx_block_tables[row, positions // self._ctx_page_size].long() + offsets = positions % self._ctx_page_size + for layer_idx, pool in enumerate(self._ctx_kv_buf): + self._ctx_k_buf[slot, layer_idx, :length] = pool[pages, 0, :, offsets, :] + self._ctx_v_buf[slot, layer_idx, :length] = pool[pages, 1, :, offsets, :] + + def prepare_managed_draft_cache( + self, draft_model, spec_metadata, attn_metadata, resource_manager + ) -> None: + """Refresh managed inputs before eager execution, capture, or graph replay.""" + manager = self.get_draft_kv_cache_manager(resource_manager) + if getattr(manager, "draft_layout", None) is None: + return + self._lazy_init_ctx_buffers(draft_model, spec_metadata, attn_metadata, manager) + spec_metadata._dflash_worker = self + restore_rows = self._prepare_managed_history(spec_metadata.request_ids) + self._refresh_ctx_block_tables( + attn_metadata, attn_metadata.num_seqs, spec_metadata.request_ids + ) + self._gather_managed_context(spec_metadata.request_ids, restore_rows) + + def snapshot_managed_draft_history(self) -> DraftHistoryUpdate | None: + """Queue iteration-owned readback; the executor publishes after completion.""" + if not self._has_unified_draft_cache() or not self._managed_request_ids: + return None + slots = torch.tensor( + [self._req_to_slot[request_id] for request_id in self._managed_request_ids], + dtype=torch.long, + device=self._ctx_len.device, + ) + lengths = self._ctx_len[slots] + positions = lengths + self._ctx_position_offset[slots] + return DraftHistoryUpdate.capture( + self._ctx_kv_manager, + self._managed_request_ids, + torch.stack((lengths, positions), dim=1), + ) + def _check_ctx_arena_fits(self, capacity, num_slots, L, nkv, hd, dtype, kv_factor=2): """Fail with the arithmetic before allocating the drafter context arena. @@ -856,7 +972,7 @@ def _managed_ctx_pool(self, draft_kv_cache_manager, L, nkv, hd, dtype, kv_factor Note the page count is not the drafter's to interpret: a V2 view spans the whole interleaved pool, not this layer's slice, so only shape[1:] is - checked here and the index space is settled in _init_ctx_block_tables. + checked here and the index space is settled in _init_ctx_block_offsets. """ if draft_kv_cache_manager is None: # No separate draft KV cache. Attention DP alone does NOT disable @@ -884,10 +1000,8 @@ def _managed_ctx_pool(self, draft_kv_cache_manager, L, nkv, hd, dtype, kv_factor ) return layers - def _init_ctx_block_tables( - self, draft_kv_cache_manager, pool, num_slots, L, nkv, hd, page_size - ) -> bool: - """Size the per-iteration block table and the offset decode constant. + def _init_ctx_block_offsets(self, draft_kv_cache_manager, pool, L, nkv, hd, page_size) -> bool: + """Validate the legacy pool layout and resolve its offset decode constant. ``kv_cache_block_offsets`` entries are ``pool_block_index * num_pool_layers * kv_factor`` plus the K/V field index. What that has to @@ -923,7 +1037,13 @@ def _init_ctx_block_tables( f"rather than guessing block indices." ) return False - max_blocks = mgr.max_blocks_per_seq + logger.info( + f"DFlash: ctx block offsets pool_idx={self._ctx_pool_idx}, " + f"divisor={self._ctx_block_divisor}, view_pages={pool[0].size(0)}" + ) + return True + + def _init_ctx_block_tables(self, num_slots: int, max_blocks: int) -> None: self._ctx_block_tables = torch.zeros( (num_slots, max_blocks), dtype=torch.int32, device="cuda" ) @@ -931,14 +1051,11 @@ def _init_ctx_block_tables( 0, (num_slots + 1) * max_blocks, max_blocks, dtype=torch.int32, device="cuda" ) self._ctx_block_counts = torch.zeros(num_slots, dtype=torch.long, device="cuda") - logger.info( - f"DFlash: ctx block tables {tuple(self._ctx_block_tables.shape)}, " - f"pool_idx={self._ctx_pool_idx}, divisor={self._ctx_block_divisor}, " - f"view_pages={pool[0].size(0)}" - ) - return True + logger.info(f"DFlash: ctx block tables {tuple(self._ctx_block_tables.shape)}") - def _refresh_ctx_block_tables(self, attn_metadata, num_seqs: int) -> bool: + def _refresh_ctx_block_tables( + self, attn_metadata, num_seqs: int, request_ids: list[int] | None = None + ) -> bool: """Decode this iteration's draft block table into the persistent buffer. In place and fixed-shape so a captured graph keeps reading the same @@ -947,6 +1064,20 @@ def _refresh_ctx_block_tables(self, attn_metadata, num_seqs: int) -> bool: """ if self._ctx_block_tables is None or num_seqs <= 0: return False + if self._has_unified_draft_cache(): + table = self._ctx_kv_manager.get_draft_block_table(request_ids[:num_seqs]) + self._ctx_block_tables[:num_seqs].copy_(table, non_blocking=True) + self._ctx_block_counts[:num_seqs].copy_( + torch.tensor( + [ + self._ctx_kv_manager.kv_cache_map[rid].num_blocks + for rid in request_ids[:num_seqs] + ], + dtype=torch.long, + device=self._ctx_block_counts.device, + ) + ) + return True src = getattr(attn_metadata, "draft_kv_cache_block_offsets", None) if src is None: # The table is bound but unfillable. Leaving it at zeros would send @@ -975,8 +1106,11 @@ def _lazy_init_ctx_buffers( # Only TrtllmAttentionMetadata builds draft_kv_cache_block_offsets. # Without it the table stays all-zero and every request would read and # write pool block 0 -- silently, at acceptance 1.0. - if draft_kv_cache_manager is not None and not hasattr( - attn_metadata, "draft_kv_cache_block_offsets" + unified = getattr(draft_kv_cache_manager, "draft_layout", None) is not None + if ( + not unified + and draft_kv_cache_manager is not None + and not hasattr(attn_metadata, "draft_kv_cache_block_offsets") ): logger.warning( f"DFlash: {type(attn_metadata).__name__} carries no draft KV block offsets; " @@ -1050,12 +1184,15 @@ def _lazy_init_ctx_buffers( self._graph_dummy_id_floor = CUDA_GRAPH_DUMMY_REQUEST_ID - self.max_draft_len self._ctx_len = torch.zeros(num_slots, dtype=torch.long, device="cuda") + self._ctx_position_offset = torch.zeros_like(self._ctx_len) self._ctx_len_host = [0] * num_slots self._batch_to_slot = torch.zeros(max_batch, dtype=torch.long, device="cuda") self._free_slots = deque(range(max_batch)) self._req_to_slot = {} self._req_ctx_pos = {} + self._managed_cache_bindings = {} + self._managed_request_ids = () # checkpoint's trained block width self._resolved_block_size = getattr(draft_model, "block_size", None) or ( @@ -1097,14 +1234,28 @@ def _lazy_init_ctx_buffers( use_paged = self._dflash_attention_backend in _PAGED_ATTENTION_BACKENDS or getattr( draft_model, "_paged_ctx_cache", False ) - self._ctx_paged = use_paged - if use_paged: - # FA4 pages, but on the private arena with its own kernel. - pool = ( - None - if self._dflash_attention_backend == "FA4" - else self._managed_ctx_pool(draft_kv_cache_manager, L, nkv, hd, dtype, kv_factor) - ) + self._ctx_paged = unified or use_paged + if self._ctx_paged: + if unified: + pool = [ + draft_kv_cache_manager.get_draft_buffers(i, kv_layout="HND") for i in range(L) + ] + self._ctx_page_size = draft_kv_cache_manager.tokens_per_block + expected = (2, nkv, self._ctx_page_size, hd) + if any( + tuple(layer.shape[1:]) != expected or layer.dtype != dtype for layer in pool + ): + raise ValueError( + "Unified DSpark draft pool does not match the drafter KV layout" + ) + else: + pool = ( + None + if self._dflash_attention_backend == "FA4" + else self._managed_ctx_pool( + draft_kv_cache_manager, L, nkv, hd, dtype, kv_factor + ) + ) # The manager's page size wins when bound to it: its pool is already # carved, so the drafter adopts the geometry rather than imposing one. page_size = self._ctx_page_size if pool is None else pool[0].size(-2) @@ -1127,13 +1278,17 @@ def _lazy_init_ctx_buffers( tokens_per_block=page_size, has_context_attention=has_context_attention, ) - elif self._dflash_attention_backend == "FA4": + elif not unified and self._dflash_attention_backend == "FA4": validate_dflash_fa4_runtime(dtype=dtype, head_dim=hd) # Settle the block table before committing to the pool: it is the # last thing that can rule the pool out, and falling back after # taking the pool branch would leave no buffer allocated at all. - if pool is not None and not self._init_ctx_block_tables( - draft_kv_cache_manager, pool, num_slots, L, nkv, hd, page_size + if ( + not unified + and pool is not None + and not self._init_ctx_block_offsets( + draft_kv_cache_manager, pool, L, nkv, hd, page_size + ) ): pool = None if pool is None: @@ -1164,17 +1319,20 @@ def _lazy_init_ctx_buffers( # block table one iteration at a time, so the footprint follows # the sequences served rather than max_batch x max_seq_len. self._ctx_kv_buf = pool + self._init_ctx_block_tables(num_slots, draft_kv_cache_manager.max_blocks_per_seq) # Only a manager that publishes and matches its own blocks can # hand back a reused prefix's drafter K/V; the private arena and # an unpaired draft pool both start every request at 0. - self._ctx_reuse_addressable = bool( - getattr(draft_kv_cache_manager, "enable_joint_kv_cache_reuse", False) - ) + if not unified: + self._ctx_reuse_addressable = bool( + getattr(draft_kv_cache_manager, "enable_joint_kv_cache_reuse", False) + ) self._ctx_kv_last_page_len = torch.full( (num_slots,), page_size, dtype=torch.int32, device="cuda" ) - else: # dense private arena (VANILLA, unpaged) - self._check_ctx_arena_fits(capacity, num_slots, L, nkv, hd, dtype, kv_factor) + if not use_paged: + if not unified: + self._check_ctx_arena_fits(capacity, num_slots, L, nkv, hd, dtype, kv_factor) kv_shape = (num_slots, L, capacity, nkv, hd) self._ctx_k_buf = torch.zeros(kv_shape, dtype=dtype, device="cuda") # An MLA drafter stores one latent per token; leaving _ctx_v_buf as @@ -1243,9 +1401,14 @@ def clear(slot: int) -> None: self._free_slots.append(old_slot) if req_id not in self._req_to_slot: if not self._free_slots: + if self._has_unified_draft_cache(): + raise RuntimeError("Unified DSpark has no free request staging slot") return None slot = self._free_slots.popleft() self._req_to_slot[req_id] = slot + if self._has_unified_draft_cache(): + self._managed_cache_bindings.pop(req_id, None) + self._ctx_position_offset[slot] = 0 clear(slot) return self._req_to_slot[req_id] @@ -1286,12 +1449,13 @@ def _store_context_kv_paged( positions_i32 = positions.to(torch.int32) kv_indices, kv_indptr = self._ctx_paged_index_args() for layer_idx in range(k.size(1)): + pool = self._ctx_kv_buf[layer_idx] append_paged_kv_cache( - append_key=k[:, layer_idx].contiguous(), - append_value=v[:, layer_idx].contiguous(), + append_key=k[:, layer_idx].to(pool.dtype).contiguous(), + append_value=v[:, layer_idx].to(pool.dtype).contiguous(), batch_indices=rows_i32, positions=positions_i32, - paged_kv_cache=self._ctx_kv_buf[layer_idx], + paged_kv_cache=pool, kv_indices=kv_indices, kv_indptr=kv_indptr, kv_last_page_len=self._ctx_kv_last_page_len, @@ -1364,6 +1528,8 @@ def _restore_ctx_len_host(self) -> None: self._ctx_len_host = list(self._saved_ctx_len_host) if self._saved_req_ctx_pos is not None: self._req_ctx_pos = dict(self._saved_req_ctx_pos) + if self._has_unified_draft_cache() and self._saved_ctx_position_offset is not None: + self._ctx_position_offset.copy_(self._saved_ctx_position_offset) def _store_prefill_context( self, @@ -1448,6 +1614,8 @@ def _store_prefill_context( # paired pool: precompute_context_kv is per-token in (hidden, p). cur = first_pos if cur + slen > cap: + if self._has_unified_draft_cache(): + raise ValueError("Unified DSpark prompt history exceeds the drafter context") # Request-level, like the no-free-slots path above: truncating # would silently draft from a stale prefix, but killing the # forward would take every other in-flight request with it. @@ -1499,10 +1667,13 @@ def _store_prefill_context( torch.full((actual,), row, dtype=torch.long, device="cuda"), local_pos, ) - else: # VANILLA DFlash backend (FlashAttention) + if self._ctx_k_buf is not None: self._ctx_k_buf[slot, :, cur:end] = chunk_k.permute(1, 0, 2, 3) if chunk_v is not None: self._ctx_v_buf[slot, :, cur:end] = chunk_v.permute(1, 0, 2, 3) + if self._has_unified_draft_cache(): + self._ctx_position_offset[slot] = self._req_ctx_pos[req_id] - end + self._managed_cache_bindings[req_id] = self._ctx_kv_manager.kv_cache_map[req_id] offset += slen self._write_ctx_len(ctx_len_updates) @@ -1558,16 +1729,16 @@ def _forward_impl( draft_model, ) - # Lazy init buffers and attach worker reference for prepare() draft_kv_cache_manager = self.get_draft_kv_cache_manager(resource_manager) - self._lazy_init_ctx_buffers( - draft_model, spec_metadata, attn_metadata, draft_kv_cache_manager - ) - spec_metadata._dflash_worker = self - # Before any store: prefill and decode both address pages through it. - # Returning False here means an empty batch -- the missing-offsets case - # raises inside, with the metadata type in the message. - self._refresh_ctx_block_tables(attn_metadata, batch_size) + if getattr(draft_kv_cache_manager, "draft_layout", None) is not None: + if not self._ctx_buf_inited or self._ctx_kv_manager is not draft_kv_cache_manager: + raise RuntimeError("Unified DSpark draft inputs must be prepared before forward") + else: + self._lazy_init_ctx_buffers( + draft_model, spec_metadata, attn_metadata, draft_kv_cache_manager + ) + spec_metadata._dflash_worker = self + self._refresh_ctx_block_tables(attn_metadata, batch_size, spec_metadata.request_ids) # Save context lengths so both warmup and a failed forward can roll # back the in-place _ctx_len updates made during drafting. @@ -1580,6 +1751,8 @@ def _forward_impl( # during capture aborts the graph itself, and captured ops do not # mutate _ctx_len until replay. self._saved_ctx_len = self._ctx_len.clone() + if self._has_unified_draft_cache(): + self._saved_ctx_position_offset = self._ctx_position_offset.clone() self._saved_ctx_len_host = list(self._ctx_len_host) self._saved_req_ctx_pos = dict(self._req_ctx_pos) self._ctx_len_restore_pending = True @@ -1661,7 +1834,9 @@ def _forward_impl( ) if num_gens > 0: - with self.draft_kv_cache_context(attn_metadata, draft_kv_cache_manager): + with self.draft_kv_cache_context( + attn_metadata, None if self._has_unified_draft_cache() else draft_kv_cache_manager + ): hidden_states_out = draft_model.dflash_forward( noise_embedding=inputs["noise_embedding"], query_positions=inputs["query_positions"], @@ -1983,17 +2158,25 @@ def prepare_1st_drafter_inputs( bonus = gen_accepted_tokens.gather(1, bonus_idx).squeeze(1).long() ctx_len_gen = self._ctx_len[slots] + ctx_position_gen = ctx_len_gen + if self._has_unified_draft_cache(): + allocated = dflash_allocated_ctx_limit( + self._ctx_block_counts[gen_rows_out], self._ctx_page_size, block_size + ).clamp(max=self._max_ctx) + torch._assert_async( + torch.all(ctx_len_gen + gen_num_accepted <= allocated), + "Unified DSpark accepted history exceeds the drafter context", + ) + ctx_position_gen = ctx_len_gen + self._ctx_position_offset[slots] j_block = torch.arange(query_tokens_per_req, dtype=torch.long, device="cuda") offsets_kp1 = torch.arange(K_plus_1, dtype=torch.long, device="cuda") - # _ctx_len is clamped to _max_ctx only AFTER this step's accepted - # tokens are folded in (see the update below), so the running length - # used here has to be clamped on its own -- otherwise a request that - # already sits at the ceiling indexes num_accepted positions past - # any position the sequence can legitimately reach. - ctx_len_now = (ctx_len_gen + gen_num_accepted.long()).clamp_(max=self._max_ctx) - query_position_ids = ctx_len_now.unsqueeze(1) + j_block.unsqueeze(0) - ctx_position_ids = ctx_len_gen.unsqueeze(1) + offsets_kp1.unsqueeze(0) + query_position = ctx_position_gen + gen_num_accepted.long() + if not self._has_unified_draft_cache(): + # Legacy warmup can advance beyond the served context ceiling. + query_position = query_position.clamp(max=self._max_ctx) + query_position_ids = query_position.unsqueeze(1) + j_block.unsqueeze(0) + ctx_position_ids = ctx_position_gen.unsqueeze(1) + offsets_kp1.unsqueeze(0) # Go through embed_tokens.forward (NOT .weight[...]) so TP-sharded # vocabs mask out ranks that don't own the token id and all-reduce. @@ -2056,7 +2239,7 @@ def prepare_1st_drafter_inputs( else: rows_long = slot_long self._store_context_kv_paged(k_new, v_new, rows_long, col_long) - else: # VANILLA DFlash backend (FlashAttention) + if self._ctx_k_buf is not None: self._ctx_k_buf[slot_long, :, col_long] = k_new if v_new is not None: self._ctx_v_buf[slot_long, :, col_long] = v_new @@ -2099,7 +2282,11 @@ def prepare_1st_drafter_inputs( "ctx_v_cache": self._ctx_v_buf, # Slots index the private arena; the manager's block table is keyed # by batch position, so the drafter reads its own rows there. - "ctx_cache_batch_idx": gen_rows_out if self._ctx_block_tables is not None else slots, + "ctx_cache_batch_idx": ( + gen_rows_out + if self._ctx_block_tables is not None and self._ctx_k_buf is None + else slots + ), # Anchor token per gen request (block slot 0): last accepted # token. The dspark Markov chain conditions its first step on it. "first_prev_tokens": bonus, diff --git a/tensorrt_llm/_torch/speculative/dspark.py b/tensorrt_llm/_torch/speculative/dspark.py index b8dff7490d4c..0c71bd1e1b81 100644 --- a/tensorrt_llm/_torch/speculative/dspark.py +++ b/tensorrt_llm/_torch/speculative/dspark.py @@ -24,12 +24,16 @@ from typing import TYPE_CHECKING, List, Optional import torch +import triton +import triton.language as tl from tensorrt_llm._utils import prefer_pinned from tensorrt_llm.logger import logger from tensorrt_llm.mapping import Mapping +from ..pyexecutor.kv_cache.standalone_draft_cache import DraftHistoryUpdate, StandaloneDraftHistory from ..pyexecutor.llm_request import ATTENTION_DP_DUMMY_REQUEST_ID +from ..pyexecutor.resource_manager import ResourceManagerType from .dflash import DFlashWorker, dflash_draft_slot_ids from .interface import SpecMetadata, SpecWorkerBase @@ -37,6 +41,68 @@ from ...llmapi.llm_args import DSparkDecodingConfig +@triton.jit +def _store_managed_window_kernel( + windows, + pool, + slots, + lengths, + positions, + write_starts, + block_tables, + real_rows, + capacities, + window_slot_stride: tl.constexpr, + window_stage_stride: tl.constexpr, + window_token_stride: tl.constexpr, + window_head_stride: tl.constexpr, + pool_page_stride: tl.constexpr, + pool_token_stride: tl.constexpr, + pool_head_stride: tl.constexpr, + table_stride: tl.constexpr, + page_size: tl.constexpr, + window_size: tl.constexpr, + head_dim: tl.constexpr, + stage: tl.constexpr, + BLOCK: tl.constexpr, +): + row = tl.program_id(0) + indices = tl.program_id(1) * BLOCK + tl.arange(0, BLOCK) + token = indices // head_dim + head = indices % head_dim + slot = tl.load(slots + row) + length = tl.load(lengths + slot) + position = tl.load(positions + slot) - length + token + write_start = tl.load(write_starts + row) + real = tl.load(real_rows + row) + capacity = tl.load(capacities + row) + valid = real & (token < length) & (token < window_size) + valid = valid & (position >= write_start) & (position >= 0) & (position < capacity) + page = tl.load( + block_tables + row * table_stride + position // page_size, + mask=valid, + other=-1, + ) + valid = valid & (page >= 0) + values = tl.load( + windows + + slot * window_slot_stride + + stage * window_stage_stride + + ((position + 1) % window_size) * window_token_stride + + head * window_head_stride, + mask=valid, + other=0, + ) + tl.store( + pool + + page * pool_page_stride + + (position % page_size) * pool_token_stride + + head * pool_head_stride, + values, + mask=valid, + ) + + def _dspark_position_ceiling(max_ctx: int, block_size: int, max_draft_len: int) -> int: """Return the number of RoPE entries needed by the DSv4 block drafter. @@ -130,15 +196,7 @@ def prepare(self): # ``DFlashSpecMetadata.prepare`` (dflash.py:96-113). worker = getattr(self, "_dspark_worker", None) if worker is not None and worker._win_inited: - current = set(self.request_ids) - for rid in list(worker._req_to_slot.keys()): - if rid not in current: - slot = worker._req_to_slot.pop(rid) - worker._ctx_len[slot] = 0 - worker._valid_len[slot] = 0 - worker._position_initialized[slot] = False - worker._kv_windows[slot].zero_() - worker._free_slots.append(slot) + worker._release_inactive_slots(self.request_ids) # Assign a persistent rolling-window slot to every real generation # request that never ran a context/seed forward on this worker. In # disaggregated serving the prompt is prefilled (and the window @@ -228,9 +286,10 @@ class DSv4DSparkWorker(SpecWorkerBase): window, refines the per-position logits with the Markov head, and predicts a per-position acceptance confidence used to truncate the proposed prefix. - Unlike DFlash, the draft does NOT use the paged KV cache or mask-token - cross-attention: its attention K/V come from the worker-owned rolling window - of projected captured context (one ``main_kv`` per decode step, per stage). + Unlike DFlash, attention reads a rolling window of projected captured + context (one ``main_kv`` per decode step, per stage). With unified V2, + manager-owned pages preserve that history and the rolling tensors are kernel + staging; legacy aggregate execution keeps the windows in the worker. Acceptance of the previous block goes through the unified :meth:`SpecWorkerBase.sample_and_accept_draft_tokens` (strict target-verify, or rejection sampling for a non-greedy batch), so greedy parity with no-spec @@ -273,6 +332,13 @@ def __init__( self._valid_len: Optional[torch.Tensor] = None # [max_batch] written window entries self._position_initialized: Optional[torch.Tensor] = None # [max_batch] bool self._win = 0 + self._draft_kv_manager = None + self._draft_kv_buffers = () + self._draft_block_tables = None + self._draft_real_rows = None + self._draft_capacities = None + self._managed_residency = {} + self._prepared_managed_request_ids = () # Set in _lazy_init from the RoPE table the drafter will build; None # leaves positions unbounded (direct construction in tests). self._position_cap: Optional[int] = None @@ -430,8 +496,11 @@ def _lazy_init(self, draft_model, spec_metadata, attn_metadata=None) -> None: def _assign_slot(self, req_id: int, reset: bool) -> int: """Get (or refresh) the slot for a request; reset clears its window.""" + if self._draft_kv_manager is not None and not self._is_managed_request(req_id): + return self._scratch_slot if reset and req_id in self._req_to_slot: old = self._req_to_slot.pop(req_id) + self._managed_residency.pop(req_id, None) self._ctx_len[old] = 0 self._valid_len[old] = 0 self._position_initialized[old] = False @@ -451,6 +520,199 @@ def _assign_slot(self, req_id: int, reset: bool) -> int: self._kv_windows[slot].zero_() return self._req_to_slot[req_id] + def _release_inactive_slots(self, request_ids: list[int]) -> None: + current = set(request_ids) + for request_id in list(self._req_to_slot): + if request_id not in current: + slot = self._req_to_slot.pop(request_id) + self._managed_residency.pop(request_id, None) + self._ctx_len[slot] = 0 + self._valid_len[slot] = 0 + self._position_initialized[slot] = False + self._kv_windows[slot].zero_() + self._free_slots.append(slot) + + def _is_managed_request(self, request_id: int) -> bool: + return ( + request_id != ATTENTION_DP_DUMMY_REQUEST_ID + and request_id < self._graph_dummy_id_floor + and request_id not in self._draft_kv_manager._draft_dummy_request_ids + ) + + def _bind_managed_history(self, resource_manager) -> None: + manager = ( + resource_manager.get_resource_manager(ResourceManagerType.KV_CACHE_MANAGER) + if resource_manager is not None + else None + ) + layout = getattr(manager, "draft_layout", None) + if layout is None: + manager = None + if manager is self._draft_kv_manager: + return + buffers = () + if manager is not None: + if ( + layout.window_size != self._win + or layout.kv_factor != 1 + or layout.num_layers != self._kv_windows.shape[1] + or layout.head_dim != self._kv_windows.shape[-1] + ): + raise ValueError("Embedded DSpark draft cache does not match its rolling window") + buffers = tuple(manager.get_draft_buffers(stage) for stage in range(layout.num_layers)) + self._draft_kv_manager = manager + self._draft_kv_buffers = buffers + self._managed_residency.clear() + if manager is not None: + self._draft_block_tables = torch.zeros( + (self._batch_to_slot.shape[0], manager.max_blocks_per_seq), + dtype=torch.int32, + device=self._kv_windows.device, + ) + self._draft_real_rows = torch.zeros( + self._batch_to_slot.shape[0], dtype=torch.bool, device=self._kv_windows.device + ) + self._draft_capacities = torch.zeros( + self._batch_to_slot.shape[0], dtype=torch.long, device=self._kv_windows.device + ) + + def prepare_managed_draft_cache( + self, draft_model, spec_metadata, attn_metadata, resource_manager + ) -> None: + """Refresh persistent draft inputs before eager execution or graph replay.""" + self._lazy_init(draft_model, spec_metadata, attn_metadata) + spec_metadata._dspark_worker = self + self._bind_managed_history(resource_manager) + if self._draft_kv_manager is not None: + self._prepare_managed_history(spec_metadata.request_ids, attn_metadata.num_contexts) + + def snapshot_managed_draft_history(self): + """Copy this iteration's history before another batch can reuse its slots.""" + if self._draft_kv_manager is None: + return None + request_ids = tuple(self._prepared_managed_request_ids) + if not request_ids: + return None + slots = torch.tensor( + [self._req_to_slot[request_id] for request_id in request_ids], + dtype=torch.long, + device=self._ctx_len.device, + ) + values = torch.stack((self._valid_len[slots], self._ctx_len[slots]), dim=1) + return DraftHistoryUpdate.capture(self._draft_kv_manager, request_ids, values) + + def _managed_window_indices(self, row: int, position: int, length: int): + """Map logical token p to a local page and DSpark's frame (p + 1) % W.""" + positions = torch.arange(position - length, position, device=self._kv_windows.device) + page_size = self._draft_kv_manager.tokens_per_block + pages = self._draft_block_tables[row, positions // page_size].long() + if torch.any(pages < 0).item(): + raise ValueError("Embedded DSpark committed window contains an unallocated page") + return pages, positions % page_size, (positions + 1) % self._win + + def _prepare_managed_history(self, request_ids: list[int], num_contexts: int) -> None: + """Restore newly resident histories without rewinding overlapping iterations.""" + # Startup probes and graph/ADP padding never publish synthetic history. + real_rows = [ + row + for row, request_id in enumerate(request_ids) + if self._is_managed_request(request_id) + ] + real_ids = [request_ids[row] for row in real_rows] + self._release_inactive_slots(real_ids) + validation_histories = [] + for request_id in real_ids: + cache = self._draft_kv_manager.kv_cache_map[request_id] + history = self._draft_kv_manager.get_draft_history(request_id) + resident = self._managed_residency.get(request_id) + if ( + resident is not None + and resident[0] is cache + and resident[1] == self._req_to_slot.get(request_id) + ): + # Context completion can retire pages before its history + # readback is published. Its retained frames are already in + # the resident device window on the execution stream. + position = max(history.position if history is not None else 0, cache.history_length) + history = StandaloneDraftHistory(min(self._win, position), position) + validation_histories.append(history or StandaloneDraftHistory(0, 0)) + table = self._draft_kv_manager.get_draft_block_table(real_ids, validation_histories) + tables_host = torch.zeros_like(self._draft_block_tables, device="cpu") + tables_host[real_rows] = table + self._draft_block_tables.copy_(tables_host, non_blocking=True) + real_host = torch.zeros_like(self._draft_real_rows, device="cpu") + real_host[real_rows] = True + self._draft_real_rows.copy_(real_host, non_blocking=True) + capacities_host = torch.zeros_like(self._draft_capacities, device="cpu") + capacities_host[real_rows] = torch.tensor( + [self._draft_kv_manager.kv_cache_map[rid].capacity for rid in real_ids], + dtype=torch.long, + ) + self._draft_capacities.copy_(capacities_host, non_blocking=True) + self._kv_windows[self._scratch_slot].zero_() + self._ctx_len[self._scratch_slot] = 0 + self._valid_len[self._scratch_slot] = 0 + self._position_initialized[self._scratch_slot] = False + batch_slots = [self._scratch_slot] * len(request_ids) + for row in real_rows: + request_id = request_ids[row] + cache = self._draft_kv_manager.kv_cache_map[request_id] + slot = self._assign_slot(request_id, reset=False) + batch_slots[row] = slot + resident = self._managed_residency.get(request_id) + if resident is not None and resident[0] is cache and resident[1] == slot: + continue + history = self._draft_kv_manager.get_draft_history(request_id) + if row >= num_contexts and history is None: + raise ValueError( + f"Embedded DSpark generation request {request_id} has no committed draft history" + ) + self._kv_windows[slot].zero_() + self._ctx_len[slot] = history.position if history is not None else 0 + self._valid_len[slot] = history.valid_length if history is not None else 0 + self._position_initialized[slot] = history is not None + if history is not None: + pages, offsets, frames = self._managed_window_indices( + row, history.position, history.valid_length + ) + for stage, pool in enumerate(self._draft_kv_buffers): + self._kv_windows[slot, stage, frames] = pool[pages, 0, 0, offsets] + self._managed_residency[request_id] = (cache, slot) + self._batch_to_slot.fill_(self._scratch_slot) + self._batch_to_slot[: len(request_ids)].copy_( + torch.tensor(batch_slots, dtype=torch.long, device=self._batch_to_slot.device) + ) + self._prepared_managed_request_ids = tuple(real_ids) + + def _write_managed_history( + self, batch_size: int, write_starts: torch.Tensor, *, batch_offset: int = 0 + ) -> None: + """Store newly appended frames without rewriting shared committed prefix pages.""" + head_dim = self._kv_windows.shape[-1] + rows = slice(batch_offset, batch_offset + batch_size) + for stage, pool in enumerate(self._draft_kv_buffers): + _store_managed_window_kernel[(batch_size, triton.cdiv(self._win * head_dim, 128))]( + self._kv_windows, + pool, + self._batch_to_slot[rows], + self._valid_len, + self._ctx_len, + write_starts, + self._draft_block_tables[rows], + self._draft_real_rows[rows], + self._draft_capacities[rows], + *self._kv_windows.stride(), + pool.stride(0), + pool.stride(3), + pool.stride(4), + self._draft_block_tables.stride(0), + self._draft_kv_manager.tokens_per_block, + self._win, + head_dim, + stage, + BLOCK=128, + ) + def _seed_context_windows( self, draft_model, @@ -477,20 +739,38 @@ def _seed_context_windows( req_id = spec_metadata.request_ids[i] first_position = int(chunk_positions[0].item()) - slot = self._assign_slot(req_id, reset=first_position == 0) - self._ctx_len[slot] = chunk_positions[-1] + 1 + slot = self._assign_slot( + req_id, reset=first_position == 0 and self._draft_kv_manager is None + ) self._position_initialized[slot] = True if captured is not None: - self._valid_len[slot] = torch.clamp( - self._valid_len[slot] + chunk_len, max=self._win - ) keep = min(self._win, chunk_len) + if self._draft_kv_manager is not None and self._draft_kv_manager.enable_block_reuse: + # Reusable prefixes can end before this chunk's final window. + # Persist earlier frames before the ring overwrites them. + prefix_end = chunk_len - keep + for begin in range(0, prefix_end, self._win): + end = min(begin + self._win, prefix_end) + self._ctx_len[slot] = chunk_positions[end - 1] + 1 + self._valid_len[slot] = torch.clamp( + self._valid_len[slot] + end - begin, max=self._win + ) + draft_model.write_context_windows( + captured[context_offset + begin : context_offset + end], + chunk_positions[begin:end] + 1, + self._kv_windows[slot], + ) + self._write_managed_history( + 1, chunk_positions[begin : begin + 1], batch_offset=i + ) + self._valid_len[slot] = torch.clamp(self._valid_len[slot] + keep, max=self._win) hidden = captured[context_offset + chunk_len - keep : context_offset + chunk_len] # A prompt token at absolute position p is stored in frame p+1, # matching the generation path's start_pos convention. window_positions = chunk_positions[-keep:] + 1 draft_model.write_context_windows(hidden, window_positions, self._kv_windows[slot]) + self._ctx_len[slot] = chunk_positions[-1] + 1 context_offset += chunk_len def _advance_generation_state( @@ -671,10 +951,18 @@ def _forward_impl( raw_logits = logits K = self.max_draft_len - self._lazy_init(draft_model, spec_metadata, attn_metadata) - # Backref so DSparkSpecMetadata.prepare() can maintain the host slot map - # and mirror it into _batch_to_slot for the CUDA-graph-safe gen path. - spec_metadata._dspark_worker = self + if self._draft_kv_manager is None: + if resource_manager is not None: + manager = resource_manager.get_resource_manager( + ResourceManagerType.KV_CACHE_MANAGER + ) + if getattr(manager, "draft_layout", None) is not None: + raise RuntimeError( + "Unified DSpark KV cache must be prepared before worker forward" + ) + self._lazy_init(draft_model, spec_metadata, attn_metadata) + spec_metadata._dspark_worker = self + self._execute_guided_decoder_if_present(logits) # Target-verify acceptance via the unified SpecWorkerBase entry: it @@ -705,6 +993,11 @@ def _forward_impl( saved_position_initialized = self._position_initialized.clone() saved_windows = self._kv_windows.clone() + if self._draft_kv_manager is not None: + # Gather before either append path advances the positions. This runs + # inside graph replay so overlapping iterations use device progress. + write_starts = self._ctx_len[self._batch_to_slot[:batch_size]] + # Assign / reset window slots for context (prefill) requests and seed each # request's rolling KV window from its prompt's captured context, so the # first generation step drafts against real context instead of an all-zero @@ -813,12 +1106,14 @@ def _forward_impl( num_accepted_tokens, ) + if self._draft_kv_manager is not None: + self._write_managed_history(batch_size, write_starts) + if is_warmup: self._ctx_len.copy_(saved_ctx_len) self._valid_len.copy_(saved_valid_len) self._position_initialized.copy_(saved_position_initialized) self._kv_windows.copy_(saved_windows) - return { "logits": raw_logits, "new_tokens": accepted_tokens, diff --git a/tensorrt_llm/_torch/speculative/interface.py b/tensorrt_llm/_torch/speculative/interface.py index 37992e2f8835..0f99a6f53c7a 100644 --- a/tensorrt_llm/_torch/speculative/interface.py +++ b/tensorrt_llm/_torch/speculative/interface.py @@ -36,7 +36,9 @@ if TYPE_CHECKING: from ..pyexecutor.guided_decoder import CapturableGuidedDecoder + from ..pyexecutor.kv_cache.standalone_draft_cache import DraftHistoryUpdate from ..pyexecutor.llm_request import LlmRequest + from ..pyexecutor.resource_manager import ResourceManager if IS_FLASHINFER_AVAILABLE: import flashinfer @@ -1602,6 +1604,19 @@ def register_auxiliary_state_handler( for registered in self._auxiliary_state_handlers): self._auxiliary_state_handlers.append(handler) + def prepare_managed_draft_cache( + self, + draft_model: nn.Module | None, + spec_metadata: "SpecMetadata", + attn_metadata: AttentionMetadata, + resource_manager: "ResourceManager", + ) -> None: + """Stage manager-owned draft history before eager execution or replay.""" + + def snapshot_managed_draft_history(self) -> Optional["DraftHistoryUpdate"]: + """Snapshot this execution's history for publication at completion.""" + return None + def commit_auxiliary_speculative_states( self, num_accepted_tokens: torch.Tensor, diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi index 91ccd9142528..d0a47da61c58 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi @@ -194,6 +194,7 @@ class AttentionLayerConfig: buffers: list[BufferConfig] sliding_window_size: int | None = None num_sink_tokens: int | None = None + cache_domain: str = "target" @property def window_size(self) -> int | None: ... diff --git a/tests/unittest/disaggregated/test_extractor.py b/tests/unittest/disaggregated/test_extractor.py index ec0a638524cb..c5e890732c53 100644 --- a/tests/unittest/disaggregated/test_extractor.py +++ b/tests/unittest/disaggregated/test_extractor.py @@ -528,6 +528,7 @@ def _make_fake_v2_manager(attrs, role_mapper_kinds, *, num_pools=1, slot_bytes_l impl=impl, pp_layers=[0, 1], num_kv_heads_per_layer=[1, 1], + _is_standalone_draft_layer=lambda _layer: False, get_disagg_role_mapper_kinds=lambda: role_mapper_kinds, get_disagg_role_layouts=lambda: {}, )