diff --git a/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/module.py b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/module.py index 68bb1ee8d1fb..2c09926c9835 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/module.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/module.py @@ -1219,6 +1219,11 @@ def _indexer_branch(): self._fused_kv_norm_active = False self._fused_kv_norm_hoisted = False + # Join the prev_topk copy forked in sparse_attn_indexer; CUDA graph + # capture requires the join within the same layer's forward. + if self.indexer is not None: + self.indexer.maybe_join_prev_topk_copy() + class DeepSeekV4Hooks(MLASparseHooks): """Typed DeepSeek-V4 adapter for the shared MLA module.""" diff --git a/tensorrt_llm/_torch/attention_backend/sparse/dsa/indexer.py b/tensorrt_llm/_torch/attention_backend/sparse/dsa/indexer.py index eae55ac98ae7..3e67594b1690 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/dsa/indexer.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/dsa/indexer.py @@ -21,7 +21,10 @@ from tensorrt_llm._torch.distributed.ops import allgather from tensorrt_llm._torch.modules.layer_norm import LayerNorm from tensorrt_llm._torch.modules.linear import Linear -from tensorrt_llm._torch.modules.multi_stream_utils import maybe_execute_in_parallel +from tensorrt_llm._torch.modules.multi_stream_utils import ( + do_multi_stream, + maybe_execute_in_parallel, +) from tensorrt_llm._torch.modules.rotary_embedding import RotaryEmbedding from tensorrt_llm._torch.utils import Fp4QuantizedTensor, maybe_compile from tensorrt_llm._utils import get_sm_version, maybe_pin_memory, prefer_pinned @@ -682,6 +685,9 @@ def __init__( self.use_fp4 = sparse_params.indexer_k_dtype == "fp4" self.aux_stream = aux_stream self.ln_events = [torch.cuda.Event(), torch.cuda.Event()] + # Fork/join events for the aux-stream prev_topk write-back. + self.prev_topk_copy_events = [torch.cuda.Event(), torch.cuda.Event()] + self._prev_topk_copy_pending = False self.use_cute_dsl_topk = sparse_params.use_cute_dsl_topk and IS_CUTLASS_DSL_AVAILABLE self.use_cute_dsl_paged_mqa_logits = ( sparse_params.use_cute_dsl_paged_mqa_logits and IS_CUTLASS_DSL_AVAILABLE @@ -1879,7 +1885,18 @@ def sparse_attn_indexer( local_layer = metadata.kv_cache_manager.layer_offsets[self.layer_idx] decode_topk = topk_indices_buffer[token_offset : token_offset + num_gen_tokens] last_mtp_topk = decode_topk[next_n - 1 :: next_n] - metadata.heuristic_prev_topk[local_layer, :num_generations].copy_(last_mtp_topk) + prev_topk_dst = metadata.heuristic_prev_topk[local_layer, :num_generations] + if do_multi_stream() and self.aux_stream is not None: + # Fork the write-back onto the aux stream to overlap with this + # layer's core sparse attention; joined by maybe_join_prev_topk_copy(). + self.prev_topk_copy_events[0].record() + with torch.cuda.stream(self.aux_stream): + self.prev_topk_copy_events[0].wait() + prev_topk_dst.copy_(last_mtp_topk) + self.prev_topk_copy_events[1].record() + self._prev_topk_copy_pending = True + else: + prev_topk_dst.copy_(last_mtp_topk) elif has_decode and metadata.skip_indexer_for_gen_reqs: # Fill topk_indices_buffer with pre-defined dense topk indices @@ -1929,6 +1946,12 @@ def _mtp_last_accepted_rows( offset = (gen_num_accepted - 1).clamp(0, next_n - 1) return gen_topk[base + offset] + def maybe_join_prev_topk_copy(self) -> None: + """Join the aux-stream heuristic prev_topk write-back, if forked.""" + if self._prev_topk_copy_pending: + self.prev_topk_copy_events[1].wait() + self._prev_topk_copy_pending = False + def _weight_scale(self, weights: torch.Tensor, q_scale: torch.Tensor) -> torch.Tensor: """Apply quantization scale to indexer attention weights.""" weights = _scale(weights, q_scale, self.weight_scale_factor) diff --git a/tensorrt_llm/_torch/attention_backend/sparse/dsa/module.py b/tensorrt_llm/_torch/attention_backend/sparse/dsa/module.py index 5202067f14a6..bf36a14655ed 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/dsa/module.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/dsa/module.py @@ -387,6 +387,11 @@ def _forward_dsa_attn( indexer_intermediates=indexer_intermediates, ) + # Join the prev_topk copy forked in sparse_attn_indexer; CUDA graph + # capture requires the join within the same layer's forward. + if self.mqa.indexer is not None: + self.mqa.indexer.maybe_join_prev_topk_copy() + def should_use_short_mha( self, attn_metadata: AttentionMetadata, position_ids: Optional[torch.Tensor]