From 1548780d8b82c1d5fb3c8052d2127179a88e18b7 Mon Sep 17 00:00:00 2001 From: Allison Lim Date: Fri, 18 Sep 2026 22:39:36 -0700 Subject: [PATCH 01/13] [None][feat] Unify standalone DSpark KV ownership and transfer Own target and standalone draft layers in one Python KVCacheManagerV2 with a common budget and separate lifecycle domains. Bind dense and paged draft consumers to managed history and restore prompt KV, lengths, and positions through the existing disaggregation transport and auxiliary buffers. Handle late receiver auxiliary dispatch, recycled request slots, scratch capacity, rewind, and cleanup. Reject unsupported execution combinations while retaining the existing MTP/Eagle, Gemma, V1 aggregate, and embedded DeepSeek paths. Validation: 446 focused tests passed; 32 acceptance-runner tests passed. The fixed Qwen trace preserves all 77 prompt positions byte-for-byte and enters the first draft forward with length 78 and query positions 78-84. Broader regressions have three documented pre-existing fixture failures. Kimi end-to-end and TRTLLM model smoke validation remain outstanding; matched aggregate/disaggregate AL measurements are left to the user. Signed-off-by: Allison Lim --- .../_torch/disaggregation/native/auxiliary.py | 134 ++++++-- .../_torch/disaggregation/native/transfer.py | 25 +- .../_torch/disaggregation/transceiver.py | 95 +++++- tensorrt_llm/_torch/pyexecutor/_util.py | 194 +++++++++-- .../kv_cache/kv_cache_manager_v2.py | 304 +++++++++++++++++- .../kv_cache/mamba_cache_manager.py | 6 +- .../kv_cache/standalone_draft_cache.py | 65 ++++ tensorrt_llm/_torch/speculative/dflash.py | 233 +++++++++++++- .../runtime/kv_cache_manager_v2/__init__.pyi | 1 + .../runtime/kv_cache_manager_v2/_config.py | 6 + .../_life_cycle_registry.py | 12 +- .../kv_cache_manager_v2/_storage/_config.py | 21 +- 12 files changed, 994 insertions(+), 102 deletions(-) create mode 100644 tensorrt_llm/_torch/pyexecutor/kv_cache/standalone_draft_cache.py diff --git a/tensorrt_llm/_torch/disaggregation/native/auxiliary.py b/tensorrt_llm/_torch/disaggregation/native/auxiliary.py index 3620b9adb8bf..633e549c24da 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,62 @@ def build_aux_transfer_layout( AuxSlot = namedtuple("AuxSlot", ["id", "buffer"]) +_DRAFT_HISTORY_VERSION = 1 +_DRAFT_HISTORY_FIELDS = 8 +_DRAFT_DTYPE_CODES = {"torch.float16": 1, "torch.bfloat16": 2} +_DRAFT_BACKEND_CODES = {"VANILLA": 1, "TRTLLM": 2} + + +def _encode_draft_history(history: dict[str, Any]) -> list[int]: + """Encode committed history and rank-local storage identity for the wire.""" + if not isinstance(history, dict): + raise ValueError("Standalone draft transfer requires draft history metadata") + layout = history.get("layout") + if not isinstance(layout, dict): + raise ValueError("Standalone draft transfer requires a storage layout") + integer_values = [ + history.get("valid_length"), + history.get("position"), + layout.get("num_layers"), + layout.get("num_kv_heads"), + layout.get("head_dim"), + ] + if any(type(value) is not int for value in integer_values): + raise ValueError( + "Standalone draft transfer metadata requires integer lengths and dimensions" + ) + valid_length, position, num_layers, num_kv_heads, head_dim = integer_values + if valid_length < 0 or position < valid_length: + raise ValueError("Invalid standalone draft history length or position") + if min(num_layers, num_kv_heads, head_dim) <= 0: + raise ValueError("Standalone draft transfer dimensions must be positive") + dtype_code = _DRAFT_DTYPE_CODES.get(layout.get("dtype")) + backend_code = _DRAFT_BACKEND_CODES.get(layout.get("attention_backend")) + if dtype_code is None or backend_code is None: + raise ValueError("Unsupported standalone draft transfer dtype or attention backend") + return [_DRAFT_HISTORY_VERSION, *integer_values, dtype_code, backend_code] + + +def _decode_draft_history(values: list[int]) -> dict[str, Any]: + if len(values) != _DRAFT_HISTORY_FIELDS or 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 = values + dtypes = {code: name for name, code in _DRAFT_DTYPE_CODES.items()} + backends = {code: name for name, code in _DRAFT_BACKEND_CODES.items()} + history = { + "valid_length": valid_length, + "position": position, + "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), + }, + } + _encode_draft_history(history) + return history + class AuxBufferBase(ABC): """ @@ -166,8 +225,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 +262,33 @@ 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 ) + # This participates in the existing auxiliary memory registration and + # transfer. Version zero denotes an unfilled or newly allocated 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]: + """Read transferred history, rejecting missing or unsupported metadata.""" + 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 b9c951f3fc82..e5a9a5f18f96 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -1678,7 +1678,8 @@ 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 atomically saves peer info and snapshots tasks against + # send() and send_aux(), including context-first auxiliary submission. info: RecvReqInfo = RecvReqInfo.from_bytes(message[1]) with self._sessions_lock: session = self._get_session(info.unique_rid) @@ -1695,6 +1696,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,)) ) @@ -1945,7 +1948,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 @@ -3049,7 +3054,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 @@ -3507,6 +3514,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) ) @@ -3522,6 +3531,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? @@ -3779,7 +3794,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))) @@ -3788,6 +3806,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/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index d91927568326..bed43106e265 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -523,10 +523,86 @@ 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 params is not None and params.schedule_style == DisaggScheduleStyle.GENERATION_FIRST + manager = getattr(self, "_kv_cache_manager", None) + return getattr(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 + if params is not None and params.schedule_style == DisaggScheduleStyle.GENERATION_FIRST: + raise ValueError( + "Standalone DSpark draft-state transfer requires context_first scheduling; " + "generation_first is not yet supported for standalone draft history." + ) + if self.pipeline_transfer_enabled: + raise ValueError( + "Standalone DSpark draft-state transfer does not support pipelined transfer." + ) + + @staticmethod + def _validate_draft_history_range(req: LlmRequest, history: dict) -> None: + # The shared transfer extent covers the complete prompt. Draft noise KV is scratch, + # never valid history. A partial history needs per-group extents before it can be sent. + if ( + not isinstance(history, dict) + or type(history.get("valid_length")) is not int + or type(history.get("position")) is not int + or not isinstance(history.get("layout"), dict) + or history["valid_length"] != req.prompt_len + or history["position"] != req.prompt_len + ): + raise ValueError( + "Standalone DSpark transfer requires valid draft history and sequence position " + f"covering the complete prompt ({req.prompt_len} tokens)." + ) + + def _pack_draft_history(self, req: LlmRequest) -> Optional[dict]: + manager = getattr(self, "_kv_cache_manager", None) + if getattr(manager, "draft_layout", None) is None: + return None + 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 + return history + + def _received_draft_history(self, req: LlmRequest) -> Optional[dict]: + self._validate_draft_transfer(req) + 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 None + if not has_draft: + raise ValueError( + "Received standalone DSpark draft history without a manager-owned draft cache." + ) + self._validate_draft_history_range(req, history) + if history["layout"] != self._kv_cache_manager.draft_layout.transfer_identity(): + raise ValueError( + "Standalone DSpark draft cache layouts do not match between prefill and " + "generation workers. Matching draft dtype, backend, and per-rank geometry " + "are required." + ) + return history + + def _restore_draft_history(self, req: LlmRequest) -> None: + history = self._received_draft_history(req) + if history is not None: + # K/V is already in this request's local pages. Only portable validity/position + # metadata crosses the wire; the manager retains the receiver's request/page map. + self._kv_cache_manager.restore_draft_history(req.py_request_id, history) def _validate_bridge_req(self, req: LlmRequest, synchronous: bool = False) -> bool: if not getattr(self, "_fp4_mla_bridge_enabled", False): @@ -811,6 +887,12 @@ 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 already carries tokens and usage in the context response. The + # existing registered auxiliary transfer carries only the additional draft state. + 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 @@ -819,7 +901,7 @@ def _apply_aux(self, session, req: LlmRequest): req.context_phase_params = ContextPhaseParams( first_gen_tokens=first_gen_tokens, req_id=req.py_request_id, - opaque_state=b"", + opaque_state=None, draft_tokens=draft_tokens, ctx_dp_rank=0, disagg_info_endpoint="", @@ -939,6 +1021,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()) @@ -987,6 +1070,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: @@ -1016,6 +1100,7 @@ def request_and_receive_sync(self, 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) req.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE else: req.state = LlmRequestState.DISAGG_TRANS_ERROR @@ -1065,6 +1150,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) @@ -1270,6 +1356,7 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]) -> GenT if self._need_aux_transfer(req): self._apply_aux(session, req) self._assert_disagg_history_declared(req) + self._restore_draft_history(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 f6e2ce17eabc..598ef17e0c14 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -67,6 +67,7 @@ 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 @@ -781,6 +782,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). @@ -933,6 +935,15 @@ def _get_kv_size_per_token(self, model_config, kv_cache_config, use_separate_draft_kv_cache=use_separate_draft_kv_cache) + if self._uses_unified_standalone_draft_cache(): + draft_layout = self._get_standalone_draft_layout() + # The unified allocator reserves an envelope for draft capture and + # noise in every attention group. Include its fixed cost as well + # as the distinct BF16 draft layers in the common byte budget. + total += CacheCost( + slope=draft_layout.bytes_per_token, + intercept=(total.slope + draft_layout.bytes_per_token) * + draft_layout.extra_tokens * self._max_batch_size) if self._is_encoder_decoder(): total += CacheCost.from_raw(self._get_cross_kv_size_per_token()) draft_cost = self._get_draft_cache_cost( @@ -1176,6 +1187,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( @@ -1220,6 +1234,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 @@ -1589,12 +1605,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 @@ -1624,6 +1644,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: @@ -1651,6 +1672,103 @@ 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 _uses_unified_standalone_draft_cache(self) -> bool: + return self._is_standalone_dspark() and self._is_kv_cache_manager_v2 + + def _validate_standalone_draft_cache(self) -> None: + """Reject unsupported standalone state ownership before profiling.""" + if not self._is_standalone_dspark(): + return + if not self._is_kv_cache_manager_v2: + if self._kv_cache_config.use_kv_cache_manager_v2 is True: + raise ValueError( + "Standalone 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( + "Standalone DSpark disaggregation requires " + "kv_cache_config.use_kv_cache_manager_v2=True and " + "TLLM_KV_CACHE_MANAGER_V2_BACKEND=python on both workers.") + return + from tensorrt_llm.runtime.kv_cache_manager_v2 import BACKEND + + if BACKEND != "python": + raise ValueError("Unified standalone DSpark KV cache requires " + "TLLM_KV_CACHE_MANAGER_V2_BACKEND=python.") + if (self._speculative_config.draft_len_schedule is not None + or self._speculative_config.max_concurrency is not None): + raise ValueError( + "Unified standalone 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._llm_args.cuda_graph_config is not None: + raise ValueError( + "Unified standalone DSpark KV cache currently requires eager " + "execution; set cuda_graph_config=None.") + if not self._disable_overlap_scheduler: + raise ValueError("Unified standalone DSpark KV cache requires " + "disable_overlap_scheduler=True.") + if self._llm_args.enable_chunked_prefill: + raise ValueError( + "Unified standalone DSpark KV cache does not yet support " + "chunked prefill; set enable_chunked_prefill=False.") + if self._mapping.pp_size != 1 or self._mapping.cp_size != 1: + raise ValueError( + "Unified standalone DSpark KV cache requires PP=1 and CP=1.") + if self._mapping.enable_attention_dp: + raise ValueError( + "Unified standalone DSpark KV cache does not yet support " + "attention data parallelism; set enable_attention_dp=False.") + if self._kv_connector_manager is not None: + raise ValueError( + "Unified standalone 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"): + raise ValueError( + "Standalone DSpark draft-state transfer requires the PYTHON " + "NIXL transceiver on both workers.") + if self._draft_config is None: + raise ValueError( + "Unified standalone DSpark KV cache requires the loaded " + "standalone draft model configuration.") + + def _get_standalone_draft_layout(self) -> StandaloneDraftLayout: + """Describe distinct standalone layers without borrowing target shapes.""" + 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.") + 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=self._speculative_config.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. @@ -1663,6 +1781,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 @@ -2519,39 +2639,44 @@ def _get_qwen4_exp_ple_cache_params(config, *, total_layers: int, def _create_kv_cache_manager( - model_engine: Optional[PyTorchModelEngine], - kv_cache_manager_cls, - mapping: Mapping, - kv_cache_config: KvCacheConfig, - tokens_per_block: int, - max_seq_len: int, - max_batch_size: int, - spec_config: Optional[SpeculativeConfig], - sparse_attention_config: Optional[SparseAttentionConfig], - max_num_tokens: int, - max_beam_width: int, - kv_connector_manager: Optional[KvCacheConnectorManager], - estimating_kv_cache: bool = False, - enable_kv_cache_stats: bool = False, - execution_stream: Optional[torch.cuda.Stream] = None, - # Optional overrides for one-model draft case (when model_engine is None) - model_config: Optional[ModelConfig] = None, - dtype: Optional[torch.dtype] = None, - is_draft: Optional[bool] = None, - layer_mask: Optional[List[bool]] = None, - num_layers: Optional[int] = None, - num_kv_heads: Optional[Union[int, List[int]]] = None, - head_dim: Optional[int] = None, - kv_cache_type=None, - is_disagg: bool = False, - disable_overlap_scheduler: bool = False, - cold_page_codec_provider: Optional[object] = None, - kv_events_config: Optional[KVEventsConfig] = None, - joint_kv_cache_reuse: bool = False) -> KVCacheManager: + model_engine: Optional[PyTorchModelEngine], + kv_cache_manager_cls, + mapping: Mapping, + kv_cache_config: KvCacheConfig, + tokens_per_block: int, + max_seq_len: int, + max_batch_size: int, + spec_config: Optional[SpeculativeConfig], + sparse_attention_config: Optional[SparseAttentionConfig], + max_num_tokens: int, + max_beam_width: int, + kv_connector_manager: Optional[KvCacheConnectorManager], + estimating_kv_cache: bool = False, + enable_kv_cache_stats: bool = False, + execution_stream: Optional[torch.cuda.Stream] = None, + # Optional overrides for one-model draft case (when model_engine is None) + model_config: Optional[ModelConfig] = None, + dtype: Optional[torch.dtype] = None, + is_draft: Optional[bool] = None, + layer_mask: Optional[List[bool]] = None, + num_layers: Optional[int] = None, + num_kv_heads: Optional[Union[int, List[int]]] = None, + head_dim: Optional[int] = None, + kv_cache_type=None, + is_disagg: bool = False, + disable_overlap_scheduler: bool = False, + cold_page_codec_provider: Optional[object] = None, + kv_events_config: Optional[KVEventsConfig] = None, + joint_kv_cache_reuse: bool = False, + standalone_draft_layout: Optional[StandaloneDraftLayout] = 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( @@ -2716,6 +2841,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 # V2 builds the block-reuse cache key of a multimodal token run from 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 4caa47a29ba7..b8d4cd13ac70 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 @@ -115,6 +115,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 @@ -988,6 +989,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, @@ -1129,6 +1132,9 @@ def _settle_context_cursor(req: LlmRequest, reuse: int, tokens_per_block: int) - 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 @@ -1168,8 +1174,36 @@ 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, **kwargs, ) -> None: + self.draft_layout = standalone_draft_layout + self.draft_layer_ids: tuple[int, ...] = () + self.draft_history: dict[int, StandaloneDraftHistory] = {} + 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: + if KV_CACHE_MANAGER_V2_BACKEND != "python": + raise ValueError( + "Unified standalone draft KV requires TLLM_KV_CACHE_MANAGER_V2_BACKEND=python" + ) + if is_draft or mapping.pp_size != 1 or mapping.cp_size != 1: + raise ValueError( + "Unified standalone draft KV requires one target manager and PP=CP=1" + ) + if kv_cache_config.enable_block_reuse: + raise ValueError("Unified standalone draft KV does not yet support prefix reuse") + if kv_connector_manager is not None: + raise ValueError("Unified standalone draft KV does not yet support KV connectors") + if kv_cache_config.enable_swa_scratch_reuse: + raise ValueError( + "Unified standalone draft KV does not yet support SWA scratch reuse" + ) + if kv_cache_config.pool_ratio is not None: + raise ValueError( + "Unified standalone draft KV does not yet support explicit pool_ratio" + ) self.mapping = mapping self.dtype = dtype self.is_disagg = is_disagg @@ -1266,6 +1300,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 @@ -1527,6 +1562,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 @@ -2007,7 +2043,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): @@ -2044,6 +2081,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 role_b = None if self.kv_cache_type == CacheTypeCpp.SELFKONLY else Role.VALUE return Role.KEY, role_b @@ -2052,6 +2093,11 @@ def _get_block_scale_role(self, role_a: DataRole) -> Optional[DataRole]: 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) + def _build_pool_mapping_tensors(self): """Build the (kv_cache_pool_pointers, kv_cache_pool_mapping) tensors. @@ -2074,7 +2120,7 @@ def _build_pool_mapping_tensors(self): ] ) if self.dtype == DataType.NVFP4: - block_scale_role = self._get_block_scale_role(role_a) + 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( @@ -2106,7 +2152,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) + 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 = ( @@ -2146,7 +2192,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) + block_scale_role = self._get_layer_block_scale_role(layer_id, role_a) if block_scale_role is None: block_scale_offset = None else: @@ -2271,7 +2317,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, ) @@ -2283,6 +2329,13 @@ 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([None] * self.draft_layout.num_layers) return layer_sizes, attention_windows def _get_max_tokens_from_quota(self, quota: int) -> float: @@ -2310,6 +2363,7 @@ def _get_max_tokens_from_quota_impl(self, quota: int) -> float: generation_capacity_headroom=self._generation_kv_capacity_headroom, ) size_per_batch = self.max_batch_size * generation_swa_size_per_request + size_per_batch += self._standalone_draft_fixed_bytes(layer_sizes) if quota < size_per_batch: return 0 context_limit_quota = self.max_num_tokens * context_size_per_token + size_per_batch @@ -2351,8 +2405,14 @@ def _get_quota_from_max_tokens_impl(self, max_tokens: int) -> int: context_tokens * context_size_per_token + generation_tokens * generation_size_per_token + self.max_batch_size * generation_swa_size_per_request + + self._standalone_draft_fixed_bytes(layer_sizes) ) + def _standalone_draft_fixed_bytes(self, layer_sizes: Sequence[int]) -> int: + # All groups share one conservative allocation envelope in V2. Its + # scratch tail is capacity, never evidence of valid draft history. + return self.max_batch_size * self._standalone_draft_reserve * sum(layer_sizes) + def _get_event_num_blocks_per_cache_level( self, cache_tiers: List[CacheTierConfig], @@ -2843,6 +2903,61 @@ 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 + ) -> 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 // 2 * self.tokens_per_block, + ) + for role in (Role.KEY, Role.VALUE) + ], + cache_domain="standalone_draft", + ) + ) + self.pp_layers.append(global_id) + self.layer_offsets[global_id] = local_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(None) + 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( @@ -2867,6 +2982,11 @@ 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} + if self.draft_layout is not None + else {} + ), ) ) else: @@ -2979,6 +3099,8 @@ 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] + if self._is_standalone_draft_layer(layer_offset): + return self.get_draft_buffers(self.draft_layer_ids.index(layer_idx), kv_layout) addr_key = self.impl.get_mem_pool_base_address(layer_offset, Role.KEY, PageIndexMode.SHARED) if self.kv_cache_type != CacheTypeCpp.SELFKONLY: addr_value = self.impl.get_mem_pool_base_address( @@ -3023,6 +3145,123 @@ def get_buffers(self, layer_idx: int, kv_layout: str = "NHD") -> Optional[torch. ) ) + def _is_standalone_draft_layer(self, local_layer_idx: int) -> bool: + return ( + getattr(self, "draft_layout", None) is not None + and self.pp_layers[local_layer_idx] in self.draft_layer_ids + ) + + def get_layer_cache_dtype(self, layer_idx: int) -> DataType: + if self.draft_layout is None: + return self.dtype + if self._is_standalone_draft_layer(self.layer_offsets[layer_idx]): + 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: + if self.draft_layout is None: + return self.kv_factor + return ( + 2 if self._is_standalone_draft_layer(self.layer_offsets[layer_idx]) else self.kv_factor + ) + + def get_draft_buffers(self, local_layer_idx: int, kv_layout: str = "HND") -> torch.Tensor: + """View authoritative draft K/V pages using the standalone 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") + if kv_layout not in ("HND", "NHD"): + raise ValueError(f"Unsupported standalone draft KV layout: {kv_layout}") + layer_id = self.layer_offsets[self.draft_layer_ids[local_layer_idx]] + key_address = self.impl.get_mem_pool_base_address(layer_id, Role.KEY, PageIndexMode.SHARED) + value_address = self.impl.get_mem_pool_base_address( + layer_id, Role.VALUE, PageIndexMode.SHARED + ) + stride = self.impl.get_page_stride(layer_id, Role.KEY) + if value_address != key_address + stride: + raise ValueError("Standalone draft K/V buffers must have adjacent equal-sized pages") + dimensions = ( + [layout.num_kv_heads, self.tokens_per_block, layout.head_dim] + if kv_layout == "HND" + else [self.tokens_per_block, layout.num_kv_heads, layout.head_dim] + ) + return convert_to_torch_tensor( + TensorWrapper( + key_address, + self.get_layer_cache_dtype(self.draft_layer_ids[local_layer_idx]), + [self.impl.get_page_index_upper_bound(layer_id, Role.KEY) // 2, 2, *dimensions], + ) + ) + + def get_draft_block_table(self, request_ids: List[int]) -> torch.Tensor: + """Current rank-local page mappings; unused tail entries point to page zero.""" + if self.draft_layout is None: + raise ValueError("No unified standalone draft cache is configured") + layer_id = self.layer_offsets[self.draft_layer_ids[0]] + pool_id = self.impl.get_layer_group_id(layer_id) + scale = self.impl.get_page_index_scale(layer_id, Role.KEY) + table = torch.zeros((len(request_ids), self.max_blocks_per_seq), dtype=torch.int32) + for row, request_id in enumerate(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") + indices = cache.get_base_page_indices(pool_id)[: cache.num_blocks] + 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") + if any(index == BAD_PAGE_INDEX for index in indices): + raise ValueError("Standalone full-attention draft cache contains missing pages") + table[row, : len(indices)] = torch.tensor( + [int(index) * int(scale) // 2 for index in 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: + if self.draft_layout is None: + raise ValueError("No unified standalone draft cache is configured") + 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.valid_length > cache.capacity: + raise ValueError("Standalone draft history exceeds allocated capacity") + 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") + if type(valid_length) is not int or type(position) is not int: + raise ValueError( + "Standalone draft transfer history must use integer lengths and positions" + ) + # Accessing the receiving mapping here verifies ownership before the + # worker can see this history. The sender's slot/page IDs are never used. + self.get_draft_block_table([request_id]) + self.set_draft_history(request_id, valid_length, position) + def get_index_k_buffer( self, layer_idx: int, @@ -3157,7 +3396,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 = ( @@ -3184,9 +3425,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)) ] ) @@ -3195,7 +3437,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): @@ -3586,7 +3828,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 @@ -3641,7 +3888,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 @@ -4901,7 +5153,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 @@ -5030,6 +5284,8 @@ 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._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. @@ -5070,6 +5326,7 @@ def get_batch_cache_indices( is_kv_aggregate=True, 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( @@ -5080,11 +5337,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 @@ -5140,7 +5398,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() @@ -5161,6 +5421,8 @@ def get_batch_cache_indices_flat( return out_tensor def get_cache_bytes_per_token(self) -> int: + if getattr(self, "draft_layout", None) 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) @@ -5176,6 +5438,12 @@ 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): + return self.draft_layout.bytes_per_layer_token // 2 + return 0 if self.dtype not in ( DataType.FP8, DataType.HALF, @@ -5286,6 +5554,8 @@ 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._request_stats_enabled_ids.clear() self._fresh_pages_filled.clear() # Drop the outstanding plans before the manager shuts down: discarding a handle applies @@ -5491,6 +5761,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.valid_length + self._standalone_draft_reserve + ) history_length = ( None # Reuse (history's consumer) is disabled under helix, and 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 8dd3622f0a6e..b68015754c93 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache/mamba_cache_manager.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/mamba_cache_manager.py @@ -3779,7 +3779,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]]: @@ -4194,7 +4196,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..6e5f92d7cedc --- /dev/null +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/standalone_draft_cache.py @@ -0,0 +1,65 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Storage and history contracts for standalone DSpark/DFlash drafters.""" + +from dataclasses import dataclass + +import torch + + +@dataclass(frozen=True) +class StandaloneDraftLayout: + """Rank-local full-attention storage, independent of target KV geometry.""" + + num_layers: int + num_kv_heads: int + head_dim: int + dtype: torch.dtype + extra_tokens: int + attention_backend: str + + 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"): + raise ValueError("Standalone draft KV requires VANILLA or TRTLLM attention") + + @property + def bytes_per_layer_token(self) -> int: + return 2 * self.num_kv_heads * self.head_dim * self.dtype.itemsize + + @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, + } + + +@dataclass(frozen=True) +class StandaloneDraftHistory: + """Committed draft tokens and their next absolute sequence position. + + These are deliberately separate from the target cache's monotonic history + watermark and from its speculative allocation capacity. + """ + + 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") diff --git a/tensorrt_llm/_torch/speculative/dflash.py b/tensorrt_llm/_torch/speculative/dflash.py index 939e7366d327..61176860df56 100644 --- a/tensorrt_llm/_torch/speculative/dflash.py +++ b/tensorrt_llm/_torch/speculative/dflash.py @@ -28,7 +28,7 @@ from ..attention.backends import AttentionMetadata from ..pyexecutor.kv_cache.mamba_cache_manager import MambaHybridCacheManager 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, @@ -43,6 +43,7 @@ if TYPE_CHECKING: from ...llmapi.llm_args import DFlashDecodingConfig + from ..pyexecutor.resource_manager import ResourceManager def dflash_draft_slot_ids( @@ -335,6 +336,95 @@ def set_draft_model(self, draft_model) -> None: super().set_draft_model(draft_model) self._validate_draft_attention_backend(draft_model) + 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(getattr(self, "_ctx_kv_manager", None), "draft_layout", None) is not None + + def _restore_managed_slots(self, request_ids: list[int], num_contexts: int) -> None: + """Bind generation requests after initialization or a manager replacement. + + Metadata preparation precedes lazy buffer binding. That binding can + replace an estimation manager and clear its staging slots, so generation + must rebind here even when metadata preparation already assigned slots. + """ + updates = {self._dummy_slot: 0} + for request_id in request_ids[num_contexts:]: + if ( + request_id != ATTENTION_DP_DUMMY_REQUEST_ID + and request_id < self._graph_dummy_id_floor + ): + self._assign_slot(request_id) + slot = self._req_to_slot.get(request_id) + if slot is not None: + history = self._ctx_kv_manager.get_draft_history(request_id) + # Request IDs can be recycled by startup probes or restarted + # requests. Only the manager knows whether pages survived. + 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 + self._write_ctx_len(updates) + 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, + ) + ) + + def _store_managed_context_kv( + self, k: torch.Tensor, v: torch.Tensor, rows: torch.Tensor, positions: torch.Tensor + ) -> None: + """Write projected accepted features to the authoritative draft pages. + + K/V have shape [tokens, layers, heads, head_dim]. Rows address this + iteration's request table, independently of the worker's staging slots. + """ + pages = self._ctx_block_tables[rows, positions // self._ctx_page_size].long() + offsets = positions % self._ctx_page_size + for layer_idx, pool in enumerate(self._ctx_kv_buf): + pool[pages, 0, :, offsets, :] = k[:, layer_idx] + pool[pages, 1, :, offsets, :] = v[:, layer_idx] + + def _gather_managed_context(self, request_ids: list[int]) -> None: + """Refresh VANILLA's dense input staging from manager-owned history. + + The dense tensors are forward workspaces: they are never exported or + used to recover request history. Noise K/V may overwrite their suffix. + """ + if self._dflash_attention_backend != "VANILLA": + return + for row, request_id in enumerate(request_ids): + 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 _publish_managed_history(self, request_ids: list[int]) -> None: + """Commit draft validity after a successful eager worker forward.""" + lengths = self._ctx_len.tolist() + for request_id in request_ids: + slot = self._req_to_slot.get(request_id) + if slot is None: + continue + length = lengths[slot] + position = self._req_ctx_pos.get(request_id, 0) + length - self._ctx_len_host[slot] + self._ctx_len_host[slot] = length + self._req_ctx_pos[request_id] = position + self._ctx_kv_manager.set_draft_history(request_id, length, position) + def _check_ctx_arena_fits(self, capacity, num_slots, L, nkv, hd, dtype): """Fail with the arithmetic before allocating the drafter context arena. @@ -473,7 +563,9 @@ def _init_ctx_block_tables( ) return True - 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: Optional[list[int]] = 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 @@ -482,6 +574,12 @@ 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(): + if request_ids is None: + raise ValueError("Unified DSpark draft pages require local request IDs") + table = self._ctx_kv_manager.get_draft_block_table(request_ids[:num_seqs]) + self._ctx_block_tables[:num_seqs].copy_(table, non_blocking=True) + 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 @@ -507,8 +605,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; " @@ -601,7 +702,42 @@ def _lazy_init_ctx_buffers( nkv = draft_model._num_kv_heads hd = draft_model._head_dim capacity = self._max_ctx + self._compute_block_size - if self._dflash_attention_backend in _PAGED_ATTENTION_BACKENDS: + if unified: + self._ctx_kv_buf = [ + 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(pool.shape[1:]) != expected or pool.dtype != dtype + for pool in self._ctx_kv_buf + ): + raise ValueError("Unified DSpark draft pool does not match the drafter KV layout") + self._ctx_block_tables = torch.zeros( + (num_slots, draft_kv_cache_manager.max_blocks_per_seq), + dtype=torch.int32, + device="cuda", + ) + if self._dflash_attention_backend == "VANILLA": + # Dense FlashAttention inputs are transient forward staging; + # every history read is refreshed from the authoritative pool. + kv_shape = (num_slots, L, capacity, nkv, hd) + self._ctx_k_buf = torch.zeros(kv_shape, dtype=dtype, device="cuda") + self._ctx_v_buf = torch.zeros(kv_shape, dtype=dtype, device="cuda") + elif self._dflash_attention_backend == "TRTLLM": + validate_dflash_trtllm_gen_runtime( + dtype=dtype, + num_heads=nh, + num_kv_heads=nkv, + head_dim=hd, + tokens_per_block=self._ctx_page_size, + has_context_attention=any( + not draft_model._get_attention_mask_args(i)[0] for i in range(L) + ), + ) + else: + raise ValueError("Unified DSpark supports VANILLA and TRTLLM draft attention") + elif self._dflash_attention_backend in _PAGED_ATTENTION_BACKENDS: pool = ( self._managed_ctx_pool(draft_kv_cache_manager, L, nkv, hd, dtype) if self._dflash_attention_backend == "TRTLLM" @@ -714,10 +850,27 @@ 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 - clear(slot) + history = ( + self._ctx_kv_manager.get_draft_history(req_id) + if self._has_unified_draft_cache() and not reset + else None + ) + if history is None: + clear(slot) + else: + # The receiving manager owns validity. A new local staging + # slot must preserve it, regardless of the sender's slot ID. + if updates is None: + self._write_ctx_len({slot: history.valid_length}) + else: + updates[slot] = history.valid_length + self._ctx_len_host[slot] = history.valid_length + self._req_ctx_pos[req_id] = history.position return self._req_to_slot[req_id] def _get_ctx_paged_append(self) -> Callable[..., None]: @@ -879,7 +1032,12 @@ def _store_prefill_context( # where the previous one ended. Everything else -- a fresh request, a # request id reused after completion, a prefill restarted after # preemption -- has to start from a clean slot. - reset = self._req_ctx_pos.get(req_id) != first_pos + previous_position = self._req_ctx_pos.get(req_id) + if self._has_unified_draft_cache(): + history = self._ctx_kv_manager.get_draft_history(req_id) + if history is not None: + previous_position = history.position + reset = previous_position != first_pos if self._assign_slot(req_id, reset=reset, updates=ctx_len_updates) is None: logger.warning("DFlash: no free slots, skipping context store") self._req_ctx_pos.pop(req_id, None) @@ -889,6 +1047,8 @@ def _store_prefill_context( slot = self._req_to_slot[req_id] cur = ctx_len_updates.get(slot, self._ctx_len_host[slot]) if cur + slen > self._max_ctx: + 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. @@ -923,7 +1083,14 @@ def _store_prefill_context( chunk_proj_cast, chunk_pos[:actual] ) # chunk_k/v: [actual, L, nkv, hd] → [L, actual, nkv, hd] - if self._dflash_attention_backend in _PAGED_ATTENTION_BACKENDS: + if self._has_unified_draft_cache(): + self._store_managed_context_kv( + chunk_k, + chunk_v, + torch.full((actual,), i, dtype=torch.long, device="cuda"), + torch.arange(cur, end, dtype=torch.long, device="cuda"), + ) + elif self._dflash_attention_backend in _PAGED_ATTENTION_BACKENDS: # Manager block tables are keyed by batch position, the # private arena's by slot. See _ctx_paged_index_args. row = i if self._ctx_block_tables is not None else slot @@ -997,8 +1164,14 @@ def _forward_impl( draft_model, spec_metadata, attn_metadata, draft_kv_cache_manager ) spec_metadata._dflash_worker = self + if self._has_unified_draft_cache(): + self._restore_managed_slots(spec_metadata.request_ids, num_contexts) # Before any store: prefill and decode both address pages through it. - self._refresh_ctx_block_tables(attn_metadata, batch_size) + self._refresh_ctx_block_tables(attn_metadata, batch_size, spec_metadata.request_ids) + if self._has_unified_draft_cache(): + # Refresh before appending accepted features; the first generation + # iteration therefore consumes the transferred prompt history. + self._gather_managed_context(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. @@ -1092,7 +1265,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"], @@ -1202,6 +1377,8 @@ def _forward_impl( if is_warmup: self._ctx_len.copy_(self._saved_ctx_len) self._restore_ctx_len_host() + elif self._has_unified_draft_cache(): + self._publish_managed_history(spec_metadata.request_ids) self._ctx_len_restore_pending = False return { @@ -1419,15 +1596,32 @@ def prepare_1st_drafter_inputs( bonus = gen_accepted_tokens.gather(1, bonus_idx).squeeze(1).long() ctx_len_gen = self._ctx_len[slots] + if ( + self._has_unified_draft_cache() + and torch.any(ctx_len_gen + gen_num_accepted > self._max_ctx).item() + ): + raise ValueError("Unified DSpark accepted history exceeds the drafter context") + ctx_position_gen = ctx_len_gen + if self._has_unified_draft_cache(): + ctx_position_gen = torch.tensor( + [ + self._req_ctx_pos.get(request_id, 0) + for request_id in spec_metadata.request_ids[ + num_contexts : num_contexts + num_gens + ] + ], + dtype=torch.long, + device=ctx_len_gen.device, + ) 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") query_position_ids = ( - ctx_len_gen.unsqueeze(1) + ctx_position_gen.unsqueeze(1) + gen_num_accepted.long().unsqueeze(1) + j_block.unsqueeze(0) ) - ctx_position_ids = ctx_len_gen.unsqueeze(1) + offsets_kp1.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. @@ -1483,7 +1677,13 @@ def prepare_1st_drafter_inputs( v_new.mul_(mask_bc) slot_long = slot_flat.long() col_long = col_flat.long() - if self._dflash_attention_backend in _PAGED_ATTENTION_BACKENDS: + if self._has_unified_draft_cache(): + rows_long = gen_rows_out.unsqueeze(1).expand(-1, K + 1).reshape(-1) + self._store_managed_context_kv(k_new, v_new, rows_long, col_long) + if self._dflash_attention_backend == "VANILLA": + self._ctx_k_buf[slot_long, :, col_long] = k_new + self._ctx_v_buf[slot_long, :, col_long] = v_new + elif self._dflash_attention_backend in _PAGED_ATTENTION_BACKENDS: if self._ctx_block_tables is not None: # Batch positions of the gen requests, matching the # per-request block table's row order. @@ -1524,7 +1724,12 @@ 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._dflash_attention_backend in _PAGED_ATTENTION_BACKENDS + 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/runtime/kv_cache_manager_v2/__init__.pyi b/tensorrt_llm/runtime/kv_cache_manager_v2/__init__.pyi index 9388dc9e84d9..1c323523fc09 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/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py index c47c59c5434c..18519741b2ec 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py @@ -112,6 +112,12 @@ class AttentionLayerConfig: # Note that we use None to represent "no sliding window". Sink tokens are excluded. sliding_window_size: int | None = None num_sink_tokens: int | None = None + cache_domain: str = "target" + """Ownership domain for layers that can share a lifecycle and physical pools. + + Standalone draft layers use a separate domain because their valid history and + speculative scratch need not advance with the target, even for identical layouts. + """ @property def window_size(self) -> int | None: diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_life_cycle_registry.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_life_cycle_registry.py index d0da6b057678..4d4c3cef7812 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_life_cycle_registry.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_life_cycle_registry.py @@ -23,17 +23,21 @@ class AttnLifeCycle(NamedTuple): window_size: SlidingWindowSize num_sink_blocks: int # div_up(num_sink_tokens, tokens_per_block) + cache_domain: str = "target" @staticmethod def make( - window_size: SlidingWindowSize, num_sink_tokens: int | None, tokens_per_block: int + window_size: SlidingWindowSize, + num_sink_tokens: int | None, + tokens_per_block: int, + cache_domain: str = "target", ) -> "AttnLifeCycle": assert tokens_per_block > 0 assert window_size is None or window_size > 0 assert num_sink_tokens is None or num_sink_tokens >= 0 assert num_sink_tokens in (None, 0) or window_size is not None num_sink_blocks = div_up(num_sink_tokens or 0, tokens_per_block) - return AttnLifeCycle(window_size, num_sink_blocks) + return AttnLifeCycle(window_size, num_sink_blocks, cache_domain) def get_stale_range( self, history_length: int, tokens_per_block: int @@ -73,7 +77,9 @@ def make_life_cycle(layer: LayerConfig, tokens_per_block: int) -> LifeCycle: return ssm_life_cycle else: assert isinstance(layer, AttentionLayerConfig) - return AttnLifeCycle.make(layer.window_size, layer.num_sink_tokens, tokens_per_block) + return AttnLifeCycle.make( + layer.window_size, layer.num_sink_tokens, tokens_per_block, layer.cache_domain + ) class LifeCycleRegistry: diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_storage/_config.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_storage/_config.py index 6a62470e582a..a42c7d7b860b 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_storage/_config.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_storage/_config.py @@ -20,7 +20,13 @@ from .._common import LayerId from .._config import CacheTierConfig, DataRole, KVCacheManagerConfig -from .._life_cycle_registry import LayerGroupId, LifeCycleId, LifeCycleRegistry, make_life_cycle +from .._life_cycle_registry import ( + AttnLifeCycle, + LayerGroupId, + LifeCycleId, + LifeCycleRegistry, + make_life_cycle, +) from .._storage._core import PoolGroupIndex, PoolIndex from .._utils import ( HomoTuple, @@ -225,15 +231,20 @@ def create_storage_config(config: KVCacheManagerConfig) -> StorageConfig: slot_groups.append( SlotDescVariant(life_cycle_id, cast(TypedIndexList[PoolIndex, CoalescedBuffer], slots)) ) - # Merge slot groups with the same slot_size_list - pool_groups_by_slot_size_list = defaultdict[HomoTuple[int], list[SlotDescVariant]]( + # Equal storage sizes permit merging only within a compatible ownership domain. + # Existing target attention/SSM groups retain their shared physical pool behavior. + pool_groups_by_layout = defaultdict[tuple[str, HomoTuple[int]], list[SlotDescVariant]]( list[SlotDescVariant] ) for slot_group in slot_groups: - pool_groups_by_slot_size_list[tuple(slot_group.slot_size_list)].append(slot_group) + life_cycle = life_cycle_registry[slot_group.life_cycle_id] + cache_domain = ( + life_cycle.cache_domain if isinstance(life_cycle, AttnLifeCycle) else "target" + ) + pool_groups_by_layout[(cache_domain, tuple(slot_group.slot_size_list))].append(slot_group) slot_desc_list = cast( TypedIndexList[PoolGroupIndex, SlotDesc], - [SlotDesc(tuple(slot_groups)) for slot_groups in pool_groups_by_slot_size_list.values()], + [SlotDesc(tuple(slot_groups)) for slot_groups in pool_groups_by_layout.values()], ) return StorageConfig( cache_tiers=tuple(config.cache_tiers), From 84091cb078ddaef45cc7cf4d8b1b8a8524b86c2f Mon Sep 17 00:00:00 2001 From: Allison Lim Date: Fri, 18 Sep 2026 22:11:04 -0700 Subject: [PATCH 02/13] [None][fix] Preserve standalone DSpark cache compatibility and cleanup Validate received draft metadata, receiver allocation, and history before rank completion consensus so invalid state follows coordinated request failure and cleanup. Preserve the existing aggregate C++ V2 and CUDA graph cache paths while keeping unsupported unified disaggregation configurations explicit. Share draft reserve sizing across creator, attention, and hybrid cache estimators without changing the runtime byte cap. Signed-off-by: Allison Lim --- .../_torch/disaggregation/transceiver.py | 25 +++++++++++-- tensorrt_llm/_torch/pyexecutor/_util.py | 30 ++++++++++------ .../kv_cache/kv_cache_manager_v2.py | 36 +++++++++++-------- .../kv_cache/mamba_cache_manager.py | 19 +++++++++- 4 files changed, 80 insertions(+), 30 deletions(-) diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index bed43106e265..b118be0100d5 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -1291,6 +1291,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: @@ -1306,6 +1309,21 @@ 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: + self._apply_aux(session, req) + self._assert_disagg_history_declared(req) + # Restore also validates the receiver's allocation and page map. + # A peer failure below leaves this request unschedulable; its + # ordinary failure cleanup releases any restored history and KV. + self._restore_draft_history(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: @@ -1353,10 +1371,11 @@ 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): + if not has_draft_history and self._need_aux_transfer(req): self._apply_aux(session, req) - self._assert_disagg_history_declared(req) - self._restore_draft_history(req) + if not has_draft_history: + self._assert_disagg_history_declared(req) + self._restore_draft_history(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 598ef17e0c14..11490c27c12e 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -930,20 +930,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) - if self._uses_unified_standalone_draft_cache(): - draft_layout = self._get_standalone_draft_layout() - # The unified allocator reserves an envelope for draft capture and - # noise in every attention group. Include its fixed cost as well - # as the distinct BF16 draft layers in the common byte budget. - total += CacheCost( - slope=draft_layout.bytes_per_token, - intercept=(total.slope + draft_layout.bytes_per_token) * - draft_layout.extra_tokens * self._max_batch_size) + 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( @@ -1680,7 +1675,18 @@ def _is_standalone_dspark(self) -> bool: and not spec_config._use_shared_kv_cache) def _uses_unified_standalone_draft_cache(self) -> bool: - return self._is_standalone_dspark() and self._is_kv_cache_manager_v2 + if not self._is_standalone_dspark() or not self._is_kv_cache_manager_v2: + return False + if self._is_disagg: + # Disaggregation must validate unified ownership, never fall back + # to draft state that the transceiver cannot transfer. + return True + from tensorrt_llm.runtime.kv_cache_manager_v2 import BACKEND + + # Preserve the existing aggregate C++/CUDA-graph draft-cache path. + # Python V2 with eager execution also supports unified aggregate runs + # for comparison with the same configuration in disaggregation. + return BACKEND == "python" and self._llm_args.cuda_graph_config is None def _validate_standalone_draft_cache(self) -> None: """Reject unsupported standalone state ownership before profiling.""" @@ -1699,6 +1705,8 @@ def _validate_standalone_draft_cache(self) -> None: "kv_cache_config.use_kv_cache_manager_v2=True and " "TLLM_KV_CACHE_MANAGER_V2_BACKEND=python on both workers.") return + if not self._uses_unified_standalone_draft_cache(): + return from tensorrt_llm.runtime.kv_cache_manager_v2 import BACKEND if BACKEND != "python": 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 b8d4cd13ac70..1d1b25c9143e 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 @@ -516,11 +516,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) @@ -538,7 +541,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), ) @@ -2354,16 +2357,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.max_batch_size * generation_swa_size_per_request - size_per_batch += self._standalone_draft_fixed_bytes(layer_sizes) + size_per_batch = self.max_batch_size * 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 @@ -2391,28 +2394,23 @@ 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.max_batch_size * generation_swa_size_per_request - + self._standalone_draft_fixed_bytes(layer_sizes) + + self.max_batch_size * generation_size_per_request ) - def _standalone_draft_fixed_bytes(self, layer_sizes: Sequence[int]) -> int: - # All groups share one conservative allocation envelope in V2. Its - # scratch tail is capacity, never evidence of valid draft history. - return self.max_batch_size * self._standalone_draft_reserve * sum(layer_sizes) - def _get_event_num_blocks_per_cache_level( self, cache_tiers: List[CacheTierConfig], @@ -5600,6 +5598,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( @@ -5632,10 +5631,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([None] * 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, @@ -5646,11 +5651,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( 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 b68015754c93..faf33472890e 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 ( @@ -2131,6 +2134,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") @@ -2193,6 +2197,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 From 5572c6e471b4c5484a2e1b1d4fea110c2ff26dda Mon Sep 17 00:00:00 2001 From: Allison Lim Date: Fri, 18 Sep 2026 23:00:05 -0700 Subject: [PATCH 03/13] [None][refactor] Simplify standalone DSpark cache validation Remove repeated creator, manager, transport, and worker checks. Serialize validated draft history directly and keep wire version, prompt coverage, storage identity, allocation, and page validation at their owning boundaries. Preserve coordinated receive failure and cleanup before completion consensus. Local validation: 230 tests passed, 34 skipped; no new GPU or AL measurement. Signed-off-by: Allison Lim --- .../_torch/disaggregation/native/auxiliary.py | 47 +++++++------------ .../_torch/disaggregation/transceiver.py | 16 +------ tensorrt_llm/_torch/pyexecutor/_util.py | 4 -- .../kv_cache/kv_cache_manager_v2.py | 18 ------- tensorrt_llm/_torch/speculative/dflash.py | 8 +--- 5 files changed, 20 insertions(+), 73 deletions(-) diff --git a/tensorrt_llm/_torch/disaggregation/native/auxiliary.py b/tensorrt_llm/_torch/disaggregation/native/auxiliary.py index 633e549c24da..e7faca0d27bd 100644 --- a/tensorrt_llm/_torch/disaggregation/native/auxiliary.py +++ b/tensorrt_llm/_torch/disaggregation/native/auxiliary.py @@ -121,42 +121,31 @@ def build_aux_transfer_layout( def _encode_draft_history(history: dict[str, Any]) -> list[int]: - """Encode committed history and rank-local storage identity for the wire.""" - if not isinstance(history, dict): + """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.get("layout") - if not isinstance(layout, dict): - raise ValueError("Standalone draft transfer requires a storage layout") - integer_values = [ - history.get("valid_length"), - history.get("position"), - layout.get("num_layers"), - layout.get("num_kv_heads"), - layout.get("head_dim"), + 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"]], ] - if any(type(value) is not int for value in integer_values): - raise ValueError( - "Standalone draft transfer metadata requires integer lengths and dimensions" - ) - valid_length, position, num_layers, num_kv_heads, head_dim = integer_values - if valid_length < 0 or position < valid_length: - raise ValueError("Invalid standalone draft history length or position") - if min(num_layers, num_kv_heads, head_dim) <= 0: - raise ValueError("Standalone draft transfer dimensions must be positive") - dtype_code = _DRAFT_DTYPE_CODES.get(layout.get("dtype")) - backend_code = _DRAFT_BACKEND_CODES.get(layout.get("attention_backend")) - if dtype_code is None or backend_code is None: - raise ValueError("Unsupported standalone draft transfer dtype or attention backend") - return [_DRAFT_HISTORY_VERSION, *integer_values, dtype_code, backend_code] def _decode_draft_history(values: list[int]) -> dict[str, Any]: - if len(values) != _DRAFT_HISTORY_FIELDS or values[0] != _DRAFT_HISTORY_VERSION: + 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 = values dtypes = {code: name for name, code in _DRAFT_DTYPE_CODES.items()} backends = {code: name for name, code in _DRAFT_BACKEND_CODES.items()} - history = { + # The receiving transceiver validates prompt coverage and the manager checks + # storage identity and local allocation before publishing this history. + return { "valid_length": valid_length, "position": position, "layout": { @@ -167,8 +156,6 @@ def _decode_draft_history(values: list[int]) -> dict[str, Any]: "attention_backend": backends.get(backend_code), }, } - _encode_draft_history(history) - return history class AuxBufferBase(ABC): @@ -409,7 +396,7 @@ def get_slot_data(self, slot: int) -> tuple[list[int], list[int], tuple[int, int return first_gen_tokens, draft_tokens, (int(prompt_tokens), int(cached_tokens)) def get_slot_draft_history(self, slot: int) -> dict[str, Any]: - """Read transferred history, rejecting missing or unsupported metadata.""" + """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: diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index b118be0100d5..58b97e1a03c0 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -549,14 +549,7 @@ def _validate_draft_transfer(self, req: LlmRequest) -> None: def _validate_draft_history_range(req: LlmRequest, history: dict) -> None: # The shared transfer extent covers the complete prompt. Draft noise KV is scratch, # never valid history. A partial history needs per-group extents before it can be sent. - if ( - not isinstance(history, dict) - or type(history.get("valid_length")) is not int - or type(history.get("position")) is not int - or not isinstance(history.get("layout"), dict) - or history["valid_length"] != req.prompt_len - or history["position"] != req.prompt_len - ): + if history["valid_length"] != req.prompt_len or history["position"] != req.prompt_len: raise ValueError( "Standalone DSpark transfer requires valid draft history and sequence position " f"covering the complete prompt ({req.prompt_len} tokens)." @@ -573,7 +566,6 @@ def _pack_draft_history(self, req: LlmRequest) -> Optional[dict]: return history def _received_draft_history(self, req: LlmRequest) -> Optional[dict]: - self._validate_draft_transfer(req) 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 @@ -589,12 +581,6 @@ def _received_draft_history(self, req: LlmRequest) -> Optional[dict]: "Received standalone DSpark draft history without a manager-owned draft cache." ) self._validate_draft_history_range(req, history) - if history["layout"] != self._kv_cache_manager.draft_layout.transfer_identity(): - raise ValueError( - "Standalone DSpark draft cache layouts do not match between prefill and " - "generation workers. Matching draft dtype, backend, and per-rank geometry " - "are required." - ) return history def _restore_draft_history(self, req: LlmRequest) -> None: diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 11490c27c12e..f213ee6537ac 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -1748,10 +1748,6 @@ def _validate_standalone_draft_cache(self) -> None: raise ValueError( "Standalone DSpark draft-state transfer requires the PYTHON " "NIXL transceiver on both workers.") - if self._draft_config is None: - raise ValueError( - "Unified standalone DSpark KV cache requires the loaded " - "standalone draft model configuration.") def _get_standalone_draft_layout(self) -> StandaloneDraftLayout: """Describe distinct standalone layers without borrowing target shapes.""" 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 1d1b25c9143e..aa1df652929d 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 @@ -1187,18 +1187,8 @@ def __init__( standalone_draft_layout.extra_tokens if standalone_draft_layout is not None else 0 ) if standalone_draft_layout is not None: - if KV_CACHE_MANAGER_V2_BACKEND != "python": - raise ValueError( - "Unified standalone draft KV requires TLLM_KV_CACHE_MANAGER_V2_BACKEND=python" - ) - if is_draft or mapping.pp_size != 1 or mapping.cp_size != 1: - raise ValueError( - "Unified standalone draft KV requires one target manager and PP=CP=1" - ) if kv_cache_config.enable_block_reuse: raise ValueError("Unified standalone draft KV does not yet support prefix reuse") - if kv_connector_manager is not None: - raise ValueError("Unified standalone draft KV does not yet support KV connectors") if kv_cache_config.enable_swa_scratch_reuse: raise ValueError( "Unified standalone draft KV does not yet support SWA scratch reuse" @@ -3193,8 +3183,6 @@ def get_draft_buffers(self, local_layer_idx: int, kv_layout: str = "HND") -> tor def get_draft_block_table(self, request_ids: List[int]) -> torch.Tensor: """Current rank-local page mappings; unused tail entries point to page zero.""" - if self.draft_layout is None: - raise ValueError("No unified standalone draft cache is configured") layer_id = self.layer_offsets[self.draft_layer_ids[0]] pool_id = self.impl.get_layer_group_id(layer_id) scale = self.impl.get_page_index_scale(layer_id, Role.KEY) @@ -3217,8 +3205,6 @@ 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: - if self.draft_layout is None: - raise ValueError("No unified standalone draft cache is configured") cache = self.kv_cache_map.get(request_id) if cache is None or not cache.is_active: raise ValueError( @@ -3251,10 +3237,6 @@ def restore_draft_history(self, request_id: int, metadata: dict) -> None: raise ValueError("Standalone draft transfer layout does not match the receiving worker") valid_length = metadata.get("valid_length") position = metadata.get("position") - if type(valid_length) is not int or type(position) is not int: - raise ValueError( - "Standalone draft transfer history must use integer lengths and positions" - ) # Accessing the receiving mapping here verifies ownership before the # worker can see this history. The sender's slot/page IDs are never used. self.get_draft_block_table([request_id]) diff --git a/tensorrt_llm/_torch/speculative/dflash.py b/tensorrt_llm/_torch/speculative/dflash.py index 61176860df56..c438bc266782 100644 --- a/tensorrt_llm/_torch/speculative/dflash.py +++ b/tensorrt_llm/_torch/speculative/dflash.py @@ -564,7 +564,7 @@ def _init_ctx_block_tables( return True def _refresh_ctx_block_tables( - self, attn_metadata, num_seqs: int, request_ids: Optional[list[int]] = None + self, attn_metadata, num_seqs: int, request_ids: list[int] ) -> bool: """Decode this iteration's draft block table into the persistent buffer. @@ -575,8 +575,6 @@ def _refresh_ctx_block_tables( if self._ctx_block_tables is None or num_seqs <= 0: return False if self._has_unified_draft_cache(): - if request_ids is None: - raise ValueError("Unified DSpark draft pages require local request IDs") table = self._ctx_kv_manager.get_draft_block_table(request_ids[:num_seqs]) self._ctx_block_tables[:num_seqs].copy_(table, non_blocking=True) return True @@ -724,7 +722,7 @@ def _lazy_init_ctx_buffers( kv_shape = (num_slots, L, capacity, nkv, hd) self._ctx_k_buf = torch.zeros(kv_shape, dtype=dtype, device="cuda") self._ctx_v_buf = torch.zeros(kv_shape, dtype=dtype, device="cuda") - elif self._dflash_attention_backend == "TRTLLM": + else: # TRTLLM; StandaloneDraftLayout validates the backend. validate_dflash_trtllm_gen_runtime( dtype=dtype, num_heads=nh, @@ -735,8 +733,6 @@ def _lazy_init_ctx_buffers( not draft_model._get_attention_mask_args(i)[0] for i in range(L) ), ) - else: - raise ValueError("Unified DSpark supports VANILLA and TRTLLM draft attention") elif self._dflash_attention_backend in _PAGED_ATTENTION_BACKENDS: pool = ( self._managed_ctx_pool(draft_kv_cache_manager, L, nkv, hd, dtype) From 020879bb1d82210320e220f49c2bbae61377890e Mon Sep 17 00:00:00 2001 From: Allison Lim Date: Sat, 19 Sep 2026 00:00:21 -0700 Subject: [PATCH 04/13] [None][refactor] Consolidate standalone DSpark draft-cache handling Reuse paged append and page-index helpers for manager-owned draft state, with destination-dtype conversion for writes. Prepare committed history once before forward and share receive validation while preserving validation before rank consensus and existing cleanup. Signed-off-by: Allison Lim --- .../_torch/disaggregation/transceiver.py | 20 ++--- .../kv_cache/kv_cache_manager_v2.py | 15 ++-- tensorrt_llm/_torch/speculative/dflash.py | 84 ++++++------------- 3 files changed, 42 insertions(+), 77 deletions(-) diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index 58b97e1a03c0..631a515f67b8 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -590,6 +590,12 @@ def _restore_draft_history(self, req: LlmRequest) -> None: # metadata crosses the wire; the manager retains the receiver's request/page map. 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): return True @@ -1083,10 +1089,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._restore_draft_history(req) + self._prepare_received_history(session, req) req.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE else: req.state = LlmRequestState.DISAGG_TRANS_ERROR @@ -1297,12 +1300,10 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]) -> GenT req = self._recv_reqs[rid] if has_draft_history: try: - self._apply_aux(session, req) - self._assert_disagg_history_declared(req) # Restore also validates the receiver's allocation and page map. # A peer failure below leaves this request unschedulable; its # ordinary failure cleanup releases any restored history and KV. - self._restore_draft_history(req) + self._prepare_received_history(session, req) except (ValueError, RuntimeError) as error: logger.warning( f"Disagg draft history validation FAILED rank={self._dist.rank} " @@ -1357,11 +1358,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 not has_draft_history and self._need_aux_transfer(req): - self._apply_aux(session, req) if not has_draft_history: - self._assert_disagg_history_declared(req) - self._restore_draft_history(req) + 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/kv_cache/kv_cache_manager_v2.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py index aa1df652929d..046116ffc91d 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 @@ -3183,22 +3183,19 @@ def get_draft_buffers(self, local_layer_idx: int, kv_layout: str = "HND") -> tor def get_draft_block_table(self, request_ids: List[int]) -> torch.Tensor: """Current rank-local page mappings; unused tail entries point to page zero.""" - layer_id = self.layer_offsets[self.draft_layer_ids[0]] - pool_id = self.impl.get_layer_group_id(layer_id) - scale = self.impl.get_page_index_scale(layer_id, Role.KEY) - table = torch.zeros((len(request_ids), self.max_blocks_per_seq), dtype=torch.int32) - for row, request_id in enumerate(request_ids): + 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") - indices = cache.get_base_page_indices(pool_id)[: cache.num_blocks] + 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") if any(index == BAD_PAGE_INDEX for index in indices): raise ValueError("Standalone full-attention draft cache contains missing pages") - table[row, : len(indices)] = torch.tensor( - [int(index) * int(scale) // 2 for index in indices], dtype=torch.int32 - ) + table[row, : len(indices)] = torch.tensor(indices, dtype=torch.int32) return table def get_draft_history(self, request_id: int) -> Optional[StandaloneDraftHistory]: diff --git a/tensorrt_llm/_torch/speculative/dflash.py b/tensorrt_llm/_torch/speculative/dflash.py index c438bc266782..ae20788ca484 100644 --- a/tensorrt_llm/_torch/speculative/dflash.py +++ b/tensorrt_llm/_torch/speculative/dflash.py @@ -348,15 +348,17 @@ def get_draft_kv_cache_manager( def _has_unified_draft_cache(self) -> bool: return getattr(getattr(self, "_ctx_kv_manager", None), "draft_layout", None) is not None - def _restore_managed_slots(self, request_ids: list[int], num_contexts: int) -> None: - """Bind generation requests after initialization or a manager replacement. + def _prepare_managed_history(self, request_ids: list[int]) -> None: + """Restore each request's committed history before this forward's writes. Metadata preparation precedes lazy buffer binding. That binding can replace an estimation manager and clear its staging slots, so generation must rebind here even when metadata preparation already assigned slots. + Context requests also restore here so continuation and restart detection + use the same authoritative state as generation. """ updates = {self._dummy_slot: 0} - for request_id in request_ids[num_contexts:]: + for request_id in request_ids: if ( request_id != ATTENTION_DP_DUMMY_REQUEST_ID and request_id < self._graph_dummy_id_floor @@ -378,20 +380,6 @@ def _restore_managed_slots(self, request_ids: list[int], num_contexts: int) -> N ) ) - def _store_managed_context_kv( - self, k: torch.Tensor, v: torch.Tensor, rows: torch.Tensor, positions: torch.Tensor - ) -> None: - """Write projected accepted features to the authoritative draft pages. - - K/V have shape [tokens, layers, heads, head_dim]. Rows address this - iteration's request table, independently of the worker's staging slots. - """ - pages = self._ctx_block_tables[rows, positions // self._ctx_page_size].long() - offsets = positions % self._ctx_page_size - for layer_idx, pool in enumerate(self._ctx_kv_buf): - pool[pages, 0, :, offsets, :] = k[:, layer_idx] - pool[pages, 1, :, offsets, :] = v[:, layer_idx] - def _gather_managed_context(self, request_ids: list[int]) -> None: """Refresh VANILLA's dense input staging from manager-owned history. @@ -716,6 +704,13 @@ def _lazy_init_ctx_buffers( dtype=torch.int32, device="cuda", ) + max_blocks = draft_kv_cache_manager.max_blocks_per_seq + self._ctx_block_indptr = torch.arange( + 0, (num_slots + 1) * max_blocks, max_blocks, dtype=torch.int32, device="cuda" + ) + self._ctx_kv_last_page_len = torch.full( + (num_slots,), self._ctx_page_size, dtype=torch.int32, device="cuda" + ) if self._dflash_attention_backend == "VANILLA": # Dense FlashAttention inputs are transient forward staging; # every history read is refreshed from the authoritative pool. @@ -851,22 +846,7 @@ def clear(slot: int) -> None: return None slot = self._free_slots.popleft() self._req_to_slot[req_id] = slot - history = ( - self._ctx_kv_manager.get_draft_history(req_id) - if self._has_unified_draft_cache() and not reset - else None - ) - if history is None: - clear(slot) - else: - # The receiving manager owns validity. A new local staging - # slot must preserve it, regardless of the sender's slot ID. - if updates is None: - self._write_ctx_len({slot: history.valid_length}) - else: - updates[slot] = history.valid_length - self._ctx_len_host[slot] = history.valid_length - self._req_ctx_pos[req_id] = history.position + clear(slot) return self._req_to_slot[req_id] def _get_ctx_paged_append(self) -> Callable[..., None]: @@ -891,12 +871,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, @@ -1029,10 +1010,6 @@ def _store_prefill_context( # request id reused after completion, a prefill restarted after # preemption -- has to start from a clean slot. previous_position = self._req_ctx_pos.get(req_id) - if self._has_unified_draft_cache(): - history = self._ctx_kv_manager.get_draft_history(req_id) - if history is not None: - previous_position = history.position reset = previous_position != first_pos if self._assign_slot(req_id, reset=reset, updates=ctx_len_updates) is None: logger.warning("DFlash: no free slots, skipping context store") @@ -1079,14 +1056,10 @@ def _store_prefill_context( chunk_proj_cast, chunk_pos[:actual] ) # chunk_k/v: [actual, L, nkv, hd] → [L, actual, nkv, hd] - if self._has_unified_draft_cache(): - self._store_managed_context_kv( - chunk_k, - chunk_v, - torch.full((actual,), i, dtype=torch.long, device="cuda"), - torch.arange(cur, end, dtype=torch.long, device="cuda"), - ) - elif self._dflash_attention_backend in _PAGED_ATTENTION_BACKENDS: + if ( + self._has_unified_draft_cache() + or self._dflash_attention_backend in _PAGED_ATTENTION_BACKENDS + ): # Manager block tables are keyed by batch position, the # private arena's by slot. See _ctx_paged_index_args. row = i if self._ctx_block_tables is not None else slot @@ -1161,7 +1134,7 @@ def _forward_impl( ) spec_metadata._dflash_worker = self if self._has_unified_draft_cache(): - self._restore_managed_slots(spec_metadata.request_ids, num_contexts) + self._prepare_managed_history(spec_metadata.request_ids) # Before any store: prefill and decode both address pages through it. self._refresh_ctx_block_tables(attn_metadata, batch_size, spec_metadata.request_ids) if self._has_unified_draft_cache(): @@ -1673,13 +1646,10 @@ def prepare_1st_drafter_inputs( v_new.mul_(mask_bc) slot_long = slot_flat.long() col_long = col_flat.long() - if self._has_unified_draft_cache(): - rows_long = gen_rows_out.unsqueeze(1).expand(-1, K + 1).reshape(-1) - self._store_managed_context_kv(k_new, v_new, rows_long, col_long) - if self._dflash_attention_backend == "VANILLA": - self._ctx_k_buf[slot_long, :, col_long] = k_new - self._ctx_v_buf[slot_long, :, col_long] = v_new - elif self._dflash_attention_backend in _PAGED_ATTENTION_BACKENDS: + if ( + self._has_unified_draft_cache() + or self._dflash_attention_backend in _PAGED_ATTENTION_BACKENDS + ): if self._ctx_block_tables is not None: # Batch positions of the gen requests, matching the # per-request block table's row order. @@ -1687,7 +1657,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._dflash_attention_backend == "VANILLA": self._ctx_k_buf[slot_long, :, col_long] = k_new self._ctx_v_buf[slot_long, :, col_long] = v_new From 71509562eccf577ae096f7e83312951ca1f64846 Mon Sep 17 00:00:00 2001 From: Allison Lim Date: Sun, 20 Sep 2026 22:44:42 -0700 Subject: [PATCH 05/13] [None][fix] Preserve DSpark draft history with default KVCache V2 Separate target and draft cache domains in the C++ backend and extend unified ownership to embedded DeepSeek rolling-window state. Reuse manager allocation and disaggregation transport to restore receiver-local draft pages and committed history before drafting. Account for independent draft geometry and window retention, validate received state before cross-rank completion, and preserve legacy aggregate configurations. Validated default C++ and Python paths, Qwen and DeepSeek state continuity and matched AL workloads, and DeepSeek GSM8K accuracy. Kimi end-to-end hardware validation remains outstanding. Signed-off-by: Allison Lim --- .../kv_cache_manager_v2/config.h | 3 + .../kv_cache_manager_v2/lifeCycleRegistry.cpp | 6 +- .../kv_cache_manager_v2/lifeCycleRegistry.h | 17 ++- .../kv_cache_manager_v2/storage/config.cpp | 12 +- .../kv_cache_manager_v2/storageManager.cpp | 8 +- .../batch_manager/kvCacheManagerV2.cpp | 12 +- .../sparse/deepseek_v4/cache_manager.py | 99 +++++++++++++++ .../_torch/disaggregation/native/auxiliary.py | 38 ++++-- .../disaggregation/resource/kv_extractor.py | 8 +- .../_torch/disaggregation/transceiver.py | 114 ++++++++++------- tensorrt_llm/_torch/pyexecutor/_util.py | 106 +++++++++------- .../kv_cache/kv_cache_manager_v2.py | 115 ++++++++++++----- .../kv_cache/standalone_draft_cache.py | 26 +++- tensorrt_llm/_torch/speculative/dspark.py | 116 +++++++++++++++++- 14 files changed, 523 insertions(+), 157 deletions(-) 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..e2b105898798 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 @@ -35,9 +35,13 @@ LifeCycle makeLifeCycle(LayerConfig const& layer, int tokensPerBlock) cfg.validate(); using T = std::decay_t; 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..bcfc8765f8ab 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,7 @@ #include #include #include +#include #include #include #include @@ -40,6 +41,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 +58,31 @@ 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; + } + if (numSinkBlocks != o.numSinkBlocks) + { + return numSinkBlocks < o.numSinkBlocks; + } + return cacheDomain < 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..17b18c36330a 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,22 @@ 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; + // Equal storage sizes permit merging only within a compatible ownership domain. + // Existing target attention/SSM groups retain their shared physical pool behavior. + 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 a13b58476b19..d4d9786540ac 100644 --- a/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp +++ b/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManagerV2.cpp @@ -1338,13 +1338,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==); @@ -1542,13 +1544,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 a367efdd011b..8491dd6aabd4 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,26 @@ NVFP4_VECTOR_SIZE = 16 +def _draft_cache_size_components( + layout: StandaloneDraftLayout | None, + 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.""" + if layout is None: + return 0, 0, 0 + 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: @@ -956,6 +979,16 @@ 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, ) + 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, + non_sliding_attn_size_per_token + context_swa_size_per_token, + ) + context_swa_size_per_token += draft_context + generation_swa_size_per_token += draft_generation + generation_swa_size_per_request += draft_per_request max_context_tokens = ( self._max_num_tokens if self._max_num_tokens is not None else max_tokens ) @@ -1009,6 +1042,16 @@ def _get_max_tokens_from_quota(self, quota: int) -> float: indexer_k_dtype=self._indexer_k_dtype, use_fp8_ds_mla=self.use_fp8_ds_mla, ) + 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, + non_sliding_attn_size_per_token + context_swa_size_per_token, + ) + context_swa_size_per_token += draft_context + generation_swa_size_per_token += draft_generation + generation_swa_size_per_request += draft_per_request padding = self._get_extra_quota_padding() size_per_batch = self.max_batch_size * generation_swa_size_per_request + padding if quota < size_per_batch: @@ -1178,6 +1221,13 @@ def _add_layer( layers=layers, ) + def _append_standalone_draft_layers( + self, config: KVCacheManagerConfigPy + ) -> KVCacheManagerConfigPy: + # Target model indices describe several virtual attention layers each. + # Draft layers are independent config entries, not extra target 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 @@ -1383,9 +1433,33 @@ def _get_cache_bytes_for_tokens(self, total_tokens: int, *, context: bool) -> in indexer_k_dtype=self._indexer_k_dtype, use_fp8_ds_mla=self.use_fp8_ds_mla, ) + draft_context, draft_generation, draft_per_request = ( + _draft_cache_size_components( + self.draft_layout, + self.tokens_per_block, + self._generation_kv_capacity_headroom, + non_sliding_attn_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=False, + indexer_k_dtype=self._indexer_k_dtype, + use_fp8_ds_mla=self.use_fp8_ds_mla, + )[0], + ) + if self.draft_layout is not None + else (0, 0, 0) + ) return int( total_tokens * (non_sliding_attn_size_per_token + swa_size_per_token) + swa_size_per_request + + total_tokens * (draft_context if context else draft_generation) + + draft_per_request ) def get_needed_resource_to_completion(self, request: llm_request.LlmRequest) -> int: @@ -1400,6 +1474,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 @@ -1674,6 +1750,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 e7faca0d27bd..debe6071e872 100644 --- a/tensorrt_llm/_torch/disaggregation/native/auxiliary.py +++ b/tensorrt_llm/_torch/disaggregation/native/auxiliary.py @@ -114,10 +114,10 @@ def build_aux_transfer_layout( AuxSlot = namedtuple("AuxSlot", ["id", "buffer"]) -_DRAFT_HISTORY_VERSION = 1 -_DRAFT_HISTORY_FIELDS = 8 +_DRAFT_HISTORY_VERSION = 2 +_DRAFT_HISTORY_FIELDS = 10 _DRAFT_DTYPE_CODES = {"torch.float16": 1, "torch.bfloat16": 2} -_DRAFT_BACKEND_CODES = {"VANILLA": 1, "TRTLLM": 2} +_DRAFT_BACKEND_CODES = {"VANILLA": 1, "TRTLLM": 2, "DSv4": 3} def _encode_draft_history(history: dict[str, Any]) -> list[int]: @@ -134,27 +134,43 @@ def _encode_draft_history(history: dict[str, Any]) -> list[int]: layout["head_dim"], _DRAFT_DTYPE_CODES[layout["dtype"]], _DRAFT_BACKEND_CODES[layout["attention_backend"]], + layout.get("kv_factor", 2), + layout.get("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 = values + ( + _, + 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()} # The receiving transceiver validates prompt coverage and the manager checks # storage identity and local allocation before publishing this history. + 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), + } + if kv_factor != 2 or window_size != 0: + layout.update(kv_factor=kv_factor, window_size=window_size or None) return { "valid_length": valid_length, "position": position, - "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), - }, + "layout": layout, } diff --git a/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py b/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py index a5fb8f56c061..a50812d2e3a1 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,11 @@ 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 getattr( + manager, "draft_layout", None + ) is not None and 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 631a515f67b8..0df6dc7cb639 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -428,47 +428,65 @@ def _describe_local(self, req: LlmRequest) -> Chunk: else np.array([], dtype=np.int64) ) continue - block_ids = adapter.get_block_ids(req, idx, lg) window_size = lg.sliding_window_size if window_size is not None: - draft_len = get_draft_token_length(req) if is_gen_only else 0 - allocated_blocks = ( - req.prompt_len - + draft_len - + self._kv_cache_manager.num_extra_kv_tokens - + tpb - - 1 - ) // tpb - if block_ids.size > allocated_blocks: - block_ids = block_ids[:allocated_blocks] - # Current PyExecutor cache managers disable KV-cache token sinks, - # so SWA block lists contain an evictable prompt prefix followed - # by the speculative scratch tail. If token sinks are enabled, - # this must use block-ordinal metadata to preserve the sink prefix. - # Remove scratch before trimming stale prompt blocks; otherwise a - # boundary-crossing allocation can displace initialized prompt KV. - scratch_blocks = max(0, allocated_blocks - prompt_blocks) - if scratch_blocks > 0: - if req.py_beam_width != 1: - raise ValueError("speculative scratch blocks require beam_width == 1") - block_ids = ( - block_ids[:-scratch_blocks] - if scratch_blocks < block_ids.size - else np.array([], dtype=np.int64) - ) - # Drop stale blocks the manager may still expose (V1 pre-eviction). stale_end = max(0, (req.prompt_len + 1 - window_size) // tpb) - expected_valid = max(0, prompt_blocks - stale_end) - if block_ids.size > expected_valid: - block_ids = ( - block_ids[-expected_valid:] - if expected_valid > 0 - else np.array([], dtype=np.int64) + if ( + isinstance(self._kv_cache_manager, KVCacheManagerV2) + and self._kv_cache_manager.draft_layout is not None + ): + # Preserve logical ordinals until both the stale prefix and + # speculative tail have been excluded. Filtering holes first + # loses their positions and can trim valid prompt pages. + cache = self._kv_cache_manager.kv_cache_map[req.py_request_id] + pages = np.fromiter( + cache.get_aggregated_page_indices(idx, valid_only=False), + dtype=np.int64, ) + block_ids = pages[stale_end:prompt_blocks] + if block_ids.size != prompt_blocks - stale_end or np.any(block_ids < 0): + raise ValueError("Missing allocated prompt pages for windowed KV transfer") + else: + block_ids = adapter.get_block_ids(req, idx, lg) + draft_len = get_draft_token_length(req) if is_gen_only else 0 + allocated_blocks = ( + req.prompt_len + + draft_len + + self._kv_cache_manager.num_extra_kv_tokens + + tpb + - 1 + ) // tpb + if block_ids.size > allocated_blocks: + block_ids = block_ids[:allocated_blocks] + # Current PyExecutor cache managers disable KV-cache token sinks, + # so SWA block lists contain an evictable prompt prefix followed + # by the speculative scratch tail. If token sinks are enabled, + # this must use block-ordinal metadata to preserve the sink prefix. + # Remove scratch before trimming stale prompt blocks; otherwise a + # boundary-crossing allocation can displace initialized prompt KV. + scratch_blocks = max(0, allocated_blocks - prompt_blocks) + if scratch_blocks > 0: + if req.py_beam_width != 1: + raise ValueError("speculative scratch blocks require beam_width == 1") + block_ids = ( + block_ids[:-scratch_blocks] + if scratch_blocks < block_ids.size + else np.array([], dtype=np.int64) + ) + # Drop stale blocks the manager may still expose (V1 pre-eviction). + stale_end = max(0, (req.prompt_len + 1 - window_size) // tpb) + expected_valid = max(0, prompt_blocks - stale_end) + if block_ids.size > expected_valid: + block_ids = ( + block_ids[-expected_valid:] + if expected_valid > 0 + else np.array([], dtype=np.int64) + ) # Skip reused blocks that remain after stale-prefix pruning. cache_skip = max(0, cached_per_lg[idx] // tpb - stale_end) else: + block_ids = adapter.get_block_ids(req, idx, lg) # Drop the speculative scratch tail; only prompt_len is transferred. if block_ids.size > prompt_blocks: block_ids = block_ids[:prompt_blocks] @@ -537,22 +555,32 @@ def _validate_draft_transfer(self, req: LlmRequest) -> None: params = req.py_disaggregated_params if params is not None and params.schedule_style == DisaggScheduleStyle.GENERATION_FIRST: raise ValueError( - "Standalone DSpark draft-state transfer requires context_first scheduling; " - "generation_first is not yet supported for standalone draft history." + "DSpark draft-state transfer requires context_first scheduling; " + "generation_first is not yet supported for draft history." ) if self.pipeline_transfer_enabled: - raise ValueError( - "Standalone DSpark draft-state transfer does not support pipelined transfer." - ) + raise ValueError("DSpark draft-state transfer does not support pipelined transfer.") @staticmethod def _validate_draft_history_range(req: LlmRequest, history: dict) -> None: - # The shared transfer extent covers the complete prompt. Draft noise KV is scratch, - # never valid history. A partial history needs per-group extents before it can be sent. - if history["valid_length"] != req.prompt_len or history["position"] != req.prompt_len: + # Full-attention drafters retain the whole prompt; rolling drafters + # retain its complete live suffix. Neither includes speculative scratch. + window_size = history["layout"].get("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 ( + type(history["valid_length"]) is not int + or type(history["position"]) is not int + or history["valid_length"] != expected_length + or history["position"] != req.prompt_len + ): raise ValueError( - "Standalone DSpark transfer requires valid draft history and sequence position " - f"covering the complete prompt ({req.prompt_len} tokens)." + "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) -> Optional[dict]: diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index f213ee6537ac..69305d00b5fe 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -1674,83 +1674,103 @@ def _is_standalone_dspark(self) -> bool: 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 not self._is_kv_cache_manager_v2: + if (not (self._is_standalone_dspark() or self._is_embedded_dspark()) + or not self._is_kv_cache_manager_v2): return False if self._is_disagg: # Disaggregation must validate unified ownership, never fall back # to draft state that the transceiver cannot transfer. return True - from tensorrt_llm.runtime.kv_cache_manager_v2 import BACKEND - - # Preserve the existing aggregate C++/CUDA-graph draft-cache path. - # Python V2 with eager execution also supports unified aggregate runs - # for comparison with the same configuration in disaggregation. - return BACKEND == "python" and self._llm_args.cuda_graph_config is None + # Keep existing aggregate execution for settings the unified draft + # lifecycle cannot support. Disaggregation must never take this path. + return self._unified_draft_cache_unsupported_reason() is None def _validate_standalone_draft_cache(self) -> None: - """Reject unsupported standalone state ownership before profiling.""" - if not self._is_standalone_dspark(): + """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( - "Standalone DSpark requested KVCacheManagerV2 but its " + "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( - "Standalone DSpark disaggregation requires " - "kv_cache_config.use_kv_cache_manager_v2=True and " - "TLLM_KV_CACHE_MANAGER_V2_BACKEND=python on both workers.") - return - if not self._uses_unified_standalone_draft_cache(): + raise ValueError("DSpark disaggregation requires " + "kv_cache_config.use_kv_cache_manager_v2=True " + "on both workers.") return - from tensorrt_llm.runtime.kv_cache_manager_v2 import BACKEND - - if BACKEND != "python": - raise ValueError("Unified standalone DSpark KV cache requires " - "TLLM_KV_CACHE_MANAGER_V2_BACKEND=python.") + 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.""" + if self._kv_cache_config.enable_block_reuse: + return "Unified DSpark draft KV does not yet support prefix reuse" + if self._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 self._kv_cache_config.pool_ratio is not None: + return "Unified DSpark draft KV does not yet support explicit pool_ratio" if (self._speculative_config.draft_len_schedule is not None or self._speculative_config.max_concurrency is not None): - raise ValueError( - "Unified standalone DSpark KV cache does not yet support " + 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._llm_args.cuda_graph_config is not None: - raise ValueError( - "Unified standalone DSpark KV cache currently requires eager " - "execution; set cuda_graph_config=None.") + return ("Unified DSpark KV cache currently requires eager " + "execution; set cuda_graph_config=None.") if not self._disable_overlap_scheduler: - raise ValueError("Unified standalone DSpark KV cache requires " - "disable_overlap_scheduler=True.") + return ("Unified DSpark KV cache requires " + "disable_overlap_scheduler=True.") if self._llm_args.enable_chunked_prefill: - raise ValueError( - "Unified standalone DSpark KV cache does not yet support " - "chunked prefill; set enable_chunked_prefill=False.") + return ("Unified DSpark KV cache does not yet support " + "chunked prefill; set enable_chunked_prefill=False.") if self._mapping.pp_size != 1 or self._mapping.cp_size != 1: - raise ValueError( - "Unified standalone DSpark KV cache requires PP=1 and CP=1.") - if self._mapping.enable_attention_dp: - raise ValueError( - "Unified standalone DSpark KV cache does not yet support " + return "Unified DSpark KV cache requires PP=1 and CP=1." + if (self._mapping.enable_attention_dp + and not self._is_embedded_dspark()): + return ( + "Unified DSpark KV cache does not yet support " "attention data parallelism; set enable_attention_dp=False.") if self._kv_connector_manager is not None: - raise ValueError( - "Unified standalone DSpark KV cache does not yet support " - "KV cache connectors.") + 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"): - raise ValueError( - "Standalone DSpark draft-state transfer requires the PYTHON " - "NIXL transceiver on both workers.") + return ("DSpark draft-state transfer requires the PYTHON " + "NIXL transceiver on both workers.") + return None def _get_standalone_draft_layout(self) -> StandaloneDraftLayout: - """Describe distinct standalone layers without borrowing target shapes.""" + """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) 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 046116ffc91d..6f8f9c7d6e01 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 @@ -1183,6 +1183,7 @@ def __init__( 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 ) @@ -1191,7 +1192,9 @@ def __init__( raise ValueError("Unified standalone draft KV does not yet support prefix reuse") if kv_cache_config.enable_swa_scratch_reuse: raise ValueError( - "Unified standalone draft KV does not yet support SWA scratch reuse" + "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: raise ValueError( @@ -2077,7 +2080,7 @@ def _get_pool_roles(self, pool_id: int) -> Tuple[DataRole, Optional[DataRole]]: 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 + 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 @@ -2328,7 +2331,9 @@ def _get_runtime_cache_size_layer_components(self) -> tuple[List[int], List[Opti layer_sizes.extend( [self.draft_layout.bytes_per_layer_token] * self.draft_layout.num_layers ) - attention_windows.extend([None] * 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_max_tokens_from_quota(self, quota: int) -> float: @@ -2892,7 +2897,7 @@ def _build_cache_config(self, config: KVCacheManagerConfigPy) -> KVCacheManagerC return config def _append_standalone_draft_layers( - self, config: KVCacheManagerConfigPy + self, config: KVCacheManagerConfigPy, *, register_model_layers: bool = True ) -> KVCacheManagerConfigPy: layout = self.draft_layout if layout is None: @@ -2909,21 +2914,26 @@ def _append_standalone_draft_layers( buffers=[ BufferConfig( role=role, - size=layout.bytes_per_layer_token // 2 * self.tokens_per_block, + size=layout.bytes_per_layer_token + // layout.kv_factor + * self.tokens_per_block, ) - for role in (Role.KEY, Role.VALUE) + for role in (Role.KEY, Role.VALUE)[: layout.kv_factor] ], + sliding_window_size=layout.retention_window_size, cache_domain="standalone_draft", ) ) - self.pp_layers.append(global_id) self.layer_offsets[global_id] = local_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(None) - self.num_local_layers = len(self.pp_layers) - self.num_layers += layout.num_layers + 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( @@ -3134,9 +3144,8 @@ def get_buffers(self, layer_idx: int, kv_layout: str = "NHD") -> Optional[torch. ) def _is_standalone_draft_layer(self, local_layer_idx: int) -> bool: - return ( - getattr(self, "draft_layout", None) is not None - and self.pp_layers[local_layer_idx] in self.draft_layer_ids + return getattr(self, "draft_layout", None) 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: @@ -3150,11 +3159,13 @@ def get_layer_kv_factor(self, layer_idx: int) -> int: if self.draft_layout is None: return self.kv_factor return ( - 2 if self._is_standalone_draft_layer(self.layer_offsets[layer_idx]) else self.kv_factor + self.draft_layout.kv_factor + if self._is_standalone_draft_layer(self.layer_offsets[layer_idx]) + else self.kv_factor ) def get_draft_buffers(self, local_layer_idx: int, kv_layout: str = "HND") -> torch.Tensor: - """View authoritative draft K/V pages using the standalone geometry.""" + """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") @@ -3162,12 +3173,15 @@ def get_draft_buffers(self, local_layer_idx: int, kv_layout: str = "HND") -> tor raise ValueError(f"Unsupported standalone draft KV layout: {kv_layout}") layer_id = self.layer_offsets[self.draft_layer_ids[local_layer_idx]] key_address = self.impl.get_mem_pool_base_address(layer_id, Role.KEY, PageIndexMode.SHARED) - value_address = self.impl.get_mem_pool_base_address( - layer_id, Role.VALUE, PageIndexMode.SHARED - ) - stride = self.impl.get_page_stride(layer_id, Role.KEY) - if value_address != key_address + stride: - raise ValueError("Standalone draft K/V buffers must have adjacent equal-sized pages") + if layout.kv_factor == 2: + value_address = self.impl.get_mem_pool_base_address( + layer_id, Role.VALUE, PageIndexMode.SHARED + ) + stride = self.impl.get_page_stride(layer_id, Role.KEY) + if value_address != key_address + stride: + raise ValueError( + "Standalone draft K/V buffers must have adjacent equal-sized pages" + ) dimensions = ( [layout.num_kv_heads, self.tokens_per_block, layout.head_dim] if kv_layout == "HND" @@ -3177,11 +3191,19 @@ def get_draft_buffers(self, local_layer_idx: int, kv_layout: str = "HND") -> tor TensorWrapper( key_address, self.get_layer_cache_dtype(self.draft_layer_ids[local_layer_idx]), - [self.impl.get_page_index_upper_bound(layer_id, Role.KEY) // 2, 2, *dimensions], + [ + self.impl.get_page_index_upper_bound(layer_id, Role.KEY) // layout.kv_factor, + layout.kv_factor, + *dimensions, + ], ) ) - def get_draft_block_table(self, request_ids: List[int]) -> torch.Tensor: + 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) @@ -3193,8 +3215,23 @@ def get_draft_block_table(self, request_ids: List[int]) -> torch.Tensor: 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") - if any(index == BAD_PAGE_INDEX for index in indices): - raise ValueError("Standalone full-attention draft cache contains missing pages") + 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 @@ -3208,8 +3245,13 @@ def set_draft_history(self, request_id: int, valid_length: int, position: int) - f"Standalone draft request {request_id} has no active cache allocation" ) history = StandaloneDraftHistory(valid_length, position) - if history.valid_length > cache.capacity: + if history.valid_length > cache.capacity or 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]: @@ -3234,9 +3276,10 @@ def restore_draft_history(self, request_id: int, metadata: dict) -> None: 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) # Accessing the receiving mapping here verifies ownership before the # worker can see this history. The sender's slot/page IDs are never used. - self.get_draft_block_table([request_id]) + self.get_draft_block_table([request_id], [history]) self.set_draft_history(request_id, valid_length, position) def get_index_k_buffer( @@ -5263,6 +5306,7 @@ def free_resources(self, request: LlmRequest, pin_on_release: bool = False): 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. @@ -5419,7 +5463,9 @@ def get_layer_bytes_per_token(self, local_layer_idx: int, data_role: Role): if data_role == Role.ALL: return self.draft_layout.bytes_per_layer_token if data_role in (Role.KEY, Role.VALUE): - return self.draft_layout.bytes_per_layer_token // 2 + 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, @@ -5533,6 +5579,7 @@ def shutdown(self): self.kv_cache_map.clear() if self.draft_layout is not None: self.draft_history.clear() + self._draft_dummy_request_ids.clear() self._request_stats_enabled_ids.clear() self._fresh_pages_filled.clear() # Drop the outstanding plans before the manager shuts down: discarding a handle applies @@ -5613,7 +5660,7 @@ def get_cache_size_per_token( 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([None] * 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 ( @@ -5752,7 +5799,7 @@ def update_resources( # Rejected target verification slots do not revoke draft # feature history. Retain it and its next-forward scratch. new_capacity = max( - new_capacity, draft_history.valid_length + self._standalone_draft_reserve + new_capacity, draft_history.position + self._standalone_draft_reserve ) history_length = ( None @@ -5854,6 +5901,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/standalone_draft_cache.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/standalone_draft_cache.py index 6e5f92d7cedc..a4ad2586fc68 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache/standalone_draft_cache.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/standalone_draft_cache.py @@ -1,7 +1,7 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""Storage and history contracts for standalone DSpark/DFlash drafters.""" +"""Storage and history contracts for manager-owned DSpark/DFlash drafters.""" from dataclasses import dataclass @@ -10,7 +10,7 @@ @dataclass(frozen=True) class StandaloneDraftLayout: - """Rank-local full-attention storage, independent of target KV geometry.""" + """Rank-local draft storage, independent of target geometry and retention.""" num_layers: int num_kv_heads: int @@ -18,6 +18,8 @@ class StandaloneDraftLayout: 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: @@ -26,25 +28,37 @@ def __post_init__(self) -> None: 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"): - raise ValueError("Standalone draft KV requires VANILLA or TRTLLM attention") + 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 2 * self.num_kv_heads * self.head_dim * self.dtype.itemsize + 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 { + identity = { "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, } + if self.kv_factor != 2 or self.window_size is not None: + identity.update(kv_factor=self.kv_factor, window_size=self.window_size) + return identity @dataclass(frozen=True) diff --git a/tensorrt_llm/_torch/speculative/dspark.py b/tensorrt_llm/_torch/speculative/dspark.py index 808ea8840e2e..1b4942468d28 100644 --- a/tensorrt_llm/_torch/speculative/dspark.py +++ b/tensorrt_llm/_torch/speculative/dspark.py @@ -30,6 +30,7 @@ from tensorrt_llm.mapping import Mapping 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 @@ -195,9 +196,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 @@ -240,6 +242,9 @@ 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 # Slot management. ``_req_to_slot`` (python dict) + ``_free_slots`` are the # source of truth, updated in prepare()/forward(); ``_batch_to_slot`` is the @@ -354,6 +359,8 @@ def _lazy_init(self, draft_model, spec_metadata) -> 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._ctx_len[old] = 0 @@ -375,6 +382,104 @@ def _assign_slot(self, req_id: int, reset: bool) -> int: self._kv_windows[slot].zero_() return self._req_to_slot[req_id] + 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 + + 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 authoritative history once, before any local accepted-token writes.""" + # 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] + table = self._draft_kv_manager.get_draft_block_table(real_ids) + self._draft_block_tables = torch.zeros( + (len(request_ids), table.shape[1]), dtype=table.dtype, device=self._kv_windows.device + ) + self._draft_block_tables[real_rows] = table.to(self._kv_windows.device) + 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] + 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" + ) + slot = self._assign_slot(request_id, reset=False) + batch_slots[row] = slot + 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._batch_to_slot[: len(request_ids)].copy_( + torch.tensor(batch_slots, dtype=torch.long, device=self._batch_to_slot.device) + ) + + def _publish_managed_history(self, request_ids: list[int]) -> None: + """Commit successful prefill/accepted-feature writes; proposal KV stays scratch.""" + lengths = self._valid_len.tolist() + positions = self._ctx_len.tolist() + for row, request_id in enumerate(request_ids): + if not self._is_managed_request(request_id): + continue + slot = self._req_to_slot[request_id] + length, position = lengths[slot], positions[slot] + if position > self._draft_kv_manager.kv_cache_map[request_id].capacity: + raise ValueError("Embedded DSpark history exceeds its allocated draft capacity") + pages, offsets, frames = self._managed_window_indices(row, position, length) + for stage, pool in enumerate(self._draft_kv_buffers): + pool[pages, 0, 0, offsets] = self._kv_windows[slot, stage, frames] + self._draft_kv_manager.set_draft_history(request_id, length, position) + def _seed_context_windows( self, draft_model, @@ -583,6 +688,9 @@ def _forward_impl( # 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 + self._bind_managed_history(resource_manager) + if self._draft_kv_manager is not None: + self._prepare_managed_history(spec_metadata.request_ids, num_contexts) self._execute_guided_decoder_if_present(logits) # Target-verify acceptance via the unified SpecWorkerBase entry: it @@ -726,6 +834,8 @@ def _forward_impl( self._valid_len.copy_(saved_valid_len) self._position_initialized.copy_(saved_position_initialized) self._kv_windows.copy_(saved_windows) + elif self._draft_kv_manager is not None: + self._publish_managed_history(spec_metadata.request_ids) return { "logits": raw_logits, From 2c09cb3fa6600f867662e355f7747fe7f62326b3 Mon Sep 17 00:00:00 2001 From: Allison Lim Date: Mon, 21 Sep 2026 11:37:23 -0700 Subject: [PATCH 06/13] [None][refactor] Minimize unified DSpark KV-cache implementation Canonicalize transferred draft layout metadata, consolidate history restoration and shared cache restrictions, and remove redundant checks and initialization. Retain receiver validation before rank consensus, lifecycle cleanup, independent draft geometry and VANILLA staging. Remove this feature's Python KVCM configuration, lifecycle and storage additions. Keep native C++ ownership and transfer support without changing backend selection or adding a C++ admission guard. Update an existing extractor test fixture for the simplified V2 predicate. Validation: default and explicit C++ each passed 660 focused tests with one pre-existing skip; Python target-only compatibility passed 26 tests. Formatting and diff checks passed. New local regressions and experiment artifacts are not included in this commit. Signed-off-by: Allison Lim --- .../kv_cache_manager_v2/storage/config.cpp | 3 +- .../sparse/deepseek_v4/cache_manager.py | 7 +- .../_torch/disaggregation/native/auxiliary.py | 13 ++- .../_torch/disaggregation/native/transfer.py | 3 +- .../disaggregation/resource/kv_extractor.py | 4 +- .../_torch/disaggregation/transceiver.py | 42 +++------ tensorrt_llm/_torch/pyexecutor/_util.py | 85 +++++++++---------- .../kv_cache/kv_cache_manager_v2.py | 46 +++++----- .../kv_cache/standalone_draft_cache.py | 13 +-- tensorrt_llm/_torch/speculative/dflash.py | 39 +++------ .../runtime/kv_cache_manager_v2/_config.py | 6 -- .../_life_cycle_registry.py | 12 +-- .../kv_cache_manager_v2/_storage/_config.py | 21 ++--- .../unittest/disaggregated/test_extractor.py | 1 + 14 files changed, 108 insertions(+), 187 deletions(-) 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 17b18c36330a..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,8 +204,7 @@ StorageConfig createStorageConfig(KVCacheManagerConfig const& config) slotGroups.push_back(std::move(var)); } - // Equal storage sizes permit merging only within a compatible ownership domain. - // Existing target attention/SSM groups retain their shared physical pool behavior. + // Merge equal storage sizes only within the same cache domain. std::map>, std::vector> poolGroupsByLayout; for (auto& sg : slotGroups) { 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 8491dd6aabd4..1fc5107f8b86 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 @@ -72,14 +72,12 @@ def _draft_cache_size_components( - layout: StandaloneDraftLayout | None, + 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.""" - if layout is None: - return 0, 0, 0 context, generation, per_request = _estimate_cache_size_components( [layout.bytes_per_layer_token] * layout.num_layers, [layout.retention_window_size] * layout.num_layers, @@ -1224,8 +1222,7 @@ def _add_layer( def _append_standalone_draft_layers( self, config: KVCacheManagerConfigPy ) -> KVCacheManagerConfigPy: - # Target model indices describe several virtual attention layers each. - # Draft layers are independent config entries, not extra target layers. + # 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: diff --git a/tensorrt_llm/_torch/disaggregation/native/auxiliary.py b/tensorrt_llm/_torch/disaggregation/native/auxiliary.py index debe6071e872..6b847d2a0336 100644 --- a/tensorrt_llm/_torch/disaggregation/native/auxiliary.py +++ b/tensorrt_llm/_torch/disaggregation/native/auxiliary.py @@ -134,8 +134,8 @@ def _encode_draft_history(history: dict[str, Any]) -> list[int]: layout["head_dim"], _DRAFT_DTYPE_CODES[layout["dtype"]], _DRAFT_BACKEND_CODES[layout["attention_backend"]], - layout.get("kv_factor", 2), - layout.get("window_size") or 0, + layout["kv_factor"], + layout["window_size"] or 0, ] @@ -156,17 +156,15 @@ def _decode_draft_history(values: list[int]) -> dict[str, Any]: ) = values dtypes = {code: name for name, code in _DRAFT_DTYPE_CODES.items()} backends = {code: name for name, code in _DRAFT_BACKEND_CODES.items()} - # The receiving transceiver validates prompt coverage and the manager checks - # storage identity and local allocation before publishing this history. 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, } - if kv_factor != 2 or window_size != 0: - layout.update(kv_factor=kv_factor, window_size=window_size or None) return { "valid_length": valid_length, "position": position, @@ -265,8 +263,7 @@ def __init__( self._prompt_token_counts_buffer = torch.zeros( self._max_slot_num, 2, dtype=data_type, device=self._device ) - # This participates in the existing auxiliary memory registration and - # transfer. Version zero denotes an unfilled or newly allocated slot. + # 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 diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index e5a9a5f18f96..fcea77bb3c2c 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -1678,8 +1678,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() and send_aux(), including context-first auxiliary submission. + # 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) diff --git a/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py b/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py index a50812d2e3a1..12c660c7f698 100644 --- a/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py +++ b/tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py @@ -742,9 +742,7 @@ 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 getattr( - manager, "draft_layout", None - ) is not None and manager._is_standalone_draft_layer(first_local_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] diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index 0df6dc7cb639..fd3f11466521 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -475,7 +475,6 @@ def _describe_local(self, req: LlmRequest) -> Chunk: else np.array([], dtype=np.int64) ) # Drop stale blocks the manager may still expose (V1 pre-eviction). - stale_end = max(0, (req.prompt_len + 1 - window_size) // tpb) expected_valid = max(0, prompt_blocks - stale_end) if block_ids.size > expected_valid: block_ids = ( @@ -543,8 +542,7 @@ def _chunk_num_bytes(self, chunk: Chunk) -> int: def _need_aux_transfer(self, req: LlmRequest) -> bool: params = req.py_disaggregated_params - manager = getattr(self, "_kv_cache_manager", None) - return getattr(manager, "draft_layout", None) is not None or ( + return getattr(self._kv_cache_manager, "draft_layout", None) is not None or ( params is not None and params.schedule_style == DisaggScheduleStyle.GENERATION_FIRST ) @@ -565,35 +563,29 @@ def _validate_draft_transfer(self, req: LlmRequest) -> None: 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"].get("window_size") + 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 ( - type(history["valid_length"]) is not int - or type(history["position"]) is not int - or history["valid_length"] != expected_length - or history["position"] != req.prompt_len - ): + 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) -> Optional[dict]: + def _pack_draft_history(self, req: LlmRequest) -> None: manager = getattr(self, "_kv_cache_manager", None) if getattr(manager, "draft_layout", None) is None: - return 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 - return history - def _received_draft_history(self, req: LlmRequest) -> Optional[dict]: + 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 @@ -603,20 +595,14 @@ def _received_draft_history(self, req: LlmRequest) -> Optional[dict]: "Standalone DSpark generation requires draft history from a prefill worker " "with matching speculative configuration; draft history metadata is missing." ) - return None + 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) - return history - - def _restore_draft_history(self, req: LlmRequest) -> None: - history = self._received_draft_history(req) - if history is not None: - # K/V is already in this request's local pages. Only portable validity/position - # metadata crosses the wire; the manager retains the receiver's request/page map. - self._kv_cache_manager.restore_draft_history(req.py_request_id, 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): @@ -909,8 +895,7 @@ 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 already carries tokens and usage in the context response. The - # existing registered auxiliary transfer carries only the additional draft state. + # Context-first tokens and usage already arrived in the context response. session.unpack_draft_history(req) return session.unpack_aux(req) @@ -921,7 +906,7 @@ def _apply_aux(self, session, req: LlmRequest): req.context_phase_params = ContextPhaseParams( first_gen_tokens=first_gen_tokens, req_id=req.py_request_id, - opaque_state=None, + opaque_state=b"", draft_tokens=draft_tokens, ctx_dp_rank=0, disagg_info_endpoint="", @@ -1328,9 +1313,8 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]) -> GenT req = self._recv_reqs[rid] if has_draft_history: try: - # Restore also validates the receiver's allocation and page map. - # A peer failure below leaves this request unschedulable; its - # ordinary failure cleanup releases any restored history and KV. + # 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( diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 69305d00b5fe..df5229a7e1ee 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -61,7 +61,8 @@ 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, @@ -1684,13 +1685,9 @@ 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 - if self._is_disagg: - # Disaggregation must validate unified ownership, never fall back - # to draft state that the transceiver cannot transfer. - return True - # Keep existing aggregate execution for settings the unified draft - # lifecycle cannot support. Disaggregation must never take this path. - return self._unified_draft_cache_unsupported_reason() is None + # 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.""" @@ -1715,14 +1712,9 @@ def _validate_standalone_draft_cache(self) -> None: def _unified_draft_cache_unsupported_reason(self) -> Optional[str]: """Shared admission requirements for unified aggregate and disagg KV.""" - if self._kv_cache_config.enable_block_reuse: - return "Unified DSpark draft KV does not yet support prefix reuse" - if self._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 self._kv_cache_config.pool_ratio is not None: - return "Unified DSpark draft KV does not yet support explicit pool_ratio" + reason = get_draft_cache_unsupported_reason(self._kv_cache_config) + if reason is not None: + return reason if (self._speculative_config.draft_len_schedule is not None or self._speculative_config.max_concurrency is not None): return ( @@ -2663,37 +2655,36 @@ def _get_qwen4_exp_ple_cache_params(config, *, total_layers: int, def _create_kv_cache_manager( - model_engine: Optional[PyTorchModelEngine], - kv_cache_manager_cls, - mapping: Mapping, - kv_cache_config: KvCacheConfig, - tokens_per_block: int, - max_seq_len: int, - max_batch_size: int, - spec_config: Optional[SpeculativeConfig], - sparse_attention_config: Optional[SparseAttentionConfig], - max_num_tokens: int, - max_beam_width: int, - kv_connector_manager: Optional[KvCacheConnectorManager], - estimating_kv_cache: bool = False, - enable_kv_cache_stats: bool = False, - execution_stream: Optional[torch.cuda.Stream] = None, - # Optional overrides for one-model draft case (when model_engine is None) - model_config: Optional[ModelConfig] = None, - dtype: Optional[torch.dtype] = None, - is_draft: Optional[bool] = None, - layer_mask: Optional[List[bool]] = None, - num_layers: Optional[int] = None, - num_kv_heads: Optional[Union[int, List[int]]] = None, - head_dim: Optional[int] = None, - kv_cache_type=None, - is_disagg: bool = False, - disable_overlap_scheduler: bool = False, - cold_page_codec_provider: Optional[object] = None, - kv_events_config: Optional[KVEventsConfig] = None, - joint_kv_cache_reuse: bool = False, - standalone_draft_layout: Optional[StandaloneDraftLayout] = None -) -> KVCacheManager: + model_engine: Optional[PyTorchModelEngine], + kv_cache_manager_cls, + mapping: Mapping, + kv_cache_config: KvCacheConfig, + tokens_per_block: int, + max_seq_len: int, + max_batch_size: int, + spec_config: Optional[SpeculativeConfig], + sparse_attention_config: Optional[SparseAttentionConfig], + max_num_tokens: int, + max_beam_width: int, + kv_connector_manager: Optional[KvCacheConnectorManager], + estimating_kv_cache: bool = False, + enable_kv_cache_stats: bool = False, + execution_stream: Optional[torch.cuda.Stream] = None, + # Optional overrides for one-model draft case (when model_engine is None) + model_config: Optional[ModelConfig] = None, + dtype: Optional[torch.dtype] = None, + is_draft: Optional[bool] = None, + layer_mask: Optional[List[bool]] = None, + num_layers: Optional[int] = None, + num_kv_heads: Optional[Union[int, List[int]]] = None, + head_dim: Optional[int] = None, + kv_cache_type=None, + is_disagg: bool = False, + 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) -> KVCacheManager: """ Returns: A KVCacheManager instance for the given model engine or model config 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 6f8f9c7d6e01..9dab0a723748 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 @@ -1134,6 +1134,21 @@ 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_block_reuse: + return "Unified DSpark draft KV does not yet support prefix reuse" + 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, ...] = () @@ -1188,18 +1203,9 @@ def __init__( standalone_draft_layout.extra_tokens if standalone_draft_layout is not None else 0 ) if standalone_draft_layout is not None: - if kv_cache_config.enable_block_reuse: - raise ValueError("Unified standalone draft KV does not yet support prefix reuse") - if kv_cache_config.enable_swa_scratch_reuse: - raise ValueError( - "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: - raise ValueError( - "Unified standalone draft KV does not yet support explicit pool_ratio" - ) + reason = get_draft_cache_unsupported_reason(kv_cache_config) + if reason is not None: + raise ValueError(reason) self.mapping = mapping self.dtype = dtype self.is_disagg = is_disagg @@ -2980,13 +2986,10 @@ 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} - if self.draft_layout is not None - else {} - ), ) ) + if hasattr(layer, "cache_domain"): + layers[-1].cache_domain = layer.cache_domain else: layers.append(SsmLayerConfig(layer_id=layer_id, buffers=buffers)) @@ -3144,7 +3147,7 @@ def get_buffers(self, layer_idx: int, kv_layout: str = "NHD") -> Optional[torch. ) def _is_standalone_draft_layer(self, local_layer_idx: int) -> bool: - return getattr(self, "draft_layout", None) is not None and any( + 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 ) @@ -3245,7 +3248,7 @@ def set_draft_history(self, request_id: int, valid_length: int, position: int) - f"Standalone draft request {request_id} has no active cache allocation" ) history = StandaloneDraftHistory(valid_length, position) - if history.valid_length > cache.capacity or history.position > cache.capacity: + if history.position > cache.capacity: raise ValueError("Standalone draft history exceeds allocated capacity") if ( self.draft_layout.window_size is not None @@ -3277,8 +3280,7 @@ def restore_draft_history(self, request_id: int, metadata: dict) -> None: valid_length = metadata.get("valid_length") position = metadata.get("position") history = StandaloneDraftHistory(valid_length, position) - # Accessing the receiving mapping here verifies ownership before the - # worker can see this history. The sender's slot/page IDs are never used. + # Validate receiver-local allocation before publishing history. self.get_draft_block_table([request_id], [history]) self.set_draft_history(request_id, valid_length, position) @@ -5442,7 +5444,7 @@ def get_batch_cache_indices_flat( return out_tensor def get_cache_bytes_per_token(self) -> int: - if getattr(self, "draft_layout", None) is not None: + 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: diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache/standalone_draft_cache.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/standalone_draft_cache.py index a4ad2586fc68..44aef500e0fa 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache/standalone_draft_cache.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/standalone_draft_cache.py @@ -49,25 +49,20 @@ def bytes_per_token(self) -> int: return self.num_layers * self.bytes_per_layer_token def transfer_identity(self) -> dict: - identity = { + 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, } - if self.kv_factor != 2 or self.window_size is not None: - identity.update(kv_factor=self.kv_factor, window_size=self.window_size) - return identity @dataclass(frozen=True) class StandaloneDraftHistory: - """Committed draft tokens and their next absolute sequence position. - - These are deliberately separate from the target cache's monotonic history - watermark and from its speculative allocation capacity. - """ + """Committed length and next absolute position, independent of target/scratch state.""" valid_length: int position: int diff --git a/tensorrt_llm/_torch/speculative/dflash.py b/tensorrt_llm/_torch/speculative/dflash.py index ae20788ca484..012b296db1a9 100644 --- a/tensorrt_llm/_torch/speculative/dflash.py +++ b/tensorrt_llm/_torch/speculative/dflash.py @@ -346,17 +346,10 @@ def get_draft_kv_cache_manager( return super().get_draft_kv_cache_manager(resource_manager) def _has_unified_draft_cache(self) -> bool: - return getattr(getattr(self, "_ctx_kv_manager", None), "draft_layout", None) is not None + return getattr(self._ctx_kv_manager, "draft_layout", None) is not None def _prepare_managed_history(self, request_ids: list[int]) -> None: - """Restore each request's committed history before this forward's writes. - - Metadata preparation precedes lazy buffer binding. That binding can - replace an estimation manager and clear its staging slots, so generation - must rebind here even when metadata preparation already assigned slots. - Context requests also restore here so continuation and restart detection - use the same authoritative state as generation. - """ + """Restore history before writes, including after profiling replaces the manager.""" updates = {self._dummy_slot: 0} for request_id in request_ids: if ( @@ -381,11 +374,7 @@ def _prepare_managed_history(self, request_ids: list[int]) -> None: ) def _gather_managed_context(self, request_ids: list[int]) -> None: - """Refresh VANILLA's dense input staging from manager-owned history. - - The dense tensors are forward workspaces: they are never exported or - used to recover request history. Noise K/V may overwrite their suffix. - """ + """Refresh VANILLA's dense staging; manager pages remain authoritative.""" if self._dflash_attention_backend != "VANILLA": return for row, request_id in enumerate(request_ids): @@ -699,25 +688,17 @@ def _lazy_init_ctx_buffers( for pool in self._ctx_kv_buf ): raise ValueError("Unified DSpark draft pool does not match the drafter KV layout") + max_blocks = draft_kv_cache_manager.max_blocks_per_seq self._ctx_block_tables = torch.zeros( - (num_slots, draft_kv_cache_manager.max_blocks_per_seq), - dtype=torch.int32, - device="cuda", + (num_slots, max_blocks), dtype=torch.int32, device="cuda" ) - max_blocks = draft_kv_cache_manager.max_blocks_per_seq self._ctx_block_indptr = torch.arange( 0, (num_slots + 1) * max_blocks, max_blocks, dtype=torch.int32, device="cuda" ) self._ctx_kv_last_page_len = torch.full( (num_slots,), self._ctx_page_size, dtype=torch.int32, device="cuda" ) - if self._dflash_attention_backend == "VANILLA": - # Dense FlashAttention inputs are transient forward staging; - # every history read is refreshed from the authoritative pool. - kv_shape = (num_slots, L, capacity, nkv, hd) - self._ctx_k_buf = torch.zeros(kv_shape, dtype=dtype, device="cuda") - self._ctx_v_buf = torch.zeros(kv_shape, dtype=dtype, device="cuda") - else: # TRTLLM; StandaloneDraftLayout validates the backend. + if self._dflash_attention_backend == "TRTLLM": validate_dflash_trtllm_gen_runtime( dtype=dtype, num_heads=nh, @@ -790,8 +771,9 @@ def _lazy_init_ctx_buffers( self._ctx_kv_last_page_len = torch.full( (num_slots,), page_size, dtype=torch.int32, device="cuda" ) - else: # VANILLA DFlash backend (FlashAttention) - self._check_ctx_arena_fits(capacity, num_slots, L, nkv, hd, dtype) + if self._dflash_attention_backend == "VANILLA": + if not unified: + self._check_ctx_arena_fits(capacity, num_slots, L, nkv, hd, dtype) kv_shape = (num_slots, L, capacity, nkv, hd) self._ctx_k_buf = torch.zeros(kv_shape, dtype=dtype, device="cuda") self._ctx_v_buf = torch.zeros(kv_shape, dtype=dtype, device="cuda") @@ -1009,8 +991,7 @@ def _store_prefill_context( # where the previous one ended. Everything else -- a fresh request, a # request id reused after completion, a prefill restarted after # preemption -- has to start from a clean slot. - previous_position = self._req_ctx_pos.get(req_id) - reset = previous_position != first_pos + reset = self._req_ctx_pos.get(req_id) != first_pos if self._assign_slot(req_id, reset=reset, updates=ctx_len_updates) is None: logger.warning("DFlash: no free slots, skipping context store") self._req_ctx_pos.pop(req_id, None) diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py index 18519741b2ec..c47c59c5434c 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_config.py @@ -112,12 +112,6 @@ class AttentionLayerConfig: # Note that we use None to represent "no sliding window". Sink tokens are excluded. sliding_window_size: int | None = None num_sink_tokens: int | None = None - cache_domain: str = "target" - """Ownership domain for layers that can share a lifecycle and physical pools. - - Standalone draft layers use a separate domain because their valid history and - speculative scratch need not advance with the target, even for identical layouts. - """ @property def window_size(self) -> int | None: diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_life_cycle_registry.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_life_cycle_registry.py index 4d4c3cef7812..d0da6b057678 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_life_cycle_registry.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_life_cycle_registry.py @@ -23,21 +23,17 @@ class AttnLifeCycle(NamedTuple): window_size: SlidingWindowSize num_sink_blocks: int # div_up(num_sink_tokens, tokens_per_block) - cache_domain: str = "target" @staticmethod def make( - window_size: SlidingWindowSize, - num_sink_tokens: int | None, - tokens_per_block: int, - cache_domain: str = "target", + window_size: SlidingWindowSize, num_sink_tokens: int | None, tokens_per_block: int ) -> "AttnLifeCycle": assert tokens_per_block > 0 assert window_size is None or window_size > 0 assert num_sink_tokens is None or num_sink_tokens >= 0 assert num_sink_tokens in (None, 0) or window_size is not None num_sink_blocks = div_up(num_sink_tokens or 0, tokens_per_block) - return AttnLifeCycle(window_size, num_sink_blocks, cache_domain) + return AttnLifeCycle(window_size, num_sink_blocks) def get_stale_range( self, history_length: int, tokens_per_block: int @@ -77,9 +73,7 @@ def make_life_cycle(layer: LayerConfig, tokens_per_block: int) -> LifeCycle: return ssm_life_cycle else: assert isinstance(layer, AttentionLayerConfig) - return AttnLifeCycle.make( - layer.window_size, layer.num_sink_tokens, tokens_per_block, layer.cache_domain - ) + return AttnLifeCycle.make(layer.window_size, layer.num_sink_tokens, tokens_per_block) class LifeCycleRegistry: diff --git a/tensorrt_llm/runtime/kv_cache_manager_v2/_storage/_config.py b/tensorrt_llm/runtime/kv_cache_manager_v2/_storage/_config.py index a42c7d7b860b..6a62470e582a 100644 --- a/tensorrt_llm/runtime/kv_cache_manager_v2/_storage/_config.py +++ b/tensorrt_llm/runtime/kv_cache_manager_v2/_storage/_config.py @@ -20,13 +20,7 @@ from .._common import LayerId from .._config import CacheTierConfig, DataRole, KVCacheManagerConfig -from .._life_cycle_registry import ( - AttnLifeCycle, - LayerGroupId, - LifeCycleId, - LifeCycleRegistry, - make_life_cycle, -) +from .._life_cycle_registry import LayerGroupId, LifeCycleId, LifeCycleRegistry, make_life_cycle from .._storage._core import PoolGroupIndex, PoolIndex from .._utils import ( HomoTuple, @@ -231,20 +225,15 @@ def create_storage_config(config: KVCacheManagerConfig) -> StorageConfig: slot_groups.append( SlotDescVariant(life_cycle_id, cast(TypedIndexList[PoolIndex, CoalescedBuffer], slots)) ) - # Equal storage sizes permit merging only within a compatible ownership domain. - # Existing target attention/SSM groups retain their shared physical pool behavior. - pool_groups_by_layout = defaultdict[tuple[str, HomoTuple[int]], list[SlotDescVariant]]( + # Merge slot groups with the same slot_size_list + pool_groups_by_slot_size_list = defaultdict[HomoTuple[int], list[SlotDescVariant]]( list[SlotDescVariant] ) for slot_group in slot_groups: - life_cycle = life_cycle_registry[slot_group.life_cycle_id] - cache_domain = ( - life_cycle.cache_domain if isinstance(life_cycle, AttnLifeCycle) else "target" - ) - pool_groups_by_layout[(cache_domain, tuple(slot_group.slot_size_list))].append(slot_group) + pool_groups_by_slot_size_list[tuple(slot_group.slot_size_list)].append(slot_group) slot_desc_list = cast( TypedIndexList[PoolGroupIndex, SlotDesc], - [SlotDesc(tuple(slot_groups)) for slot_groups in pool_groups_by_layout.values()], + [SlotDesc(tuple(slot_groups)) for slot_groups in pool_groups_by_slot_size_list.values()], ) return StorageConfig( cache_tiers=tuple(config.cache_tiers), 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: {}, ) From 698bb939b5e9de1e34856c008c64154a882cf892 Mon Sep 17 00:00:00 2001 From: allisonlim-nv Date: Tue, 22 Sep 2026 15:16:29 -0700 Subject: [PATCH 07/13] Refactor lifeCycleRegistry to simplify SsmLayerConfig handling Removed unnecessary conditional block for SsmLayerConfig. Signed-off-by: allisonlim-nv --- .../batch_manager/kv_cache_manager_v2/lifeCycleRegistry.cpp | 4 ---- 1 file changed, 4 deletions(-) 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 e2b105898798..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 @@ -35,13 +35,9 @@ LifeCycle makeLifeCycle(LayerConfig const& layer, int tokensPerBlock) cfg.validate(); using T = std::decay_t; if constexpr (std::is_same_v) - { return SsmLifeCycle{}; - } else - { return AttnLifeCycle::make(cfg.slidingWindowSize, cfg.numSinkTokens, tokensPerBlock, cfg.cacheDomain); - } }, layer); } From 52c8e56516be53e7a71240f851f169727f0b29f0 Mon Sep 17 00:00:00 2001 From: Allison Lim Date: Tue, 22 Sep 2026 15:31:29 -0700 Subject: [PATCH 08/13] [None][refactor] Simplify unified draft cache helpers Signed-off-by: Allison Lim --- .../kv_cache_manager_v2/lifeCycleRegistry.h | 12 +++--------- .../pyexecutor/kv_cache/kv_cache_manager_v2.py | 15 +++------------ tensorrt_llm/_torch/speculative/dflash.py | 3 +-- 3 files changed, 7 insertions(+), 23 deletions(-) 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 bcfc8765f8ab..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 @@ -26,6 +26,7 @@ #include #include #include +#include #include #include #include @@ -63,15 +64,8 @@ struct AttnLifeCycle bool operator<(AttnLifeCycle const& o) const noexcept { - if (windowSize != o.windowSize) - { - return windowSize < o.windowSize; - } - if (numSinkBlocks != o.numSinkBlocks) - { - return numSinkBlocks < o.numSinkBlocks; - } - return cacheDomain < o.cacheDomain; + return std::tie(windowSize, numSinkBlocks, cacheDomain) + < std::tie(o.windowSize, o.numSinkBlocks, o.cacheDomain); } static AttnLifeCycle make( 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 594812610d46..8f47c18685fc 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 @@ -2988,10 +2988,9 @@ 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, ) ) - if hasattr(layer, "cache_domain"): - layers[-1].cache_domain = layer.cache_domain else: layers.append(SsmLayerConfig(layer_id=layer_id, buffers=buffers)) @@ -3154,20 +3153,12 @@ def _is_standalone_draft_layer(self, local_layer_idx: int) -> bool: ) def get_layer_cache_dtype(self, layer_idx: int) -> DataType: - if self.draft_layout is None: - return self.dtype - if self._is_standalone_draft_layer(self.layer_offsets[layer_idx]): + 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: - if self.draft_layout is None: - return self.kv_factor - return ( - self.draft_layout.kv_factor - if self._is_standalone_draft_layer(self.layer_offsets[layer_idx]) - else self.kv_factor - ) + 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.""" diff --git a/tensorrt_llm/_torch/speculative/dflash.py b/tensorrt_llm/_torch/speculative/dflash.py index 0630ec7ee063..56c610e64d33 100644 --- a/tensorrt_llm/_torch/speculative/dflash.py +++ b/tensorrt_llm/_torch/speculative/dflash.py @@ -2116,14 +2116,13 @@ 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) if torch.any(ctx_len_gen + gen_num_accepted > allocated).item(): raise ValueError("Unified DSpark accepted history exceeds the drafter context") - ctx_position_gen = ctx_len_gen - if self._has_unified_draft_cache(): ctx_position_gen = torch.tensor( [ self._req_ctx_pos.get(request_id, 0) From d53f771c8b108b405ac50a322911aa3008ed5f7d Mon Sep 17 00:00:00 2001 From: Allison Lim Date: Wed, 23 Sep 2026 16:48:24 -0700 Subject: [PATCH 09/13] [None][fix] Allow overlap scheduling with unified DSpark KV cache Signed-off-by: Allison Lim --- tensorrt_llm/_torch/pyexecutor/_util.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 334b3676a6cd..3d437935a9fd 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -1876,9 +1876,6 @@ def _unified_draft_cache_unsupported_reason(self) -> Optional[str]: if self._llm_args.cuda_graph_config is not None: return ("Unified DSpark KV cache currently requires eager " "execution; set cuda_graph_config=None.") - if not self._disable_overlap_scheduler: - return ("Unified DSpark KV cache requires " - "disable_overlap_scheduler=True.") if self._llm_args.enable_chunked_prefill: return ("Unified DSpark KV cache does not yet support " "chunked prefill; set enable_chunked_prefill=False.") From 61317e94b7b5e62e489ced94aa5d853e6d150b29 Mon Sep 17 00:00:00 2001 From: Allison Lim Date: Fri, 25 Sep 2026 10:14:46 -0700 Subject: [PATCH 10/13] [None][fix] Allow chunked prefill with unified DSpark KV cache Signed-off-by: Allison Lim --- tensorrt_llm/_torch/pyexecutor/_util.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index e64994d9ca6c..96745060a1db 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -1876,9 +1876,6 @@ def _unified_draft_cache_unsupported_reason(self) -> Optional[str]: if self._llm_args.cuda_graph_config is not None: return ("Unified DSpark KV cache currently requires eager " "execution; set cuda_graph_config=None.") - if self._llm_args.enable_chunked_prefill: - return ("Unified DSpark KV cache does not yet support " - "chunked prefill; set enable_chunked_prefill=False.") 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._mapping.enable_attention_dp From 234ee16411a0ced12cbeb242f2fd1eb5390e4d61 Mon Sep 17 00:00:00 2001 From: Allison Lim Date: Fri, 25 Sep 2026 14:11:55 -0700 Subject: [PATCH 11/13] [None][fix] Support CUDA graphs with unified DSpark KV cache Prepare manager-owned draft cache metadata before execution and publish each iteration's history after sampling completion. Keep draft KV writes and history progression in persistent device buffers so graph replay and overlap scheduling use current state. Signed-off-by: Allison Lim --- tensorrt_llm/_torch/pyexecutor/_util.py | 3 - .../kv_cache/standalone_draft_cache.py | 44 ++++ .../_torch/pyexecutor/model_engine.py | 18 ++ tensorrt_llm/_torch/pyexecutor/py_executor.py | 14 +- .../_torch/pyexecutor/sampler/sampler.py | 2 + tensorrt_llm/_torch/speculative/dflash.py | 151 +++++++---- tensorrt_llm/_torch/speculative/dspark.py | 244 +++++++++++++++--- tensorrt_llm/_torch/speculative/interface.py | 15 ++ 8 files changed, 394 insertions(+), 97 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 96745060a1db..15371c62be57 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -1873,9 +1873,6 @@ def _unified_draft_cache_unsupported_reason(self) -> Optional[str]: "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._llm_args.cuda_graph_config is not None: - return ("Unified DSpark KV cache currently requires eager " - "execution; set cuda_graph_config=None.") 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._mapping.enable_attention_dp diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache/standalone_draft_cache.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/standalone_draft_cache.py index 44aef500e0fa..8b2e76df2eff 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache/standalone_draft_cache.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/standalone_draft_cache.py @@ -4,9 +4,13 @@ """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: @@ -72,3 +76,43 @@ def __post_init__(self) -> None: 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 ebc4177eab27..503c73a2fb01 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 @@ -6428,6 +6429,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 @@ -6511,6 +6519,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 62ccc9abeb35..3d9a97a44684 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -2921,7 +2921,9 @@ def _executor_loop_pp(self): # Copy the batch outputs as sampler inputs # to avoid next forward step overwriting them. batch_outputs_copy = { - name: tensor.clone() + name: + tensor.clone() if isinstance( + tensor, torch.Tensor) else tensor for name, tensor in batch_outputs.items() } self.sample_stream.wait_stream( @@ -7706,6 +7708,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 @@ -7732,10 +7737,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 a66449af8bcf..a0f792239c98 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 56c610e64d33..b70b73404401 100644 --- a/tensorrt_llm/_torch/speculative/dflash.py +++ b/tensorrt_llm/_torch/speculative/dflash.py @@ -28,6 +28,7 @@ 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, ResourceManagerType from .accept_stats import maybe_create_recorder @@ -543,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()): @@ -667,12 +668,17 @@ def __init__( # graph compatible. self._ctx_buf_inited = False self._ctx_len = None + self._ctx_position_offset = None + self._managed_cache_bindings = {} + self._managed_restore_rows = set() + 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 @@ -799,22 +805,42 @@ 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]) -> None: - """Restore history before writes, including after profiling replaces the manager.""" + """Restore newly resident histories without rewinding overlapping iterations.""" updates = {self._dummy_slot: 0} - for request_id in request_ids: - if ( - request_id != ATTENTION_DP_DUMMY_REQUEST_ID - and request_id < self._graph_dummy_id_floor - ): + offsets = {self._dummy_slot: 0} + self._managed_restore_rows = set() + 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) - # Request IDs can be recycled by startup probes or restarted - # requests. Only the manager knows whether pages survived. 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 + self._managed_restore_rows.add(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], @@ -822,12 +848,19 @@ def _prepare_managed_history(self, request_ids: list[int]) -> None: 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 + ) - def _gather_managed_context(self, request_ids: list[int]) -> None: + def _gather_managed_context( + self, request_ids: list[int], *, restored_only: bool = False + ) -> None: """Refresh VANILLA's dense staging; manager pages remain authoritative.""" if self._dflash_attention_backend != "VANILLA": return for row, request_id in enumerate(request_ids): + if restored_only and row not in self._managed_restore_rows: + continue slot = self._req_to_slot.get(request_id) if slot is None: continue @@ -839,18 +872,37 @@ def _gather_managed_context(self, request_ids: list[int]) -> None: 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 _publish_managed_history(self, request_ids: list[int]) -> None: - """Commit draft validity after a successful eager worker forward.""" - lengths = self._ctx_len.tolist() - for request_id in request_ids: - slot = self._req_to_slot.get(request_id) - if slot is None: - continue - length = lengths[slot] - position = self._req_ctx_pos.get(request_id, 0) + length - self._ctx_len_host[slot] - self._ctx_len_host[slot] = length - self._req_ctx_pos[request_id] = position - self._ctx_kv_manager.set_draft_history(request_id, length, position) + 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 + 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, restored_only=True) + + 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. @@ -1136,12 +1188,16 @@ 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_restore_rows = set() + self._managed_request_ids = () # checkpoint's trained block width self._resolved_block_size = getattr(draft_model, "block_size", None) or ( @@ -1368,6 +1424,9 @@ def clear(slot: int) -> None: 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] @@ -1487,6 +1546,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, @@ -1624,10 +1685,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) @@ -1683,20 +1747,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 - if self._has_unified_draft_cache(): - self._prepare_managed_history(spec_metadata.request_ids) - # Before any store: prefill and decode both address pages through it. - self._refresh_ctx_block_tables(attn_metadata, batch_size, spec_metadata.request_ids) - if self._has_unified_draft_cache(): - # Refresh before appending accepted features; the first generation - # iteration therefore consumes the transferred prompt history. - self._gather_managed_context(spec_metadata.request_ids) + 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. @@ -1709,6 +1769,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 @@ -1902,8 +1964,6 @@ def _forward_impl( if is_warmup: self._ctx_len.copy_(self._saved_ctx_len) self._restore_ctx_len_host() - elif self._has_unified_draft_cache(): - self._publish_managed_history(spec_metadata.request_ids) self._ctx_len_restore_pending = False return { @@ -2121,18 +2181,11 @@ def prepare_1st_drafter_inputs( allocated = dflash_allocated_ctx_limit( self._ctx_block_counts[gen_rows_out], self._ctx_page_size, block_size ).clamp(max=self._max_ctx) - if torch.any(ctx_len_gen + gen_num_accepted > allocated).item(): - raise ValueError("Unified DSpark accepted history exceeds the drafter context") - ctx_position_gen = torch.tensor( - [ - self._req_ctx_pos.get(request_id, 0) - for request_id in spec_metadata.request_ids[ - num_contexts : num_contexts + num_gens - ] - ], - dtype=torch.long, - device=ctx_len_gen.device, + 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") diff --git a/tensorrt_llm/_torch/speculative/dspark.py b/tensorrt_llm/_torch/speculative/dspark.py index 12957071c70e..77829de0d91b 100644 --- a/tensorrt_llm/_torch/speculative/dspark.py +++ b/tensorrt_llm/_torch/speculative/dspark.py @@ -24,11 +24,14 @@ 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 @@ -38,6 +41,66 @@ from ...llmapi.llm_args import DSparkDecodingConfig +@triton.jit +def _store_managed_window_kernel( + windows, + pool, + slots, + lengths, + positions, + 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 + real = tl.load(real_rows + row) + capacity = tl.load(capacities + row) + valid = real & (token < length) & (token < window_size) + valid = valid & (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. @@ -131,15 +194,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 @@ -278,6 +333,10 @@ def __init__( 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 @@ -439,6 +498,7 @@ def _assign_slot(self, req_id: int, reset: bool) -> int: 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 @@ -458,6 +518,18 @@ 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 @@ -488,6 +560,44 @@ def _bind_managed_history(self, resource_manager) -> None: 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.""" @@ -499,7 +609,7 @@ def _managed_window_indices(self, row: int, position: int, length: int): return pages, positions % page_size, (positions + 1) % self._win def _prepare_managed_history(self, request_ids: list[int], num_contexts: int) -> None: - """Restore authoritative history once, before any local accepted-token writes.""" + """Restore newly resident histories without rewinding overlapping iterations.""" # Startup probes and graph/ADP padding never publish synthetic history. real_rows = [ row @@ -507,11 +617,36 @@ def _prepare_managed_history(self, request_ids: list[int], num_contexts: int) -> if self._is_managed_request(request_id) ] real_ids = [request_ids[row] for row in real_rows] - table = self._draft_kv_manager.get_draft_block_table(real_ids) - self._draft_block_tables = torch.zeros( - (len(request_ids), table.shape[1]), dtype=table.dtype, device=self._kv_windows.device + 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_block_tables[real_rows] = table.to(self._kv_windows.device) + 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 @@ -519,13 +654,17 @@ def _prepare_managed_history(self, request_ids: list[int], num_contexts: int) -> 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" ) - slot = self._assign_slot(request_id, reset=False) - batch_slots[row] = slot 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 @@ -536,25 +675,37 @@ def _prepare_managed_history(self, request_ids: list[int], num_contexts: int) -> ) 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) ) - - def _publish_managed_history(self, request_ids: list[int]) -> None: - """Commit successful prefill/accepted-feature writes; proposal KV stays scratch.""" - lengths = self._valid_len.tolist() - positions = self._ctx_len.tolist() - for row, request_id in enumerate(request_ids): - if not self._is_managed_request(request_id): - continue - slot = self._req_to_slot[request_id] - length, position = lengths[slot], positions[slot] - if position > self._draft_kv_manager.kv_cache_map[request_id].capacity: - raise ValueError("Embedded DSpark history exceeds its allocated draft capacity") - pages, offsets, frames = self._managed_window_indices(row, position, length) - for stage, pool in enumerate(self._draft_kv_buffers): - pool[pages, 0, 0, offsets] = self._kv_windows[slot, stage, frames] - self._draft_kv_manager.set_draft_history(request_id, length, position) + self._prepared_managed_request_ids = tuple(real_ids) + + def _write_managed_history(self, batch_size: int) -> None: + """Store accepted rolling frames in their pages with replay-time device indices.""" + head_dim = self._kv_windows.shape[-1] + 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, + self._valid_len, + self._ctx_len, + self._draft_block_tables, + self._draft_real_rows, + self._draft_capacities, + *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, @@ -582,7 +733,9 @@ 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) + slot = self._assign_slot( + req_id, reset=first_position == 0 and self._draft_kv_manager is None + ) self._ctx_len[slot] = chunk_positions[-1] + 1 self._position_initialized[slot] = True @@ -776,13 +929,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 - self._bind_managed_history(resource_manager) - if self._draft_kv_manager is not None: - self._prepare_managed_history(spec_metadata.request_ids, num_contexts) + 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 @@ -921,14 +1079,14 @@ def _forward_impl( num_accepted_tokens, ) + if self._draft_kv_manager is not None: + self._write_managed_history(batch_size) + 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) - elif self._draft_kv_manager is not None: - self._publish_managed_history(spec_metadata.request_ids) - 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 be8ee510f645..3234918145a1 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 @@ -1590,6 +1592,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, From 3a78654901a305455c6656774bcabdfc40839e16 Mon Sep 17 00:00:00 2001 From: Allison Lim Date: Sat, 26 Sep 2026 19:21:28 -0700 Subject: [PATCH 12/13] [None][fix] Enable unified DSpark block reuse and attention DP Signed-off-by: Allison Lim --- tensorrt_llm/_torch/pyexecutor/_util.py | 5 -- .../kv_cache/kv_cache_manager_v2.py | 11 +++- tensorrt_llm/_torch/speculative/dspark.py | 51 ++++++++++++++----- 3 files changed, 48 insertions(+), 19 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 15371c62be57..78eecc0dd4d0 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -1875,11 +1875,6 @@ def _unified_draft_cache_unsupported_reason(self) -> Optional[str]: "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._mapping.enable_attention_dp - and not self._is_embedded_dspark()): - return ( - "Unified DSpark KV cache does not yet support " - "attention data parallelism; set enable_attention_dp=False.") if self._kv_connector_manager is not None: return ("Unified DSpark KV cache does not yet support " "KV cache connectors.") 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 faa488f67552..576130f7522d 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 @@ -1112,8 +1112,6 @@ def _settle_context_cursor(req: LlmRequest, reuse: int, tokens_per_block: int) - 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_block_reuse: - return "Unified DSpark draft KV does not yet support prefix reuse" if kv_cache_config.enable_swa_scratch_reuse: return ( "Unified DSpark draft KV cannot use SWA scratch reuse; " @@ -3807,6 +3805,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 diff --git a/tensorrt_llm/_torch/speculative/dspark.py b/tensorrt_llm/_torch/speculative/dspark.py index 77829de0d91b..0c71bd1e1b81 100644 --- a/tensorrt_llm/_torch/speculative/dspark.py +++ b/tensorrt_llm/_torch/speculative/dspark.py @@ -48,6 +48,7 @@ def _store_managed_window_kernel( slots, lengths, positions, + write_starts, block_tables, real_rows, capacities, @@ -72,10 +73,11 @@ def _store_managed_window_kernel( 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 >= 0) & (position < capacity) + valid = valid & (position >= write_start) & (position >= 0) & (position < capacity) page = tl.load( block_tables + row * table_stride + position // page_size, mask=valid, @@ -682,19 +684,23 @@ def _prepare_managed_history(self, request_ids: list[int], num_contexts: int) -> ) self._prepared_managed_request_ids = tuple(real_ids) - def _write_managed_history(self, batch_size: int) -> None: - """Store accepted rolling frames in their pages with replay-time device indices.""" + 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, + self._batch_to_slot[rows], self._valid_len, self._ctx_len, - self._draft_block_tables, - self._draft_real_rows, - self._draft_capacities, + 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), @@ -736,19 +742,35 @@ def _seed_context_windows( slot = self._assign_slot( req_id, reset=first_position == 0 and self._draft_kv_manager is None ) - self._ctx_len[slot] = chunk_positions[-1] + 1 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( @@ -971,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 @@ -1080,7 +1107,7 @@ def _forward_impl( ) if self._draft_kv_manager is not None: - self._write_managed_history(batch_size) + self._write_managed_history(batch_size, write_starts) if is_warmup: self._ctx_len.copy_(saved_ctx_len) From 99e470188a4dc0db858cc772c386a14f7910d43a Mon Sep 17 00:00:00 2001 From: Allison Lim Date: Sun, 27 Sep 2026 21:02:19 -0700 Subject: [PATCH 13/13] [None][refactor] Consolidate unified DSpark cache helpers Signed-off-by: Allison Lim --- .../sparse/deepseek_v4/cache_manager.py | 169 +++++------------- .../kv_cache/kv_cache_manager_v2.py | 105 +++++------ tensorrt_llm/_torch/pyexecutor/py_executor.py | 4 +- tensorrt_llm/_torch/speculative/dflash.py | 122 ++++++------- 4 files changed, 145 insertions(+), 255 deletions(-) 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 0fa0dcf261b1..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 @@ -944,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, @@ -987,85 +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, - non_sliding_attn_size_per_token + context_swa_size_per_token, + context_size_per_token, ) - context_swa_size_per_token += draft_context - generation_swa_size_per_token += draft_generation + 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, - ) - 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, - non_sliding_attn_size_per_token + context_swa_size_per_token, - ) - context_swa_size_per_token += draft_context - generation_swa_size_per_token += draft_generation - generation_swa_size_per_request += draft_per_request + 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 @@ -1073,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 @@ -1418,58 +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, - ) - draft_context, draft_generation, draft_per_request = ( - _draft_cache_size_components( - self.draft_layout, - self.tokens_per_block, - self._generation_kv_capacity_headroom, - non_sliding_attn_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=False, - indexer_k_dtype=self._indexer_k_dtype, - use_fp8_ds_mla=self.use_fp8_ds_mla, - )[0], - ) - if self.draft_layout is not None - else (0, 0, 0) - ) - return int( - total_tokens * (non_sliding_attn_size_per_token + swa_size_per_token) - + swa_size_per_request - + total_tokens * (draft_context if context else draft_generation) - + draft_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): 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 40b2e5497b4c..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 @@ -3130,47 +3130,58 @@ def get_buffers(self, layer_idx: int, kv_layout: str = "NHD") -> Optional[torch. layer_offset = self.layer_offsets[layer_idx] if self._is_standalone_draft_layer(layer_offset): return self.get_draft_buffers(self.draft_layer_ids.index(layer_idx), kv_layout) - addr_key = self.impl.get_mem_pool_base_address(layer_offset, Role.KEY, PageIndexMode.SHARED) - if self.kv_cache_type != CacheTypeCpp.SELFKONLY: + 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, + ], ) ) @@ -3192,34 +3203,14 @@ def get_draft_buffers(self, local_layer_idx: int, kv_layout: str = "HND") -> tor 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") - if kv_layout not in ("HND", "NHD"): - raise ValueError(f"Unsupported standalone draft KV layout: {kv_layout}") - layer_id = self.layer_offsets[self.draft_layer_ids[local_layer_idx]] - key_address = self.impl.get_mem_pool_base_address(layer_id, Role.KEY, PageIndexMode.SHARED) - if layout.kv_factor == 2: - value_address = self.impl.get_mem_pool_base_address( - layer_id, Role.VALUE, PageIndexMode.SHARED - ) - stride = self.impl.get_page_stride(layer_id, Role.KEY) - if value_address != key_address + stride: - raise ValueError( - "Standalone draft K/V buffers must have adjacent equal-sized pages" - ) - dimensions = ( - [layout.num_kv_heads, self.tokens_per_block, layout.head_dim] - if kv_layout == "HND" - else [self.tokens_per_block, layout.num_kv_heads, layout.head_dim] - ) - return convert_to_torch_tensor( - TensorWrapper( - key_address, - self.get_layer_cache_dtype(self.draft_layer_ids[local_layer_idx]), - [ - self.impl.get_page_index_upper_bound(layer_id, Role.KEY) // layout.kv_factor, - layout.kv_factor, - *dimensions, - ], - ) + 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( diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 179298d07853..08b9aae49d92 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -2921,9 +2921,7 @@ def _executor_loop_pp(self): # Copy the batch outputs as sampler inputs # to avoid next forward step overwriting them. batch_outputs_copy = { - name: - tensor.clone() if isinstance( - tensor, torch.Tensor) else tensor + name: tensor.clone() for name, tensor in batch_outputs.items() } self.sample_stream.wait_stream( diff --git a/tensorrt_llm/_torch/speculative/dflash.py b/tensorrt_llm/_torch/speculative/dflash.py index b70b73404401..8302bd8c1bbb 100644 --- a/tensorrt_llm/_torch/speculative/dflash.py +++ b/tensorrt_llm/_torch/speculative/dflash.py @@ -670,7 +670,6 @@ def __init__( self._ctx_len = None self._ctx_position_offset = None self._managed_cache_bindings = {} - self._managed_restore_rows = set() self._managed_request_ids = () # Host shadows of _ctx_len and of each request's prompt progress. self._ctx_len_host = None @@ -804,11 +803,11 @@ def get_draft_kv_cache_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]) -> 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} - self._managed_restore_rows = set() + restore_rows = [] real_request_ids = { request_id for request_id in request_ids @@ -837,7 +836,7 @@ def _prepare_managed_history(self, request_ids: list[int]) -> None: 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 - self._managed_restore_rows.add(row) + restore_rows.append(row) self._write_ctx_len(updates) for slot, offset in offsets.items(): self._ctx_position_offset[slot] = offset @@ -851,16 +850,14 @@ def _prepare_managed_history(self, request_ids: list[int]) -> None: 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], *, restored_only: bool = False - ) -> None: + 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, request_id in enumerate(request_ids): - if restored_only and row not in self._managed_restore_rows: - continue + for row in restore_rows: + request_id = request_ids[row] slot = self._req_to_slot.get(request_id) if slot is None: continue @@ -881,11 +878,11 @@ def prepare_managed_draft_cache( return self._lazy_init_ctx_buffers(draft_model, spec_metadata, attn_metadata, manager) spec_metadata._dflash_worker = self - self._prepare_managed_history(spec_metadata.request_ids) + 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, restored_only=True) + 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.""" @@ -975,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 @@ -1003,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 @@ -1042,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" ) @@ -1050,12 +1051,7 @@ 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, request_ids: list[int] | None = None @@ -1196,7 +1192,6 @@ def _lazy_init_ctx_buffers( self._req_to_slot = {} self._req_ctx_pos = {} self._managed_cache_bindings = {} - self._managed_restore_rows = set() self._managed_request_ids = () # checkpoint's trained block width @@ -1240,46 +1235,27 @@ def _lazy_init_ctx_buffers( draft_model, "_paged_ctx_cache", False ) self._ctx_paged = unified or use_paged - if unified: - self._ctx_kv_buf = [ - 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(pool.shape[1:]) != expected or pool.dtype != dtype - for pool in self._ctx_kv_buf - ): - raise ValueError("Unified DSpark draft pool does not match the drafter KV layout") - max_blocks = draft_kv_cache_manager.max_blocks_per_seq - self._ctx_block_tables = torch.zeros( - (num_slots, max_blocks), dtype=torch.int32, device="cuda" - ) - self._ctx_block_indptr = torch.arange( - 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") - self._ctx_kv_last_page_len = torch.full( - (num_slots,), self._ctx_page_size, dtype=torch.int32, device="cuda" - ) - if self._dflash_attention_backend == "TRTLLM": - draft_model.validate_block_attention_windows() - validate_dflash_trtllm_gen_runtime( - dtype=dtype, - num_heads=nh, - num_kv_heads=nkv, - head_dim=hd, - tokens_per_block=self._ctx_page_size, - has_context_attention=any( - not draft_model._get_attention_mask_args(i)[0] for i in range(L) - ), + 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 + ) ) - elif use_paged: - 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) @@ -1302,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: @@ -1339,12 +1319,14 @@ 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" )