Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,8 @@
if TYPE_CHECKING:
from tensorrt_llm.llmapi.llm_args import SparseAttentionConfig

from .cache_manager import DeepseekV4CacheManager

DEEPSEEK_V4_SPARSE_RATIO = 4
DEEPSEEK_V4_OVERLAP_COMPRESSOR_RATIO = 4

Expand Down Expand Up @@ -516,6 +518,9 @@ def __post_init__(self):
# so compute them once during initialization instead of every prepare().
self._init_cache_buffer_data_pointers()

# Draft-sized sparse buffers for one-model MTP separate draft KV cache.
self._init_draft_sparse_buffers()

def prepare_for_indexer_k_cache(self):
"""Prepare the shared indexer K-cache decode table for DSA kernels."""
# INDEXER_COMPRESS uses shared page indices, so the generic DSA
Expand Down Expand Up @@ -647,32 +652,133 @@ def _prepare_deepseek_v4_indices_compiled(
raise ValueError(f"Unsupported compress_ratio: {compress_ratio}")
sparse_mla_topk_lens_bufs[compress_ratio][:num_tokens] = total_count.to(torch.int32)

def _build_cache_buffer_data_pointers(
self, manager: "DeepseekV4CacheManager", compress_ratios_by_layer: list[int]
) -> tuple[dict[int, int], dict[int, int], dict[int, int]]:
"""Build (sparse_mla_base_ptrs, swa_buffer_ptrs, compressed_buffer_ptrs)
for a DeepSeek-V4 cache manager; serves both target and draft managers.
``compress_ratios_by_layer`` is indexed by global layer index."""
sparse_mla_base_ptrs = {
1: manager.swa_pool_ptr,
}
for ratio, compress_pool_ptr in manager.compress_pool_ptrs.items():
sparse_mla_base_ptrs[ratio] = compress_pool_ptr

swa_buffer_ptrs = {layer_idx: manager.swa_pool_ptr for layer_idx in manager.pp_layers}
compressed_buffer_ptrs = {
layer_idx: manager.get_buffers(layer_idx, DeepseekV4AttentionType.COMPRESS).data_ptr()
for layer_idx in manager.pp_layers
if is_compress_layer(compress_ratios_by_layer[layer_idx])
}
return sparse_mla_base_ptrs, swa_buffer_ptrs, compressed_buffer_ptrs

def _init_cache_buffer_data_pointers(self):
# If MTP is enabled, enlarge the compress ratios by max_draft_tokens - 1
extend_compress_ratios = self.compress_ratios + [self.compress_ratios[-1]] * (
self.max_draft_tokens - 1
)
# SWA uses PER_LAYER indices; COMPRESS uses SHARED indices. The sparse
# MLA conversion kernel receives a representative base pointer per pool
# and a per-layer buffer pointer so it can account for any layer offset.
self.sparse_mla_base_ptrs = {
1: self.kv_cache_manager.swa_pool_ptr,
}
for ratio, compress_pool_ptr in self.kv_cache_manager.compress_pool_ptrs.items():
self.sparse_mla_base_ptrs[ratio] = compress_pool_ptr
(
self.sparse_mla_base_ptrs,
self.swa_buffer_ptrs,
self.compressed_buffer_ptrs,
) = self._build_cache_buffer_data_pointers(self.kv_cache_manager, extend_compress_ratios)

def _init_draft_sparse_buffers(self):
"""Allocate draft-sized sparse block tables + draft pool pointers for
one-model MTP with a separate draft KV cache manager (SWA-only draft).
The draft needs its own tensors because it has fewer layers than the
target and cannot share them within one captured CUDA graph. No-op when
there is no separate draft manager."""
self.draft_sliding_block_tables = None
self.draft_compress_block_tables = None
self.draft_sparse_mla_base_ptrs = None
self.draft_swa_buffer_ptrs = None
self.draft_compressed_buffer_ptrs = None
draft_mgr = self.draft_kv_cache_manager
if draft_mgr is None or not hasattr(draft_mgr, "compute_sliding_block_tables"):
return
# Only SWA-only (compress_ratio 1) draft layers are supported; compress
# (128) / indexer (4) draft layers are intentionally unsupported.
draft_ratio = self.compress_ratios[-1]
if draft_ratio != 1:
raise NotImplementedError(
"Separate DeepSeek-V4 draft KV cache supports only SWA-only "
f"(compress_ratio 1) MTP draft layers; got ratio {draft_ratio}."
)
capture_graph = self.is_cuda_graph
draft_block_table_shape = (
draft_mgr.num_local_layers,
len(DEEPSEEK_V4_SLIDING_ATTENTION),
self.max_num_sequences,
draft_mgr.max_blocks_per_seq,
)
self.draft_sliding_block_tables = self.get_empty(
self.cuda_graph_buffers,
draft_block_table_shape,
cache_name="draft_sliding_block_tables",
dtype=torch.int32,
capture_graph=capture_graph,
)
# SWA-only draft (asserted above): no compress tables.
self.draft_compress_block_tables = {}

extend_compress_ratios = self.compress_ratios + [draft_ratio] * (self.max_draft_tokens - 1)
# Draft pointers keyed by pp_layers (global MTP indices == forward layer_idx).
(
self.draft_sparse_mla_base_ptrs,
self.draft_swa_buffer_ptrs,
self.draft_compressed_buffer_ptrs,
) = self._build_cache_buffer_data_pointers(draft_mgr, extend_compress_ratios)

_DRAFT_SPARSE_FIELDS = (
"sliding_block_tables",
"compress_block_tables",
"sparse_mla_base_ptrs",
"swa_buffer_ptrs",
"compressed_buffer_ptrs",
)

self.swa_buffer_ptrs = {
layer_idx: self.kv_cache_manager.swa_pool_ptr
for layer_idx in self.kv_cache_manager.pp_layers
}
self.compressed_buffer_ptrs = {
layer_idx: self.kv_cache_manager.get_buffers(
layer_idx, DeepseekV4AttentionType.COMPRESS
).data_ptr()
for layer_idx in self.kv_cache_manager.pp_layers
if is_compress_layer(extend_compress_ratios[layer_idx])
def save_target_sparse_state(self) -> dict | None:
"""Snapshot the target fields before applying the draft sparse state.
Returns the snapshot dict, or None when there is no separate draft."""
if getattr(self, "draft_sliding_block_tables", None) is None:
return None
return {f: getattr(self, f) for f in self._DRAFT_SPARSE_FIELDS}

def apply_draft_sparse_state(self) -> None:
"""Repoint the forward-read sparse fields at the draft-sized buffers and
copy the draft tables for the current batch. Precondition: kv_cache_manager
is already the draft manager; save_target_sparse_state was called first so
the caller's finally can restore. Draft tables were computed in prepare();
this copies only."""
for f in self._DRAFT_SPARSE_FIELDS:
setattr(self, f, getattr(self, f"draft_{f}"))
self.prepare_for_block_tables()

def restore_target_sparse_state(self, saved: dict | None) -> None:
"""Restore the target sparse fields from save_target_sparse_state.
No-op when saved is None."""
if saved is None:
return
for f in self._DRAFT_SPARSE_FIELDS:
setattr(self, f, saved[f])

def prepare_for_draft_replay(self) -> dict:
dsa_saved = super().prepare_for_draft_replay()
dsv4_saved = self.save_target_sparse_state()
if dsv4_saved is not None:
self.apply_draft_sparse_state()
return {
"dsa": dsa_saved,
"dsv4": dsv4_saved,
}

def restore_after_draft_replay(self, saved_state: dict | None) -> None:
if saved_state is None:
return
super().restore_after_draft_replay(saved_state["dsa"])
self.restore_target_sparse_state(saved_state["dsv4"])

def prepare(self):
assert self.kv_cache_manager is not None
assert self.request_ids is not None
Expand All @@ -682,6 +788,15 @@ def prepare(self):
self.num_contexts,
)

# Set the draft manager's _num_tables before TrtllmAttentionMetadata.prepare()
# calls its copy_batch_block_offsets; the draft-forward swap then copy-only.
draft_mgr = self.draft_kv_cache_manager
if draft_mgr is not None and hasattr(draft_mgr, "compute_sliding_block_tables"):
draft_mgr.compute_sliding_block_tables(
self.request_ids,
self.num_contexts,
)

TrtllmAttentionMetadata.prepare(self)

num_requests = self.num_contexts + self.num_generations
Expand Down
64 changes: 64 additions & 0 deletions tensorrt_llm/_torch/attention_backend/sparse/dsa.py
Original file line number Diff line number Diff line change
Expand Up @@ -709,6 +709,70 @@ def __post_init__(self):
self.create_buffers_for_mla_rope_append(capture_graph=capture_graph)
self.create_buffers_for_indexer(capture_graph=capture_graph)

def prepare_for_draft_replay(self) -> dict | None:
if (self.kv_cache_manager is None
or not hasattr(self.kv_cache_manager, "index_head_dim")):
return None

saved = {
"host_indexer_k_cache_block_offsets":
self.host_indexer_k_cache_block_offsets.clone(),
"indexer_k_cache_block_offsets":
self.indexer_k_cache_block_offsets.clone(),
"host_slot_mapping_fp8":
self.host_slot_mapping_fp8.clone(),
"host_slot_mapping_scale":
self.host_slot_mapping_scale.clone(),
"slot_mapping_fp8":
self.slot_mapping_fp8.clone(),
"slot_mapping_scale":
self.slot_mapping_scale.clone(),
}

# Derive pool indices from the draft manager's encoded block
# offsets (via _get_pool_block_indices) instead of using raw block
# IDs. With host cache offload, block IDs can exceed
# blocks_in_primary_pool after offload swaps (the block keeps its
# original high ID even though its memory now lives in the primary
# GPU pool). Using raw block IDs as pool indices causes OOB access
# in the indexer k-cache buffers. _get_pool_block_indices correctly
# decodes memPoolBlockIndex from the C++ encoded offsets.
# Note: kv_cache_manager was already swapped to draft above
# in prepare_attn_metadata_for_draft_replay() in _torch/speculative/interface.py
pool_indices = self._get_pool_block_indices()
num_blocks = pool_indices.shape[1]
self.host_indexer_k_cache_block_offsets[:self.num_seqs, :
num_blocks].copy_(pool_indices)
self.indexer_k_cache_block_offsets[:self.num_seqs].copy_(
self.host_indexer_k_cache_block_offsets[:self.num_seqs],
non_blocking=True,
)
# Safety clamp: sanitize stale padding entries beyond num_seqs
# that may contain negative or out-of-range values, matching the
# regular DSA prepare() flow.
self.indexer_k_cache_block_offsets.clamp_(min=0)
Indexer.recompute_slot_mappings(self)

return saved

def restore_after_draft_replay(self, saved_state: dict | None) -> None:
if saved_state is None:
return

self.host_indexer_k_cache_block_offsets.copy_(
saved_state["host_indexer_k_cache_block_offsets"],
non_blocking=True,
)
self.indexer_k_cache_block_offsets.copy_(
saved_state["indexer_k_cache_block_offsets"],
non_blocking=True,
)
self.host_slot_mapping_fp8.copy_(saved_state["host_slot_mapping_fp8"])
self.host_slot_mapping_scale.copy_(
saved_state["host_slot_mapping_scale"])
self.slot_mapping_fp8.copy_(saved_state["slot_mapping_fp8"])
self.slot_mapping_scale.copy_(saved_state["slot_mapping_scale"])

def prepare(self):
super().prepare()
self._invalidate_pool_view_cache()
Expand Down
8 changes: 8 additions & 0 deletions tensorrt_llm/_torch/attention_backend/trtllm.py
Original file line number Diff line number Diff line change
Expand Up @@ -517,6 +517,14 @@ def _bind_runtime_views(
self.prompt_lens_cpu_runtime = prompt_lens_cpu
self.host_request_types_runtime = host_request_types

def prepare_for_draft_replay(self) -> dict | None:
"""Prepare backend-specific state after switching to the draft manager."""
return None

def restore_after_draft_replay(self, saved_state: dict | None) -> None:
"""Restore backend-specific state after switching back to the target
manager."""

def prepare(self) -> None:
super().prepare()
extra_attrs = get_model_extra_attrs()
Expand Down
10 changes: 10 additions & 0 deletions tensorrt_llm/_torch/pyexecutor/_util.py
Original file line number Diff line number Diff line change
Expand Up @@ -1319,6 +1319,16 @@ def _should_create_separate_draft_kv_cache(self) -> bool:
"Attention DP is enabled, separate draft KV cache is not supported."
)
return False

sparse_cfg = self._sparse_attention_config
if (sparse_cfg is not None
and getattr(sparse_cfg, "algorithm", None) == "deepseek_v4"
and self._mapping.pp_size > 1):
logger.info(
"DeepSeek-V4 separate draft KV cache is only supported for PP=1; "
"folding draft layers into the unified manager for pp_size=%d.",
self._mapping.pp_size)
return False
return should_use_separate_draft_kv_cache(self._speculative_config)

def _get_effective_draft_config(self) -> ModelConfig:
Expand Down
65 changes: 14 additions & 51 deletions tensorrt_llm/_torch/speculative/interface.py
Original file line number Diff line number Diff line change
Expand Up @@ -154,46 +154,9 @@ def prepare_attn_metadata_for_draft_replay(attn_metadata,
if attn_metadata.enable_flash_mla:
attn_metadata.prepare_flash_mla()

from ..attention_backend.sparse.dsa import (DSAtrtllmAttentionMetadata,
Indexer)
if (isinstance(attn_metadata, DSAtrtllmAttentionMetadata)
and hasattr(draft_kv_cache_manager, 'index_head_dim')):
m = attn_metadata
saved['saved_dsa_state'] = {
'host_indexer_k_cache_block_offsets':
m.host_indexer_k_cache_block_offsets.clone(),
'indexer_k_cache_block_offsets':
m.indexer_k_cache_block_offsets.clone(),
'host_slot_mapping_fp8':
m.host_slot_mapping_fp8.clone(),
'host_slot_mapping_scale':
m.host_slot_mapping_scale.clone(),
'slot_mapping_fp8':
m.slot_mapping_fp8.clone(),
'slot_mapping_scale':
m.slot_mapping_scale.clone(),
}
# Derive pool indices from the draft manager's encoded block
# offsets (via _get_pool_block_indices) instead of using raw block
# IDs. With host cache offload, block IDs can exceed
# blocks_in_primary_pool after offload swaps (the block keeps its
# original high ID even though its memory now lives in the primary
# GPU pool). Using raw block IDs as pool indices causes OOB access
# in the indexer k-cache buffers. _get_pool_block_indices correctly
# decodes memPoolBlockIndex from the C++ encoded offsets.
# Note: kv_cache_manager was already swapped to draft above (line 67).
pool_indices = m._get_pool_block_indices()
num_blocks = pool_indices.shape[1]
m.host_indexer_k_cache_block_offsets[:m.num_seqs, :num_blocks].copy_(
pool_indices)
m.indexer_k_cache_block_offsets[:m.num_seqs].copy_(
m.host_indexer_k_cache_block_offsets[:m.num_seqs],
non_blocking=True)
# Safety clamp: sanitize stale padding entries beyond num_seqs
# that may contain negative or out-of-range values, matching the
# regular DSA prepare() flow.
m.indexer_k_cache_block_offsets.clamp_(min=0)
Indexer.recompute_slot_mappings(m)
backend_saved = attn_metadata.prepare_for_draft_replay()
if backend_saved is not None:
saved['saved_backend_state'] = backend_saved
return saved


Expand All @@ -208,17 +171,8 @@ def restore_attn_metadata_after_draft_replay(attn_metadata, saved_state):
saved_state['target_host_kv_cache_block_offsets'])
if attn_metadata.enable_flash_mla:
attn_metadata.prepare_flash_mla()
saved_dsa = saved_state.get('saved_dsa_state')
if saved_dsa is not None:
m = attn_metadata
m.host_indexer_k_cache_block_offsets.copy_(
saved_dsa['host_indexer_k_cache_block_offsets'], non_blocking=True)
m.indexer_k_cache_block_offsets.copy_(
saved_dsa['indexer_k_cache_block_offsets'], non_blocking=True)
m.host_slot_mapping_fp8.copy_(saved_dsa['host_slot_mapping_fp8'])
m.host_slot_mapping_scale.copy_(saved_dsa['host_slot_mapping_scale'])
m.slot_mapping_fp8.copy_(saved_dsa['slot_mapping_fp8'])
m.slot_mapping_scale.copy_(saved_dsa['slot_mapping_scale'])
attn_metadata.restore_after_draft_replay(
saved_state.get('saved_backend_state'))


def get_force_num_accepted_tokens() -> int:
Expand Down Expand Up @@ -2171,9 +2125,18 @@ def draft_kv_cache_context(self, attn_metadata, draft_kv_cache_manager):
if attn_metadata.enable_flash_mla:
attn_metadata.prepare_flash_mla()

# DeepSeek-V4: repoint SWA/compress tables and pool pointers to the draft.
if hasattr(attn_metadata, "save_target_sparse_state"):
saved_dsv4_state = attn_metadata.save_target_sparse_state()
if saved_dsv4_state is not None:
Comment thread
lfr-0531 marked this conversation as resolved.
attn_metadata.apply_draft_sparse_state()

try:
yield
finally:
# DeepSeek-V4: restore the target SWA/compress tables and pool pointers.
if hasattr(attn_metadata, "save_target_sparse_state"):
attn_metadata.restore_target_sparse_state(saved_dsv4_state)
# Restore main KV cache manager and block offsets
attn_metadata.kv_cache_manager = target_kv_cache_manager
attn_metadata.kv_cache_block_offsets = target_kv_cache_block_offsets
Expand Down
4 changes: 4 additions & 0 deletions tests/integration/defs/accuracy/references/gsm8k.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -145,6 +145,10 @@ deepseek-ai/DeepSeek-V4-Flash:
# 95.11 reference still holds for the hypothesis test.
- quant_algo: FP8_BLOCK_SCALES
accuracy: 95.11
- quant_algo: FP8_BLOCK_SCALES
kv_cache_quant_algo: FP8
spec_dec_algo: MTP
accuracy: 95.11
deepseek-ai/DeepSeek-V4-Pro:
# Full GSM8K aggregate gate for the Pro deployment path: TP=8, EP=8,
# attention DP, TRTLLM MoE, FP8 KV cache, MTP max_draft_len=1, padded CUDA
Expand Down
Loading
Loading