From e56cedb6686f0d73ac4009e991b50acca175484d Mon Sep 17 00:00:00 2001 From: Xuanyu Chen Date: Wed, 15 Jul 2026 16:03:30 -0700 Subject: [PATCH] [https://nvbugs/6463967][fix] DeepSeek-V4 one-model MTP separate draft kv cache (TEP) Signed-off-by: Xuanyu Chen --- .../sparse/deepseek_v4/deepseek_v4.py | 151 +++++++++++++++--- .../_torch/attention_backend/sparse/dsa.py | 64 ++++++++ .../_torch/attention_backend/trtllm.py | 8 + tensorrt_llm/_torch/pyexecutor/_util.py | 10 ++ tensorrt_llm/_torch/speculative/interface.py | 65 ++------ .../defs/accuracy/references/gsm8k.yaml | 4 + .../defs/accuracy/test_llm_api_pytorch.py | 21 +++ .../test_lists/test-db/l0_dgx_b200.yml | 1 + .../attention/sparse/dsa/test_dsa_indexer.py | 5 +- 9 files changed, 259 insertions(+), 70 deletions(-) diff --git a/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/deepseek_v4.py b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/deepseek_v4.py index 1f76e5a73e51..207efbb3967e 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/deepseek_v4.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/deepseek_v4.py @@ -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 @@ -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 @@ -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 @@ -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 diff --git a/tensorrt_llm/_torch/attention_backend/sparse/dsa.py b/tensorrt_llm/_torch/attention_backend/sparse/dsa.py index 26d8a96e11ce..f593ea7428aa 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/dsa.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/dsa.py @@ -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() diff --git a/tensorrt_llm/_torch/attention_backend/trtllm.py b/tensorrt_llm/_torch/attention_backend/trtllm.py index cb2d5e308d61..3ad89605ad1a 100644 --- a/tensorrt_llm/_torch/attention_backend/trtllm.py +++ b/tensorrt_llm/_torch/attention_backend/trtllm.py @@ -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() diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index 7b4ffa4d4b2f..0e3e3dae540a 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -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: diff --git a/tensorrt_llm/_torch/speculative/interface.py b/tensorrt_llm/_torch/speculative/interface.py index 721b3942ed04..a2b7ee82c72e 100644 --- a/tensorrt_llm/_torch/speculative/interface.py +++ b/tensorrt_llm/_torch/speculative/interface.py @@ -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 @@ -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: @@ -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: + 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 diff --git a/tests/integration/defs/accuracy/references/gsm8k.yaml b/tests/integration/defs/accuracy/references/gsm8k.yaml index ef32f74f1ba0..acaaeac05cef 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-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 diff --git a/tests/integration/defs/accuracy/test_llm_api_pytorch.py b/tests/integration/defs/accuracy/test_llm_api_pytorch.py index df42b0615e65..0c9409439441 100644 --- a/tests/integration/defs/accuracy/test_llm_api_pytorch.py +++ b/tests/integration/defs/accuracy/test_llm_api_pytorch.py @@ -3876,6 +3876,27 @@ 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 (folded under attention_dp=True everywhere else). + # CUDA graph enabled to cover the draft-forward capture/replay swap. + 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", [ pytest.param( 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 77ca7e7f50bf..07ab2df8e88d 100644 --- a/tests/integration/test_lists/test-db/l0_dgx_b200.yml +++ b/tests/integration/test_lists/test-db/l0_dgx_b200.yml @@ -55,6 +55,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 66578c3d7707..c9a6658d24ef 100644 --- a/tests/unittest/_torch/attention/sparse/dsa/test_dsa_indexer.py +++ b/tests/unittest/_torch/attention/sparse/dsa/test_dsa_indexer.py @@ -3290,6 +3290,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_replay.return_value = None return meta @staticmethod @@ -3324,13 +3325,15 @@ def test_prepare_swaps_and_restore_recovers(self): 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_replay.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_replay.assert_called_once_with(None) @pytest.mark.skipif(not has_deep_gemm(), reason="DeepGEMM not available")