From 52d3f9b6d502fc5f01f9485efce116a4d1351fb6 Mon Sep 17 00:00:00 2001 From: longcheng-nv <243710427+longcheng-nv@users.noreply.github.com> Date: Sat, 8 Aug 2026 07:36:47 +0000 Subject: [PATCH] [None][perf] Overlap DSA heuristic prev_topk write-back on the aux stream The per-layer heuristic top-k feedback copy (this step's decode top-k -> next step's pre_idx hint) is a strided gather sitting on the main stream's critical path, once per indexer layer per decode step. Nothing in the current step consumes it, so fork it onto the Indexer's existing aux stream right after the top-k kernel and join it in the same layer's MLA forward once core sparse attention is enqueued -- the copy overlaps with the layer's heaviest decode work. Same-layer fork/join keeps CUDA graph capture free of unjoined forks (cudaStreamEndCapture rejects them) and restores ordering before the next layer overwrites the shared topk_indices_buffer rows the copy reads. Source and destination are persistent stable-address buffers, so replays stay valid with no record_stream bookkeeping. The fork engages only under do_multi_stream() (i.e. inside CUDA graph capture, where replay makes the stream/event host overhead free); eager execution keeps the original inline copy unchanged. Validated with a pattern-level CUDA graph smoke test: capture with the join succeeds, replayed feedback values are step-correct, and capture without the join fails with cudaErrorStreamCaptureUnjoined. Ported onto the sparse-attention framework refactor (#12733): the Indexer changes moved from sparse/dsa.py to sparse/dsa/indexer.py, and the two MLA join sites moved from modules/mla.py to sparse/dsa/module.py (_forward_dsa_attn) and sparse/deepseek_v4/module.py (forward_sparse_attn). Made-with: Claude Code (Fable 5) Co-Authored-By: Claude Fable 5 Signed-off-by: longcheng-nv <243710427+longcheng-nv@users.noreply.github.com> --- .../sparse/deepseek_v4/module.py | 7 +++ .../attention_backend/sparse/dsa/indexer.py | 49 ++++++++++++++++++- .../attention_backend/sparse/dsa/module.py | 7 +++ 3 files changed, 61 insertions(+), 2 deletions(-) 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 7d0ee04aff18..85f808d8fa03 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/module.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/deepseek_v4/module.py @@ -802,6 +802,13 @@ def _indexer_branch(): sparse_epilogue_output=sparse_epilogue_output, ) + # Join the aux-stream heuristic prev_topk write-back forked in + # sparse_attn_indexer, now that this layer's core attention is + # enqueued (the copy overlaps with it). Must stay within this + # layer's forward: CUDA graph capture rejects unjoined forks. + 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 ded1bda7694f..5d9cd9042a20 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 @@ -679,6 +682,11 @@ 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 pair for the aux-stream heuristic prev_topk write-back: + # [0] orders the copy after the top-k kernel, [1] is waited on by + # maybe_join_prev_topk_copy() in the owning MLA layer. + 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 @@ -1846,7 +1854,29 @@ 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 so the strided + # gather copy overlaps with this layer's core sparse + # attention instead of sitting on the critical path. + # Nothing in this step reads it back — the next consumer + # is the next decode step's pre_idx for this same layer. + # Source and destination are persistent stable-address + # buffers, so no record_stream bookkeeping is needed. The + # fork MUST be joined within this layer's forward via + # maybe_join_prev_topk_copy(): that both keeps CUDA graph + # capture free of unjoined forks (cudaStreamEndCapture + # rejects them) and restores ordering before the next + # layer overwrites the shared topk_indices_buffer rows + # this copy reads. + 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 @@ -1896,6 +1926,21 @@ 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. + + Called by the owning MLA layer after this layer's core sparse + attention has been enqueued, so the copy forked in + sparse_attn_indexer overlaps with it. Joining within the same + layer's forward keeps every fork matched with a join inside a + single captured region (CUDA graph capture rejects unjoined + forks) and orders the copy's read of topk_indices_buffer before + the next layer's indexer overwrites those rows. + """ + 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 5f07977a505d..624d57f1c560 100644 --- a/tensorrt_llm/_torch/attention_backend/sparse/dsa/module.py +++ b/tensorrt_llm/_torch/attention_backend/sparse/dsa/module.py @@ -387,6 +387,13 @@ def _forward_dsa_attn( indexer_intermediates=indexer_intermediates, ) + # Join the aux-stream heuristic prev_topk write-back forked in + # sparse_attn_indexer, now that this layer's core attention is + # enqueued (the copy overlaps with it). Must stay within this + # layer's forward: CUDA graph capture rejects unjoined forks. + 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]