diff --git a/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/cache_manager.py b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/cache_manager.py index 925c289b8251..e6ea47db9acd 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/cache_manager.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/cache_manager.py @@ -1283,9 +1283,9 @@ def copy_batch_block_offsets( assert beam_width == 1, "DSV4 only supports beam width 1 now" assert dst_tensor.is_cuda, "copy_batch_block_offsets expects a CUDA destination" dst_tensor.fill_(BAD_PAGE_INDEX) - dst_tensor[:, : self._num_tables, 0, :].copy_( + dst_tensor[:, :num_seqs, 0, :].copy_( self._precomputed_sliding_block_tables[ - :, DeepseekV4AttentionType.SWA.value, : self._num_tables, : + :, DeepseekV4AttentionType.SWA.value, :num_seqs, : ], non_blocking=True, ) @@ -1303,8 +1303,8 @@ def copy_batch_sliding_block_tables( """ assert dst_tensor.is_cuda, "copy_batch_sliding_block_tables expects a CUDA destination" dst_tensor.fill_(BAD_PAGE_INDEX) - dst_tensor[:, :, : self._num_tables, :].copy_( - self._precomputed_sliding_block_tables[:, :, : self._num_tables, :], + dst_tensor[:, :, :num_seqs, :].copy_( + self._precomputed_sliding_block_tables[:, :, :num_seqs, :], non_blocking=True, ) diff --git a/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/metadata.py b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/metadata.py index 5dc1aa9c15d2..cfd22a4da152 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/metadata.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/metadata.py @@ -5,7 +5,7 @@ from __future__ import annotations import math -from typing import Dict, Optional, Set, Tuple +from typing import TYPE_CHECKING, Dict, Optional, Set, Tuple import torch @@ -23,6 +23,9 @@ is_compress_layer, ) +if TYPE_CHECKING: + from .cache_manager import DeepseekV4CacheManager + class DeepseekV4TrtllmAttentionMetadata(DSAtrtllmAttentionMetadata): # The set of compress ratios for the layers @@ -282,6 +285,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 @@ -413,31 +419,95 @@ 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 cache pointers for a target or draft manager.""" + 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 + # 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): + """Initialize sparse buffers for a separate one-model MTP draft cache.""" + self.draft_sliding_block_tables = None + self.draft_sparse_mla_base_ptrs = None + self.draft_swa_buffer_ptrs = None + + draft_mgr = self.draft_kv_cache_manager + from .cache_manager import DeepseekV4CacheManager + + if not isinstance(draft_mgr, DeepseekV4CacheManager): + return + + # Current DSv4 MTP layers are SWA-only. Fail fast if a future model + # introduces a compressed or indexer MTP layer. + 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}." + ) - 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]) - } + 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=self.is_cuda_graph, + ) + extend_compress_ratios = self.compress_ratios + [draft_ratio] * (self.max_draft_tokens - 1) + ( + self.draft_sparse_mla_base_ptrs, + self.draft_swa_buffer_ptrs, + _, + ) = self._build_cache_buffer_data_pointers(draft_mgr, extend_compress_ratios) + + _DRAFT_SPARSE_FIELDS = ( + "sliding_block_tables", + "sparse_mla_base_ptrs", + "swa_buffer_ptrs", + ) + + def prepare_for_draft_forward(self) -> dict | None: + """Repoint sparse fields to the draft buffers for a draft forward.""" + if self.draft_sliding_block_tables is None: + return None + saved_state = {field: getattr(self, field) for field in self._DRAFT_SPARSE_FIELDS} + for field in self._DRAFT_SPARSE_FIELDS: + setattr(self, field, getattr(self, f"draft_{field}")) + return saved_state + + def restore_after_draft_forward(self, saved_state: dict | None) -> None: + """Restore the target sparse fields after a draft forward.""" + if saved_state is None: + return + for field in self._DRAFT_SPARSE_FIELDS: + setattr(self, field, saved_state[field]) def prepare(self): assert self.kv_cache_manager is not None @@ -448,6 +518,22 @@ def prepare(self): self.num_contexts, ) + # Prepare the draft manager's tables before the generic prepare copies + # its block offsets, and populate the dedicated DSv4 draft sparse table. + 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, + ) + if self.draft_sliding_block_tables is not None: + draft_mgr.copy_batch_sliding_block_tables( + self.draft_sliding_block_tables, + self.request_ids, + self.num_contexts, + self.num_seqs, + ) + TrtllmAttentionMetadata.prepare(self) num_requests = self.num_contexts + self.num_generations diff --git a/tensorrt_llm/_torch/attention_backend/sparse/dsa/metadata.py b/tensorrt_llm/_torch/attention_backend/sparse/dsa/metadata.py index a41631cda0e1..998366944daf 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/dsa/metadata.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/dsa/metadata.py @@ -218,6 +218,90 @@ def prepare(self): # Prepare metadata for indexer Indexer.prepare(metadata=self) + def prepare_for_draft_forward(self) -> dict | None: + """Select native DSA indexer metadata for a draft forward.""" + # DeepSeek-V4 metadata inherits DSA metadata, but its cache manager uses a + # different dual-pool layout. Only native DSA cache managers use the DSA + # draft-replay buffers below. + if not is_dsa_cache_manager(self.kv_cache_manager): + return None + + saved_state = { + "host_indexer_k_cache_block_offsets": self.host_indexer_k_cache_block_offsets, + "indexer_k_cache_block_offsets": self.indexer_k_cache_block_offsets, + "host_slot_mapping_fp8": self.host_slot_mapping_fp8, + "host_slot_mapping_scale": self.host_slot_mapping_scale, + "slot_mapping_fp8": self.slot_mapping_fp8, + "slot_mapping_scale": self.slot_mapping_scale, + "block_table": self.block_table, + "block_table_expanded": self.block_table_expanded, + "host_block_table_expanded": self.host_block_table_expanded, + } + # The cached-KV feature owns these references even when an optimized + # path aliases them to slot_mapping_*. With the feature disabled, the + # aliases are lazy and may not exist on the first generation replay. + if self.enable_context_mla_with_cached_kv: + saved_state.update( + { + "slot_mapping_fp8_fullkv": self.slot_mapping_fp8_fullkv, + "slot_mapping_scale_fullkv": self.slot_mapping_scale_fullkv, + } + ) + + # Rebind to the draft manager's dedicated buffers instead of + # overwriting the target tensors in place. Rebinding is invisible to + # CUDA graph capture, so the target and draft segments of the graph + # bake distinct addresses (like draft_kv_cache_block_offsets) and no + # graph-recorded copy from a transient host buffer is needed. + self.host_indexer_k_cache_block_offsets = self.host_draft_indexer_k_cache_block_offsets + self.indexer_k_cache_block_offsets = self.draft_indexer_k_cache_block_offsets + self.host_slot_mapping_fp8 = self.host_draft_slot_mapping_fp8 + self.slot_mapping_fp8 = self.draft_slot_mapping_fp8 + self.host_slot_mapping_scale = self.host_draft_slot_mapping_scale + self.slot_mapping_scale = self.draft_slot_mapping_scale + self.block_table = self.draft_block_table + self.block_table_expanded = self.draft_block_table_expanded + self.host_block_table_expanded = self.host_draft_block_table_expanded + self._invalidate_pool_view_cache() + + # Recording a capture executes no kernels, so the draft mappings only + # need refreshing when the transfers actually run: eager forwards + # (warmup) and the pre-replay call from model_engine. The per-step + # advance inside the captured graph re-derives slot mappings on + # device from the rebound block-offset buffer. + # kv_cache_manager was already swapped to the draft manager above. + if not torch.cuda.is_current_stream_capturing(): + self.prepare_for_indexer_k_cache() + self._refresh_expanded_block_table() + Indexer.recompute_slot_mappings(self) + Indexer.recompute_context_kv_gather_mappings(self) + + return saved_state + + def restore_after_draft_forward(self, saved_state: dict | None) -> None: + """Restore native DSA indexer metadata after a draft forward.""" + if saved_state is None: + return + + self.host_indexer_k_cache_block_offsets = saved_state["host_indexer_k_cache_block_offsets"] + self.indexer_k_cache_block_offsets = saved_state["indexer_k_cache_block_offsets"] + self.host_slot_mapping_fp8 = saved_state["host_slot_mapping_fp8"] + self.host_slot_mapping_scale = saved_state["host_slot_mapping_scale"] + self.slot_mapping_fp8 = saved_state["slot_mapping_fp8"] + self.slot_mapping_scale = saved_state["slot_mapping_scale"] + self.block_table = saved_state["block_table"] + self.block_table_expanded = saved_state["block_table_expanded"] + self.host_block_table_expanded = saved_state["host_block_table_expanded"] + self._invalidate_pool_view_cache() + if "slot_mapping_fp8_fullkv" in saved_state: + self.slot_mapping_fp8_fullkv = saved_state["slot_mapping_fp8_fullkv"] + self.slot_mapping_scale_fullkv = saved_state["slot_mapping_scale_fullkv"] + else: + # The draft recomputation rebound the aliases to the draft tensors; + # point them back at the restored target tensors. + self.slot_mapping_fp8_fullkv = self.slot_mapping_fp8 + self.slot_mapping_scale_fullkv = self.slot_mapping_scale + def get_indexer_kv_lens(self, kv_lens: torch.Tensor) -> torch.Tensor: if self._indexer_compress_ratio <= 1: return kv_lens diff --git a/tensorrt_llm/_torch/attention_backend/trtllm.py b/tensorrt_llm/_torch/attention_backend/trtllm.py index 72e99c6e1066..70f8394ecf7b 100644 --- a/tensorrt_llm/_torch/attention_backend/trtllm.py +++ b/tensorrt_llm/_torch/attention_backend/trtllm.py @@ -652,6 +652,14 @@ def mla_prepare_ctx_cu_seqlens(self) -> Optional[torch.Tensor]: self._mla_ctx_cu_seqlens_valid = True return self.mla_ctx_cu_q_seqlens[:num_ctx + 1] + def prepare_for_draft_forward(self) -> dict | None: + """Prepare backend state shared by draft-forward execution paths.""" + return None + + def restore_after_draft_forward(self, saved_state: dict | None) -> None: + """Restore backend state modified for draft-forward execution.""" + return None + def prepare(self) -> None: super().prepare() # Recomputed on first use this iteration; see mla_prepare_scheduler_buffers. diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 620253e6800f..bddef85d0304 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -1391,6 +1391,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: diff --git a/tensorrt_llm/_torch/speculative/interface.py b/tensorrt_llm/_torch/speculative/interface.py index 3208c6972f1d..f160d1ad3099 100644 --- a/tensorrt_llm/_torch/speculative/interface.py +++ b/tensorrt_llm/_torch/speculative/interface.py @@ -170,64 +170,11 @@ 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, is_dsa_cache_manager) - - # DeepSeek-V4 metadata inherits DSA metadata, but its cache manager uses a - # different dual-pool layout. Only native DSA cache managers use the DSA - # draft-replay buffers below. - if (isinstance(attn_metadata, DSAtrtllmAttentionMetadata) - and is_dsa_cache_manager(draft_kv_cache_manager)): - m = attn_metadata - saved['saved_dsa_state'] = { - 'host_indexer_k_cache_block_offsets': - m.host_indexer_k_cache_block_offsets, - 'indexer_k_cache_block_offsets': m.indexer_k_cache_block_offsets, - 'host_slot_mapping_fp8': m.host_slot_mapping_fp8, - 'host_slot_mapping_scale': m.host_slot_mapping_scale, - 'slot_mapping_fp8': m.slot_mapping_fp8, - 'slot_mapping_scale': m.slot_mapping_scale, - 'block_table': m.block_table, - 'block_table_expanded': m.block_table_expanded, - 'host_block_table_expanded': m.host_block_table_expanded, - } - # The cached-KV feature owns these references even when an optimized - # path aliases them to slot_mapping_*. With the feature disabled, the - # aliases are lazy and may not exist on the first generation replay. - if m.enable_context_mla_with_cached_kv: - saved['saved_dsa_state'].update({ - 'slot_mapping_fp8_fullkv': - m.slot_mapping_fp8_fullkv, - 'slot_mapping_scale_fullkv': - m.slot_mapping_scale_fullkv, - }) - # Rebind to the draft manager's dedicated buffers instead of - # overwriting the target tensors in place. Rebinding is invisible to - # CUDA graph capture, so the target and draft segments of the graph - # bake distinct addresses (like draft_kv_cache_block_offsets) and no - # graph-recorded copy from a transient host buffer is needed. - m.host_indexer_k_cache_block_offsets = ( - m.host_draft_indexer_k_cache_block_offsets) - m.indexer_k_cache_block_offsets = m.draft_indexer_k_cache_block_offsets - m.host_slot_mapping_fp8 = m.host_draft_slot_mapping_fp8 - m.slot_mapping_fp8 = m.draft_slot_mapping_fp8 - m.host_slot_mapping_scale = m.host_draft_slot_mapping_scale - m.slot_mapping_scale = m.draft_slot_mapping_scale - m.block_table = m.draft_block_table - m.block_table_expanded = m.draft_block_table_expanded - m.host_block_table_expanded = m.host_draft_block_table_expanded - m._invalidate_pool_view_cache() - # Recording a capture executes no kernels, so the draft mappings only - # need refreshing when the transfers actually run: eager forwards - # (warmup) and the pre-replay call from model_engine. The per-step - # advance inside the captured graph re-derives slot mappings on - # device from the rebound block-offset buffer. - # kv_cache_manager was already swapped to the draft manager above. - if not torch.cuda.is_current_stream_capturing(): - m.prepare_for_indexer_k_cache() - m._refresh_expanded_block_table() - Indexer.recompute_slot_mappings(m) - Indexer.recompute_context_kv_gather_mappings(m) + # Backends select any additional draft-forward state, such as native DSA + # indexer buffers or DeepSeek-V4 sparse tables and pool pointers. + backend_saved = attn_metadata.prepare_for_draft_forward() + if backend_saved is not None: + saved['saved_backend_state'] = backend_saved return saved @@ -249,29 +196,8 @@ def restore_attn_metadata_after_draft_replay(attn_metadata, saved_state): # needs to invalidate the scheduler metadata; refreshing the unchanged # target buffers would repeat request-specific H2D work. attn_metadata._flash_mla_metadata_valid = False - saved_dsa = saved_state.get('saved_dsa_state') - if saved_dsa is not None: - m = attn_metadata - m.host_indexer_k_cache_block_offsets = saved_dsa[ - 'host_indexer_k_cache_block_offsets'] - m.indexer_k_cache_block_offsets = saved_dsa[ - 'indexer_k_cache_block_offsets'] - m.host_slot_mapping_fp8 = saved_dsa['host_slot_mapping_fp8'] - m.host_slot_mapping_scale = saved_dsa['host_slot_mapping_scale'] - m.slot_mapping_fp8 = saved_dsa['slot_mapping_fp8'] - m.slot_mapping_scale = saved_dsa['slot_mapping_scale'] - m.block_table = saved_dsa['block_table'] - m.block_table_expanded = saved_dsa['block_table_expanded'] - m.host_block_table_expanded = saved_dsa['host_block_table_expanded'] - m._invalidate_pool_view_cache() - if 'slot_mapping_fp8_fullkv' in saved_dsa: - m.slot_mapping_fp8_fullkv = saved_dsa['slot_mapping_fp8_fullkv'] - m.slot_mapping_scale_fullkv = saved_dsa['slot_mapping_scale_fullkv'] - else: - # The draft recomputation rebound the aliases to the draft tensors; - # point them back at the restored target tensors. - m.slot_mapping_fp8_fullkv = m.slot_mapping_fp8 - m.slot_mapping_scale_fullkv = m.slot_mapping_scale + attn_metadata.restore_after_draft_forward( + saved_state.get('saved_backend_state')) def get_force_num_accepted_tokens() -> int: diff --git a/tests/integration/defs/accuracy/references/gsm8k.yaml b/tests/integration/defs/accuracy/references/gsm8k.yaml index 9f62b2bc501e..4244d15b56fe 100644 --- a/tests/integration/defs/accuracy/references/gsm8k.yaml +++ b/tests/integration/defs/accuracy/references/gsm8k.yaml @@ -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-Flash-Base: # Base (pretrained, non-instruct) checkpoint, so GSM8K lands well below the # instruct DeepSeek-V4-Flash above. GSM8K measurements from diff --git a/tests/integration/defs/accuracy/test_llm_api_pytorch.py b/tests/integration/defs/accuracy/test_llm_api_pytorch.py index 75cb7b411afc..46db9c0d4c02 100644 --- a/tests/integration/defs/accuracy/test_llm_api_pytorch.py +++ b/tests/integration/defs/accuracy/test_llm_api_pytorch.py @@ -4053,6 +4053,28 @@ def test_auto_dtype(self): task = GSM8K(self.MODEL_NAME) task.evaluate(llm) + @pytest.mark.skip_less_mpi_world_size(4) + def test_tep_mtp_separate_draft_kv_cache(self): + # TEP (attention_dp=False) + one-model MTP exercises the separate draft + # KV cache manager path. CUDA graphs cover draft capture and replay. + kv_cache_config = KvCacheConfig(free_gpu_memory_fraction=0.5, + dtype="fp8") + with LLM( + self.MODEL_PATH, + tensor_parallel_size=4, + moe_expert_parallel_size=4, + moe_config=MoeConfig(backend="TRTLLM"), + enable_attention_dp=False, + speculative_config=MTPDecodingConfig(max_draft_len=1), + max_batch_size=DEEPSEEKV4_TEST_MAX_BATCH_SIZE, + max_seq_len=4096, + max_num_tokens=4096, + cuda_graph_config=CudaGraphConfig(enable_padding=True), + kv_cache_config=kv_cache_config, + ) as llm: + task = GSM8K(self.MODEL_NAME) + task.evaluate(llm) + @pytest.mark.skip_less_mpi_world_size(4) @parametrize_with_ids("moe_backend", ["TRTLLM", "MEGAMOE_DEEPGEMM"]) def test_nvfp4_4gpus_static_eplb(self, moe_backend): diff --git a/tests/integration/test_lists/test-db/l0_dgx_b200.yml b/tests/integration/test_lists/test-db/l0_dgx_b200.yml index a46a74690d77..12e01f12341e 100644 --- a/tests/integration/test_lists/test-db/l0_dgx_b200.yml +++ b/tests/integration/test_lists/test-db/l0_dgx_b200.yml @@ -53,6 +53,7 @@ l0_dgx_b200: - accuracy/test_llm_api_pytorch.py::TestDeepSeekV3Lite::test_bfloat16_4gpus[pp4-mtp_nextn=0-attention_dp=False-cuda_graph=False-overlap_scheduler=False-torch_compile=False] TIMEOUT (60) - unittest/_torch/modeling/test_modeling_deepseekv4.py - accuracy/test_llm_api_pytorch.py::TestDeepSeekV4Flash::test_auto_dtype TIMEOUT (60) + - accuracy/test_llm_api_pytorch.py::TestDeepSeekV4Flash::test_tep_mtp_separate_draft_kv_cache TIMEOUT (60) # ------------- NVBug 6025177: trtllm-serve cross-request KV contamination (OpenAI) --------------- - test_e2e.py::test_openai_kv_cache_contamination TIMEOUT (120) # ------------- DSA FP4 indexer (Blackwell-only) --------------- diff --git a/tests/unittest/_torch/attention/sparse/dsa/test_dsa_indexer.py b/tests/unittest/_torch/attention/sparse/dsa/test_dsa_indexer.py index 9755b4534af9..76552749651f 100644 --- a/tests/unittest/_torch/attention/sparse/dsa/test_dsa_indexer.py +++ b/tests/unittest/_torch/attention/sparse/dsa/test_dsa_indexer.py @@ -3513,6 +3513,7 @@ def _make_mock_metadata(): meta.kv_cache_block_offsets = torch.tensor([10, 20, 30]) meta.host_kv_cache_block_offsets = torch.tensor([10, 20, 30]) meta.draft_kv_cache_block_offsets = torch.tensor([100, 200, 300]) + meta.prepare_for_draft_forward.return_value = None return meta @staticmethod @@ -3547,13 +3548,15 @@ def is_attn_metadata(obj, cls): assert saved is not None assert saved["target_kv_cache_manager"] is original_kv_mgr assert meta.kv_cache_manager is mgr - assert "saved_dsa_state" not in saved + assert "saved_backend_state" not in saved + meta.prepare_for_draft_forward.assert_called_once_with() restore_attn_metadata_after_draft_replay(meta, saved) assert meta.kv_cache_manager is original_kv_mgr torch.testing.assert_close(meta.kv_cache_block_offsets, original_offsets) torch.testing.assert_close(meta.host_kv_cache_block_offsets, original_host_offsets) + meta.restore_after_draft_forward.assert_called_once_with(None) def test_native_dsa_replay_swaps_and_restores_buffers(self): """Switch native DSA metadata to draft buffers and restore it.""" @@ -3583,6 +3586,15 @@ def test_native_dsa_replay_swaps_and_restores_buffers(self): del meta.slot_mapping_fp8_fullkv del meta.slot_mapping_scale_fullkv + meta.prepare_for_draft_forward.side_effect = ( + lambda: DSAtrtllmAttentionMetadata.prepare_for_draft_forward(meta) + ) + meta.restore_after_draft_forward.side_effect = ( + lambda saved_state: DSAtrtllmAttentionMetadata.restore_after_draft_forward( + meta, saved_state + ) + ) + def is_attn_metadata(obj, cls): if cls in (TrtllmAttentionMetadata, DSAtrtllmAttentionMetadata): return obj is meta @@ -3594,13 +3606,13 @@ def is_attn_metadata(obj, cls): side_effect=is_attn_metadata, ), patch( - "tensorrt_llm._torch.speculative.interface.torch.cuda.is_current_stream_capturing", + "tensorrt_llm._torch.attention_backend.sparse.dsa.metadata.torch.cuda.is_current_stream_capturing", return_value=True, ), ): saved = prepare_attn_metadata_for_draft_replay(meta, mgr) - assert "saved_dsa_state" in saved + assert "saved_backend_state" in saved assert meta.slot_mapping_fp8 is draft_buffers["slot_mapping_fp8"] restore_attn_metadata_after_draft_replay(meta, saved) assert meta.slot_mapping_fp8 is target_buffers["slot_mapping_fp8"]