diff --git a/python/sglang/srt/layers/attention/hybrid_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_attn_backend.py index f57a9410d..af5bd4de2 100644 --- a/python/sglang/srt/layers/attention/hybrid_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_attn_backend.py @@ -59,6 +59,10 @@ class HybridAttnBackend(AttentionBackend): backend = self._select_backend(forward_batch.forward_mode) backend.init_forward_metadata_out_graph(forward_batch, in_capture=in_capture) + def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch): + backend = self._select_backend(forward_batch.forward_mode) + backend.init_forward_metadata_in_graph(forward_batch) + def init_forward_metadata(self, forward_batch: ForwardBatch): backend = self._select_backend(forward_batch.forward_mode) backend.init_forward_metadata(forward_batch) diff --git a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py index 6c6f9911f..efefc6bb3 100644 --- a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py @@ -838,6 +838,10 @@ class HybridLinearAttnBackend(AttentionBackend): forward_batch, in_capture=in_capture ) + def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch): + for attn_backend in self.attn_backend_list: + attn_backend.init_forward_metadata_in_graph(forward_batch) + def init_forward_metadata(self, forward_batch: ForwardBatch): if forward_batch.forward_mode.is_draft_extend_v2(): # DRAFT_EXTEND_V2 runs only full-attn layers in the draft model; skip diff --git a/python/sglang/srt/layers/attention/triton_ops/trtllm_mha_graph_metadata.py b/python/sglang/srt/layers/attention/triton_ops/trtllm_mha_graph_metadata.py new file mode 100644 index 000000000..73a828658 --- /dev/null +++ b/python/sglang/srt/layers/attention/triton_ops/trtllm_mha_graph_metadata.py @@ -0,0 +1,198 @@ +"""Fused CUDA-graph metadata update for the TRTLLM MHA backend. + +`TRTLLMHAAttnBackend._apply_cuda_graph_metadata` used to rebuild the +page table(s) and seqlen buffers with ~25 small aten ops per graph +replay (index gathers, floor_divide, cumsum, dtype casts, copies). +On some CPUs that is ~0.7-1.0 ms of pure host dispatch, repeated 4x +per decode step (2 draft-decode steps + target-verify + draft-extend) +on every TP rank. The resulting per-rank CPU jitter skews the +cudaGraphLaunch across ranks and is paid as spin time inside the first +custom all-reduce of every replayed graph. + +This kernel performs the whole update in ONE launch: + - cache_seqlens[i] = seq_lens[i] + seqlen_offset (int32) + - cu_seqlens_k[1:] = cumsum(cache_seqlens) (int32) + - cu_seqlens_q[1:] = cumsum(qlens) or arange*q_stride (optional) + - page_table[i, p] = req_to_token[req_pool_indices[i], + p * page_size] // page_size + - swa_page_table = full_to_swa_mapping[token] // page_size (optional) + - swa_out_cache_loc = full_to_swa_mapping[out_cache_loc], zero padded + (optional) +""" + +import triton +import triton.language as tl + +# cu_seqlens_q handling inside the fused kernel +Q_MODE_NONE = 0 # cu_seqlens_q is preset (decode / target-verify) +Q_MODE_CUMSUM = 1 # cu_seqlens_q[1:] = cumsum(qlens) (draft-extend) +Q_MODE_STRIDED = 2 # cu_seqlens_q[1:] = arange*q_stride (draft-extend v2) + + +@triton.jit +def update_trtllm_mha_graph_metadata_kernel( + # inputs + req_pool_indices_ptr, # [bs] int + seq_lens_ptr, # [bs] int + req_to_token_ptr, # [pool_size, req_to_token_stride] int32 + swa_mapping_ptr, # [full_size + page_size + 1] int64, or None + out_cache_loc_ptr, # [num_out_tokens] int64, or None + qlens_ptr, # [bs] int, or None (Q_MODE_CUMSUM only) + # outputs + cache_seqlens_ptr, # [bs] int32 + cu_seqlens_k_ptr, # [bs + 1] int32 + cu_seqlens_q_ptr, # [bs + 1] int32, or None + page_table_ptr, # [bs, page_table_stride] int32 + swa_page_table_ptr, # [bs, swa_page_table_stride] int32, or None + swa_out_cache_loc_ptr, # [swa_out_len] int64, or None + # scalars + bs, + seqlen_offset, # added to seq_lens for cache_seqlens / cu_seqlens_k + max_seq_pages, # page-table columns to (re)write per row + q_stride, # Q_MODE_STRIDED stride + num_out_tokens, # valid prefix of out_cache_loc + swa_out_len, # full swa_out_cache_loc length (zero-padded tail) + req_to_token_stride, + page_table_stride, + swa_page_table_stride, + # constexpr + PAGE_SIZE: tl.constexpr, + HAS_SWA: tl.constexpr, + HAS_SWA_OUT: tl.constexpr, + Q_MODE: tl.constexpr, + PAGE_BLOCK: tl.constexpr, + BS_BLOCK: tl.constexpr, +): + pid = tl.program_id(axis=0) + + if pid < bs: + # One program per batch row: cache_seqlens + page table row(s). + req_pool_index = tl.load(req_pool_indices_ptr + pid).to(tl.int64) + seqlen = (tl.load(seq_lens_ptr + pid) + seqlen_offset).to(tl.int32) + tl.store(cache_seqlens_ptr + pid, seqlen) + + row_in = req_to_token_ptr + req_pool_index * req_to_token_stride + row_out = page_table_ptr + pid.to(tl.int64) * page_table_stride + if HAS_SWA: + swa_row_out = swa_page_table_ptr + pid.to(tl.int64) * swa_page_table_stride + for i in range(tl.cdiv(max_seq_pages, PAGE_BLOCK)): + page_idx = i * PAGE_BLOCK + tl.arange(0, PAGE_BLOCK) + mask = page_idx < max_seq_pages + token = tl.load( + row_in + page_idx.to(tl.int64) * PAGE_SIZE, mask=mask, other=0 + ) + tl.store(row_out + page_idx, token // PAGE_SIZE, mask=mask) + if HAS_SWA: + token64 = token.to(tl.int64) + # Real req_to_token slots are >=0; the token>=0 guard + other=-1 mirror + # the swa_out_cache_loc -1 sentinel (uniform handling, no wrap). + swa_token = tl.load( + swa_mapping_ptr + token64, mask=mask & (token64 >= 0), other=-1 + ) + swa_page = tl.where(swa_token < 0, -1, swa_token // PAGE_SIZE) + tl.store(swa_row_out + page_idx, swa_page.to(tl.int32), mask=mask) + elif pid == bs: + # Single program: cu_seqlens_k (+ optional cu_seqlens_q) cumsum. + offs = tl.arange(0, BS_BLOCK) + mask = offs < bs + seqlens = (tl.load(seq_lens_ptr + offs, mask=mask, other=0)).to(tl.int32) + seqlens = tl.where(mask, seqlens + seqlen_offset, 0) + tl.store(cu_seqlens_k_ptr + 1 + offs, tl.cumsum(seqlens, axis=0), mask=mask) + if Q_MODE == 1: # Q_MODE_CUMSUM + qlens = tl.load(qlens_ptr + offs, mask=mask, other=0).to(tl.int32) + qlens = tl.where(mask, qlens, 0) + tl.store(cu_seqlens_q_ptr + 1 + offs, tl.cumsum(qlens, axis=0), mask=mask) + if Q_MODE == 2: # Q_MODE_STRIDED + tl.store( + cu_seqlens_q_ptr + 1 + offs, + ((offs + 1) * q_stride).to(tl.int32), + mask=mask, + ) + else: + # Remaining programs: swa_out_cache_loc translate + zero padding. + if HAS_SWA_OUT: + out_idx = (pid - bs - 1) * PAGE_BLOCK + tl.arange(0, PAGE_BLOCK) + in_range = out_idx < swa_out_len + is_real = in_range & (out_idx < num_out_tokens) + loc = tl.load(out_cache_loc_ptr + out_idx, mask=is_real, other=0) + translated = tl.load( + swa_mapping_ptr + loc, mask=is_real & (loc >= 0), other=0 + ) + translated = tl.where(is_real & (loc < 0), -1, translated) + tl.store(swa_out_cache_loc_ptr + out_idx, translated, mask=in_range) + + +def update_trtllm_mha_graph_metadata( + *, + req_pool_indices, + seq_lens, + req_to_token, + cache_seqlens, + cu_seqlens_k, + page_table, + bs: int, + seqlen_offset: int, + max_seq_pages: int, + page_size: int, + swa_mapping=None, + swa_page_table=None, + out_cache_loc=None, + swa_out_cache_loc=None, + cu_seqlens_q=None, + qlens=None, + q_stride: int = 0, + q_mode: int = Q_MODE_NONE, +): + """Launch the fused metadata update (one kernel for the whole replay init).""" + if bs == 0: + return + + # Launch-block width: page-table columns each program writes per iteration + # (also the swa_out_cache_loc tile width). 512 keeps the per-program working + # set small enough to stay off the register-pressure / occupancy cliff while + # being wide enough to cover the static page-table width in few iterations. + PAGE_BLOCK = 512 + has_swa = swa_page_table is not None + has_swa_out = swa_out_cache_loc is not None + + swa_out_len = swa_out_cache_loc.shape[0] if has_swa_out else 0 + if has_swa_out and out_cache_loc is not None: + num_out_tokens = min(swa_out_len, out_cache_loc.shape[0]) + else: + num_out_tokens = 0 + if num_out_tokens == 0: + # All loads are masked out; pass a valid dummy pointer for codegen. + out_cache_loc = swa_out_cache_loc + + grid_extra = triton.cdiv(swa_out_len, PAGE_BLOCK) if has_swa_out else 0 + grid = (bs + 1 + grid_extra,) + + update_trtllm_mha_graph_metadata_kernel[grid]( + req_pool_indices, + seq_lens, + req_to_token, + swa_mapping, + out_cache_loc, + qlens, + cache_seqlens, + cu_seqlens_k, + cu_seqlens_q, + page_table, + swa_page_table, + swa_out_cache_loc, + bs, + seqlen_offset, + max_seq_pages, + q_stride, + num_out_tokens, + swa_out_len, + req_to_token.stride(0), + page_table.stride(0), + swa_page_table.stride(0) if has_swa else 0, + PAGE_SIZE=page_size, + HAS_SWA=has_swa, + HAS_SWA_OUT=has_swa_out, + Q_MODE=q_mode, + PAGE_BLOCK=PAGE_BLOCK, + BS_BLOCK=triton.next_power_of_2(bs), + ) diff --git a/python/sglang/srt/layers/attention/trtllm_mha_backend.py b/python/sglang/srt/layers/attention/trtllm_mha_backend.py index aa0ff2588..50ec976cf 100644 --- a/python/sglang/srt/layers/attention/trtllm_mha_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mha_backend.py @@ -19,6 +19,11 @@ from sglang.srt.layers.attention.flashinfer_backend import ( from sglang.srt.layers.attention.triton_ops.trtllm_fp8_kv_kernel import ( fused_fp8_set_kv_buffer, ) +from sglang.srt.layers.attention.triton_ops.trtllm_mha_graph_metadata import ( + Q_MODE_NONE, + Q_MODE_STRIDED, + update_trtllm_mha_graph_metadata, +) from sglang.srt.layers.attention.triton_ops.trtllm_mha_page_table import ( build_trtllm_mha_page_table, ) @@ -136,6 +141,16 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): # SWA hybrid models split the KV cache into full and SWA pools with # separate index spaces; SWA layers need a translated page_table. self._swa_kv_pool: Optional[SWAKVPool] = self._resolve_swa_kv_pool(model_runner) + # Raw full->swa index mapping tensor for the fused cuda-graph + # metadata kernel (gather + // page_size happen on device). + if self._swa_kv_pool is not None: + self._swa_full_to_swa_mapping = self._swa_kv_pool.full_to_swa_index_mapping + assert self._swa_full_to_swa_mapping is not None, ( + "SWA pool must register full_to_swa_index_mapping before " + "TRTLLMHAAttnBackend is constructed" + ) + else: + self._swa_full_to_swa_mapping = None # Static page-table width (upper bound). The CUDA-graph path builds the # page table on-device sized to this constant, so it never reads a runtime @@ -307,8 +322,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): } # SWA write-target buffer; bound as a [:num_tokens] view in - # _build_cuda_graph_metadata, refilled before each replay in - # init_forward_metadata_out_graph. + # _build_cuda_graph_metadata and refilled by the fused metadata kernel. self.cuda_graph_swa_out_cache_loc = ( torch.zeros(max_num_tokens, dtype=torch.int64, device=self.device) if self.use_sliding_window_kv_pool @@ -471,7 +485,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): ) self.draft_extend_metadata[bs] = metadata - # Bind the SWA write-target buffer slice (refilled at replay). + # Bind the SWA write-target buffer slice (refilled by in-graph metadata). if self.use_sliding_window_kv_pool: metadata.swa_out_cache_loc = self.cuda_graph_swa_out_cache_loc[:num_tokens] @@ -484,73 +498,86 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): seq_lens: torch.Tensor, forward_mode: ForwardMode, spec_info: Optional[SpecInput], + out_cache_loc: Optional[torch.Tensor] = None, ): """Shared capture+replay body for the cuda-graph init path. - Public entry: :py:meth:`init_forward_metadata_out_graph`. + One fused triton kernel (update_trtllm_mha_graph_metadata) rebuilds + cache_seqlens, cu_seqlens_k/q, the page table(s), and swa_out_cache_loc. + The previous aten-op implementation issued ~25 host dispatches per graph + replay, whose per-rank jitter was paid as spin time inside the first + all-reduce of every replayed graph. + + The page table is rewritten to the static ``max_num_pages`` width (the + same upper bound ``_fill_page_table_device`` uses); the kernel + bounds the actual KV reads by the on-device ``cache_seqlens``, so no + runtime host max / seq_lens_cpu D2H sync is needed. + + Public entry: :py:meth:`init_forward_metadata_in_graph`. """ seq_lens = seq_lens[:bs] req_pool_indices = req_pool_indices[:bs] - # The device-side build (_fill_page_table_device) sizes to the static - # max_num_pages and bounds the actual writes by cache_seqlens, so no - # runtime host max is needed. + metadata = None + seqlen_offset = 0 + cu_seqlens_q = None + qlens = None + q_stride = 0 + q_mode = Q_MODE_NONE if forward_mode.is_decode_or_idle(): if spec_info is not None: # Draft Decode # Here we only support topk = 1 for now. metadata = self.decode_cuda_graph_metadata[bs] - metadata.cache_seqlens_int32 = self.decode_cuda_graph_metadata[ - "cache_seqlens" - ][:bs] - metadata.cache_seqlens_int32.copy_( - seq_lens + self.speculative_step_id + 1 - ) + seqlen_offset = self.speculative_step_id + 1 else: # Normal Decode metadata = self.decode_cuda_graph_metadata[bs] - metadata.cache_seqlens_int32.copy_(seq_lens) - - metadata.cu_seqlens_k[1:].copy_( - torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32) - ) - self._fill_page_table_device( - metadata, req_pool_indices, metadata.cache_seqlens_int32 - ) elif forward_mode.is_target_verify(): # Here we only support topk = 1 for now. metadata = self.target_verify_metadata[bs] - metadata.cache_seqlens_int32.copy_(seq_lens + metadata.max_seq_len_q) - metadata.cu_seqlens_k[1:].copy_( - torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32) - ) - self._fill_page_table_device( - metadata, req_pool_indices, metadata.cache_seqlens_int32 - ) + seqlen_offset = metadata.max_seq_len_q elif forward_mode.is_draft_extend_v2(): metadata = self.draft_extend_metadata[bs] - metadata.cache_seqlens_int32.copy_(seq_lens) - metadata.cu_seqlens_k[1:].copy_( - torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32) - ) - num_tokens_per_bs = spec_info.num_tokens_per_req - if num_tokens_per_bs <= 0: - # Capture uses a synthetic EagleDraftExtendInput; infer the - # fixed V2 stride from the capture buffer when it is unset. - num_tokens_per_bs = int(spec_info.num_accept_tokens[:bs].max().item()) - metadata.max_seq_len_q = num_tokens_per_bs - metadata.cu_seqlens_q[1:].copy_( - torch.arange( - num_tokens_per_bs, - bs * num_tokens_per_bs + 1, - num_tokens_per_bs, - dtype=torch.int32, - device=metadata.cu_seqlens_q.device, - ) - ) - self._fill_page_table_device( - metadata, req_pool_indices, metadata.cache_seqlens_int32 + # Static per-request query width, fixed by the captured graph shape. + # Do not inspect replay-time tensors here; this body is recorded into + # the CUDA graph. + num_tokens_per_bs = metadata.max_seq_len_q + cu_seqlens_q = metadata.cu_seqlens_q + q_stride = num_tokens_per_bs + q_mode = Q_MODE_STRIDED + else: + raise ValueError( + "TRTLLM-MHA CUDA graph metadata build got an unsupported forward " + f"mode: {forward_mode}" ) + + assert metadata is not None + # Static upper-bound page-table width (see docstring); the kernel + # bounds real KV reads by cache_seqlens, so this is a fixed loop + # bound only — never a host max / seq_lens_cpu D2H sync. + max_seq_pages = self.max_num_pages + update_trtllm_mha_graph_metadata( + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + req_to_token=self.req_to_token, + cache_seqlens=metadata.cache_seqlens_int32, + cu_seqlens_k=metadata.cu_seqlens_k, + page_table=metadata.page_table, + bs=bs, + seqlen_offset=seqlen_offset, + max_seq_pages=max_seq_pages, + page_size=self.page_size, + swa_mapping=self._swa_full_to_swa_mapping, + swa_page_table=metadata.swa_page_table, + out_cache_loc=out_cache_loc, + swa_out_cache_loc=metadata.swa_out_cache_loc, + cu_seqlens_q=cu_seqlens_q, + qlens=qlens, + q_stride=q_stride, + q_mode=q_mode, + ) + self.forward_metadata = metadata def update_verify_buffers_to_fill_after_draft( @@ -598,44 +625,36 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): in_capture: bool = False, ): bs = forward_batch.batch_size - req_pool_indices = forward_batch.req_pool_indices - seq_lens = forward_batch.seq_lens - encoder_lens = forward_batch.encoder_lens forward_mode = forward_batch.forward_mode spec_info = forward_batch.spec_info if in_capture: num_tokens = forward_batch.positions.numel() self._build_cuda_graph_metadata( - bs, num_tokens, forward_mode, spec_info, seq_lens.device - ) - self._apply_cuda_graph_metadata( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - forward_mode=forward_mode, - spec_info=spec_info, - ) - else: - self._apply_cuda_graph_metadata( - bs=bs, - req_pool_indices=req_pool_indices, - seq_lens=seq_lens, - forward_mode=forward_mode, - spec_info=spec_info, + bs, num_tokens, forward_mode, spec_info, forward_batch.seq_lens.device ) - # Refill the SWA write-target buffer from the live out_cache_loc before - # replay (the per-bs metadata holds a view bound in _build). - if self.use_sliding_window_kv_pool and forward_batch.out_cache_loc is not None: - n = forward_batch.out_cache_loc.shape[0] - self.cuda_graph_swa_out_cache_loc[n:].zero_() - self.cuda_graph_swa_out_cache_loc[:n].copy_( - self.token_to_kv_pool.translate_loc_from_full_to_swa( - forward_batch.out_cache_loc - ) + if forward_mode.is_decode_or_idle(): + self.forward_metadata = self.decode_cuda_graph_metadata[bs] + elif forward_mode.is_target_verify(): + self.forward_metadata = self.target_verify_metadata[bs] + elif forward_mode.is_draft_extend_v2(): + self.forward_metadata = self.draft_extend_metadata[bs] + else: + raise ValueError( + f"Invalid forward mode: {forward_mode=} for CUDA Graph replay." ) + def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch): + self._apply_cuda_graph_metadata( + bs=forward_batch.batch_size, + req_pool_indices=forward_batch.req_pool_indices, + seq_lens=forward_batch.seq_lens, + forward_mode=forward_batch.forward_mode, + spec_info=forward_batch.spec_info, + out_cache_loc=forward_batch.out_cache_loc, + ) + def init_forward_metadata(self, forward_batch: ForwardBatch): """Initialize the metadata for a forward pass.""" @@ -980,3 +999,18 @@ class TRTLLMHAAttnMultiStepDraftBackend(FlashInferMultiStepDraftBackend): self.attn_backends[i].init_forward_metadata_out_graph( inner_fb, in_capture=in_capture ) + + def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch) -> None: + from sglang.srt.model_executor.forward_batch_info import build_inner_fb_view + + assert forward_batch.spec_info is not None + assert forward_batch.spec_info.is_draft_input() + + inner_fb = build_inner_fb_view( + forward_batch, + bs=forward_batch.batch_size, + forward_mode=ForwardMode.DECODE, + encoder_lens=forward_batch.encoder_lens, + ) + for i in range(self.speculative_num_steps - 1): + self.attn_backends[i].init_forward_metadata_in_graph(inner_fb) diff --git a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py index f573bfe2e..03662c5da 100644 --- a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py @@ -385,6 +385,8 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): ) def run_once(): + self.draft_extend_attn_backend.init_forward_metadata_in_graph(forward_batch) + # Clean intermediate result cache for DP attention forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None set_dp_buffer_len( diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py index 58ef342cb..4ad723bcc 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py @@ -294,6 +294,15 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner): ) def run_once(): + # Record the metadata rebuild against the committed target-prefix + # geometry (spec_info nulled → plain target-length decode), matching + # every other frozen-KV metadata init. Without the view, backends + # that key seqlen offsets off spec_info (trtllm_mha's draft-decode + # branch adds speculative_step_id + 1) bake a +1 offset into the + # captured graph and replay reads one extra, never-written KV slot. + with self.frozen_kv_mtp_worker._frozen_kv_target_view(forward_batch): + self.draft_attn_backend.init_forward_metadata_in_graph(forward_batch) + forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None set_dp_buffer_len( global_dp_buffer_len, diff --git a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index 5fc143e7d..1f325c779 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -337,6 +337,8 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): attn_backend = self.eagle_worker.draft_extend_attn_backend_list[self.step] def run_once(): + attn_backend.init_forward_metadata_in_graph(forward_batch) + # Clean intermediate result cache for DP attention forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None set_dp_buffer_len( diff --git a/test/registered/attention/test_trtllm_mha_graph_metadata.py b/test/registered/attention/test_trtllm_mha_graph_metadata.py new file mode 100644 index 000000000..eaf9b64fd --- /dev/null +++ b/test/registered/attention/test_trtllm_mha_graph_metadata.py @@ -0,0 +1,432 @@ +"""Correctness tests for the fused TRTLLM-MHA cuda-graph metadata kernel. + +Validates the single-launch triton kernel against a pure-aten reference that +mirrors the exact semantics of the triton port: cache_seqlens / cu_seqlens_k / +cu_seqlens_q (all 3 q-modes) / page_table / swa_page_table / swa_out_cache_loc, +with the SWA -1 sentinel guard. +""" + +from types import SimpleNamespace + +import pytest +import torch + +import sglang.srt.layers.attention.trtllm_mha_backend as trtllm_mha_backend +from sglang.srt.layers.attention.triton_ops.trtllm_mha_graph_metadata import ( + Q_MODE_CUMSUM, + Q_MODE_NONE, + Q_MODE_STRIDED, + update_trtllm_mha_graph_metadata, +) +from sglang.srt.layers.attention.trtllm_mha_backend import TRTLLMHAAttnBackend +from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.test.ci.ci_register import register_cuda_ci + +# trtllm_mha kernels are sm100-only; run this kernel-unit test on Blackwell. +register_cuda_ci(est_time=30, stage="base-b-kernel-unit", runner_config="4-gpu-b200") + +DEVICE = "cuda" +PAGE_SIZE = 128 + + +def _make_backend_for_hook_test(speculative_num_draft_tokens=None): + backend = TRTLLMHAAttnBackend.__new__(TRTLLMHAAttnBackend) + backend.device = torch.device("cpu") + backend.max_context_len = 1024 + backend.page_size = PAGE_SIZE + backend.max_num_pages = 8 + backend.req_to_token = torch.zeros(4, 1024, dtype=torch.int32) + backend.use_sliding_window_kv_pool = False + backend._swa_kv_pool = None + backend._swa_full_to_swa_mapping = None + backend.speculative_step_id = 0 + backend.speculative_num_draft_tokens = speculative_num_draft_tokens + backend.decode_cuda_graph_metadata = {} + backend.target_verify_metadata = {} + backend.draft_extend_metadata = {} + backend.init_cuda_graph_state(max_bs=4, max_num_tokens=16) + return backend + + +def test_cuda_graph_metadata_launch_runs_in_graph_hook(monkeypatch): + calls = [] + + def fake_update(**kwargs): + calls.append(kwargs) + + monkeypatch.setattr( + trtllm_mha_backend, "update_trtllm_mha_graph_metadata", fake_update + ) + backend = _make_backend_for_hook_test() + fb = SimpleNamespace( + batch_size=2, + req_pool_indices=torch.arange(2, dtype=torch.int64), + seq_lens=torch.ones(2, dtype=torch.int32), + forward_mode=ForwardMode.DECODE, + spec_info=None, + positions=torch.arange(2, dtype=torch.int64), + out_cache_loc=torch.arange(2, dtype=torch.int64), + ) + + backend.init_forward_metadata_out_graph(fb, in_capture=True) + assert calls == [] + assert backend.forward_metadata is backend.decode_cuda_graph_metadata[2] + + backend.init_forward_metadata_in_graph(fb) + assert len(calls) == 1 + assert calls[0]["out_cache_loc"] is fb.out_cache_loc + + calls.clear() + backend.init_forward_metadata_out_graph(fb) + assert calls == [] + assert backend.forward_metadata is backend.decode_cuda_graph_metadata[2] + + +def test_draft_extend_in_graph_uses_captured_static_q_stride(monkeypatch): + calls = [] + + def fake_update(**kwargs): + calls.append(kwargs) + + class ExplodingAcceptTokens: + def __getitem__(self, key): + raise AssertionError("in-graph metadata must not inspect accept tokens") + + monkeypatch.setattr( + trtllm_mha_backend, "update_trtllm_mha_graph_metadata", fake_update + ) + backend = _make_backend_for_hook_test(speculative_num_draft_tokens=4) + fb = SimpleNamespace( + batch_size=2, + req_pool_indices=torch.arange(2, dtype=torch.int64), + seq_lens=torch.ones(2, dtype=torch.int32), + forward_mode=ForwardMode.DRAFT_EXTEND_V2, + spec_info=SimpleNamespace( + num_tokens_per_req=0, + num_accept_tokens=ExplodingAcceptTokens(), + ), + positions=torch.arange(8, dtype=torch.int64), + out_cache_loc=torch.arange(8, dtype=torch.int64), + ) + + backend.init_forward_metadata_out_graph(fb, in_capture=True) + backend.init_forward_metadata_in_graph(fb) + + assert len(calls) == 1 + assert calls[0]["q_mode"] == Q_MODE_STRIDED + assert calls[0]["q_stride"] == 4 + + +def test_hybrid_wrappers_forward_in_graph_hook(): + """Hybrid wrappers must forward init_forward_metadata_in_graph to the + wrapped backend(s) — the inherited no-op would leave the fused metadata + rebuild out of the captured graph (stale page table on every replay).""" + from sglang.srt.layers.attention.hybrid_attn_backend import HybridAttnBackend + from sglang.srt.layers.attention.hybrid_linear_attn_backend import ( + HybridLinearAttnBackend, + ) + + def make_fake(name, calls): + return SimpleNamespace( + token_to_kv_pool=None, + req_to_token_pool=None, + needs_cpu_seq_lens=False, + init_forward_metadata_in_graph=lambda fb: calls.append(name), + ) + + fb = SimpleNamespace(forward_mode=ForwardMode.DECODE) + + calls = [] + hybrid = HybridAttnBackend( + SimpleNamespace( + kv_cache_dtype=torch.bfloat16, + token_to_kv_pool=None, + req_to_token_pool=None, + ), + prefill_backend=make_fake("prefill", calls), + decode_backend=make_fake("decode", calls), + ) + hybrid.init_forward_metadata_in_graph(fb) + assert calls == ["decode"] + + calls = [] + hybrid_linear = HybridLinearAttnBackend( + full_attn_backend=make_fake("full", calls), + linear_attn_backend=make_fake("linear", calls), + full_attn_layers=[0], + ) + hybrid_linear.init_forward_metadata_in_graph(fb) + assert calls == ["full", "linear"] + + +def test_metadata_update_records_inside_cuda_graph(): + if not torch.cuda.is_available(): + pytest.skip("CUDA required") + + backend = _make_backend_for_hook_test() + backend.device = torch.device(DEVICE) + backend.page_size = 2 + backend.max_num_pages = 4 + backend.req_to_token = torch.arange(32, dtype=torch.int32, device=DEVICE).reshape( + 4, 8 + ) + backend.init_cuda_graph_state(max_bs=2, max_num_tokens=2) + + fb = SimpleNamespace( + batch_size=2, + req_pool_indices=torch.arange(2, dtype=torch.int64, device=DEVICE), + seq_lens=torch.tensor([3, 4], dtype=torch.int32, device=DEVICE), + forward_mode=ForwardMode.DECODE, + spec_info=None, + positions=torch.arange(2, dtype=torch.int64, device=DEVICE), + out_cache_loc=torch.arange(2, dtype=torch.int64, device=DEVICE), + ) + + backend.init_forward_metadata_out_graph(fb, in_capture=True) + backend.init_forward_metadata_in_graph(fb) + torch.cuda.synchronize() + + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + backend.init_forward_metadata_in_graph(fb) + + fb.seq_lens.copy_(torch.tensor([5, 6], dtype=torch.int32, device=DEVICE)) + graph.replay() + torch.cuda.synchronize() + + torch.testing.assert_close( + backend.forward_metadata.cache_seqlens_int32, + torch.tensor([5, 6], dtype=torch.int32, device=DEVICE), + rtol=0, + atol=0, + ) + + +def _build_inputs(bs, pool_size, max_num_pages, max_seq_pages, seq_max, seed): + """Build random pool / indices / seq_lens consistent with backend buffers.""" + g = torch.Generator(device="cpu").manual_seed(seed) + req_to_token_stride = max_num_pages * PAGE_SIZE + # int32 token ids in [0, pool_token_cap); -1 allowed in unused tails. + pool_token_cap = pool_size * req_to_token_stride + req_to_token = torch.randint( + 0, + pool_token_cap, + (pool_size, req_to_token_stride), + generator=g, + dtype=torch.int32, + ).to(DEVICE) + req_pool_indices = torch.randperm(pool_size, generator=g)[:bs].to( + DEVICE, dtype=torch.int64 + ) + seq_lens = torch.randint(1, seq_max + 1, (bs,), generator=g, dtype=torch.int32).to( + DEVICE + ) + return req_to_token, req_pool_indices, seq_lens, req_to_token_stride, pool_token_cap + + +def _ref_cache_seqlens(seq_lens, seqlen_offset): + return (seq_lens.to(torch.int32) + seqlen_offset).to(torch.int32) + + +def _ref_page_table(req_to_token, req_pool_indices, max_seq_pages): + strided = torch.arange(0, max_seq_pages * PAGE_SIZE, PAGE_SIZE, device=DEVICE) + gathered = req_to_token[req_pool_indices[:, None], strided[None, :]] + return gathered // PAGE_SIZE, gathered + + +def _ref_swa_page_table(gathered_tokens, swa_mapping): + # mimic mapping[-1]=-1 sentinel: token<0 -> -1, else mapping[token]//page + tok = gathered_tokens.to(torch.int64) + safe = torch.where(tok >= 0, tok, torch.zeros_like(tok)) + swa_token = swa_mapping[safe] + swa_token = torch.where(tok >= 0, swa_token, torch.full_like(swa_token, -1)) + swa_page = torch.where( + swa_token < 0, + torch.full_like(swa_token, -1), + swa_token // PAGE_SIZE, + ) + return swa_page.to(torch.int32) + + +def _make_swa_mapping(pool_token_cap, seed): + g = torch.Generator(device="cpu").manual_seed(seed + 7) + # Random non-negative SWA pool ids, with a -1 sentinel appended (index -1). + mapping = torch.randint( + 0, pool_token_cap, (pool_token_cap + PAGE_SIZE + 1,), generator=g + ).to(DEVICE, dtype=torch.int64) + mapping[-1] = -1 # sentinel for wrapped -1 index + return mapping + + +@pytest.mark.parametrize("bs", [1, 3, 8, 17]) +@pytest.mark.parametrize("seqlen_offset", [0, 1, 4]) +@pytest.mark.parametrize("q_mode", [Q_MODE_NONE, Q_MODE_CUMSUM, Q_MODE_STRIDED]) +@pytest.mark.parametrize("with_swa", [False, True]) +# static_width=True exercises the production path: the backend passes the STATIC +# max_num_pages (full upper bound), not a per-batch dynamic width, so the kernel +# rewrites the whole page-table width every replay (tail pages beyond a request's +# seq are gathered but ignored by the attention kernel via cache_seqlens). The +# aten reference gathers the same full width, so this stays a bit-exact check. +@pytest.mark.parametrize("static_width", [False, True]) +def test_metadata_correctness(bs, seqlen_offset, q_mode, with_swa, static_width): + if not torch.cuda.is_available(): + pytest.skip("CUDA required") + + seed = ( + 1234 + + bs * 31 + + seqlen_offset * 7 + + q_mode * 3 + + int(with_swa) + + 1000 * int(static_width) + ) + pool_size = 64 + max_num_pages = 16 + seq_max = (max_num_pages - 2) * PAGE_SIZE # leave headroom for seqlen_offset + seq_max = min(seq_max, 1500) + + ( + req_to_token, + req_pool_indices, + seq_lens, + req_to_token_stride, + pool_token_cap, + ) = _build_inputs(bs, pool_size, max_num_pages, None, seq_max, seed) + + cache_seqlens_ref = _ref_cache_seqlens(seq_lens, seqlen_offset) + max_seq_len_k = int(cache_seqlens_ref.max().item()) + if static_width: + # Production passes the static upper bound (self.max_num_pages), not a + # dynamic per-batch width — write the whole table every replay. + max_seq_pages = max_num_pages + else: + max_seq_pages = (max_seq_len_k + PAGE_SIZE - 1) // PAGE_SIZE + + # Pre-allocated output buffers (mirror init_cuda_graph_state). + cache_seqlens = torch.zeros(bs, dtype=torch.int32, device=DEVICE) + cu_seqlens_k = torch.zeros(bs + 1, dtype=torch.int32, device=DEVICE) + page_table = torch.zeros(bs, max_num_pages, dtype=torch.int32, device=DEVICE) + + cu_seqlens_q = None + qlens = None + q_stride = 0 + if q_mode == Q_MODE_CUMSUM: + cu_seqlens_q = torch.zeros(bs + 1, dtype=torch.int32, device=DEVICE) + g = torch.Generator(device="cpu").manual_seed(seed + 99) + qlens = torch.randint(1, 6, (bs,), generator=g, dtype=torch.int32).to(DEVICE) + elif q_mode == Q_MODE_STRIDED: + cu_seqlens_q = torch.zeros(bs + 1, dtype=torch.int32, device=DEVICE) + q_stride = 4 + + swa_mapping = None + swa_page_table = None + swa_out_cache_loc = None + out_cache_loc = None + if with_swa: + swa_mapping = _make_swa_mapping(pool_token_cap, seed) + swa_page_table = torch.zeros( + bs, max_num_pages, dtype=torch.int32, device=DEVICE + ) + num_out = bs # one written token per request (decode-like) + swa_out_len = num_out + 5 # extra padding tail to validate zero-fill + swa_out_cache_loc = torch.full( + (swa_out_len,), 123, dtype=torch.int64, device=DEVICE + ) + g = torch.Generator(device="cpu").manual_seed(seed + 555) + out_cache_loc = torch.randint( + 0, pool_token_cap, (num_out,), generator=g, dtype=torch.int64 + ).to(DEVICE) + # inject a -1 entry to exercise the sentinel path + out_cache_loc[0] = -1 + + update_trtllm_mha_graph_metadata( + req_pool_indices=req_pool_indices, + seq_lens=seq_lens, + req_to_token=req_to_token, + cache_seqlens=cache_seqlens, + cu_seqlens_k=cu_seqlens_k, + page_table=page_table, + bs=bs, + seqlen_offset=seqlen_offset, + max_seq_pages=max_seq_pages, + page_size=PAGE_SIZE, + swa_mapping=swa_mapping, + swa_page_table=swa_page_table, + out_cache_loc=out_cache_loc, + swa_out_cache_loc=swa_out_cache_loc, + cu_seqlens_q=cu_seqlens_q, + qlens=qlens, + q_stride=q_stride, + q_mode=q_mode, + ) + torch.cuda.synchronize() + + # ---- cache_seqlens ---- + torch.testing.assert_close(cache_seqlens, cache_seqlens_ref, rtol=0, atol=0) + + # ---- cu_seqlens_k ---- + cu_k_ref = torch.zeros(bs + 1, dtype=torch.int32, device=DEVICE) + cu_k_ref[1:] = torch.cumsum(cache_seqlens_ref, dim=0, dtype=torch.int32) + torch.testing.assert_close(cu_seqlens_k, cu_k_ref, rtol=0, atol=0) + + # ---- page_table ---- + pt_ref, gathered = _ref_page_table(req_to_token, req_pool_indices, max_seq_pages) + torch.testing.assert_close(page_table[:, :max_seq_pages], pt_ref, rtol=0, atol=0) + + # ---- cu_seqlens_q ---- + if q_mode == Q_MODE_CUMSUM: + cu_q_ref = torch.zeros(bs + 1, dtype=torch.int32, device=DEVICE) + cu_q_ref[1:] = torch.cumsum(qlens, dim=0, dtype=torch.int32) + torch.testing.assert_close(cu_seqlens_q, cu_q_ref, rtol=0, atol=0) + elif q_mode == Q_MODE_STRIDED: + cu_q_ref = torch.zeros(bs + 1, dtype=torch.int32, device=DEVICE) + cu_q_ref[1:] = ( + torch.arange(1, bs + 1, device=DEVICE, dtype=torch.int32) * q_stride + ) + torch.testing.assert_close(cu_seqlens_q, cu_q_ref, rtol=0, atol=0) + + # ---- swa_page_table / swa_out_cache_loc ---- + if with_swa: + swa_pt_ref = _ref_swa_page_table(gathered, swa_mapping) + torch.testing.assert_close( + swa_page_table[:, :max_seq_pages], swa_pt_ref, rtol=0, atol=0 + ) + + # swa_out_cache_loc reference: translate real prefix, zero-fill tail. + num_out = out_cache_loc.shape[0] + swa_out_len = swa_out_cache_loc.shape[0] + num_real = min(num_out, swa_out_len) + out_ref = torch.zeros(swa_out_len, dtype=torch.int64, device=DEVICE) + loc = out_cache_loc[:num_real].to(torch.int64) + safe = torch.where(loc >= 0, loc, torch.zeros_like(loc)) + translated = swa_mapping[safe] + translated = torch.where(loc >= 0, translated, torch.full_like(translated, -1)) + out_ref[:num_real] = translated + torch.testing.assert_close(swa_out_cache_loc, out_ref, rtol=0, atol=0) + + +def test_bs_zero_noop(): + if not torch.cuda.is_available(): + pytest.skip("CUDA required") + # bs == 0 should be a no-op (early return). + cache_seqlens = torch.zeros(0, dtype=torch.int32, device=DEVICE) + cu_seqlens_k = torch.zeros(1, dtype=torch.int32, device=DEVICE) + page_table = torch.zeros(0, 4, dtype=torch.int32, device=DEVICE) + update_trtllm_mha_graph_metadata( + req_pool_indices=torch.zeros(0, dtype=torch.int64, device=DEVICE), + seq_lens=torch.zeros(0, dtype=torch.int32, device=DEVICE), + req_to_token=torch.zeros(4, 4 * PAGE_SIZE, dtype=torch.int32, device=DEVICE), + cache_seqlens=cache_seqlens, + cu_seqlens_k=cu_seqlens_k, + page_table=page_table, + bs=0, + seqlen_offset=0, + max_seq_pages=0, + page_size=PAGE_SIZE, + ) + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/test/registered/attention/unittests/dense/test_trtllm_mha.py b/test/registered/attention/unittests/dense/test_trtllm_mha.py index 7115fff1e..0bb4cc3d8 100644 --- a/test/registered/attention/unittests/dense/test_trtllm_mha.py +++ b/test/registered/attention/unittests/dense/test_trtllm_mha.py @@ -21,6 +21,7 @@ from sglang.test.kits.attention_unittest.runner_modes.cuda_graph_decode_runner i ) from sglang.test.kits.attention_unittest.runner_modes.speculative_draft_runner import ( run_dense_eagle_draft_cuda_graph_runner_case, + run_dense_frozen_kv_mtp_cuda_graph_runner_case, ) register_cuda_ci(est_time=20, stage="base-b", runner_config="4-gpu-b200") @@ -143,6 +144,20 @@ class TestTRTLLMMHADenseAttentionBackendCorrectness(CustomTestCase): ), ) + # Frozen-KV MTP draft CG runner (chain, topk=1) — records the fused + # in-graph metadata rebuild inside FrozenKVMTPCudaGraphRunner's capture. + FROZEN_KV_MTP_RUNNER_CASES = ( + DenseAttentionCase( + name="runner_frozen_kv_mtp_decode_trtllm_mha_cuda_graph", + backend="trtllm_mha", + forward_mode=ForwardMode.DECODE, + num_heads=4, + num_kv_heads=4, + page_size=16, + prefix_lens=(4, 7), + ), + ) + def test_projected_dense_decode_cases(self): for case in self.DECODE_CASES: with self.subTest(case=case.name, backend=case.backend): @@ -175,6 +190,16 @@ class TestTRTLLMMHADenseAttentionBackendCorrectness(CustomTestCase): hidden_size=self.HIDDEN_SIZE, ) + def test_runner_mode_frozen_kv_mtp_cuda_graph_runner_cases(self): + for case in self.FROZEN_KV_MTP_RUNNER_CASES: + with self.subTest(case=case.name, backend=case.backend): + run_dense_frozen_kv_mtp_cuda_graph_runner_case( + self, + case, + head_dim=self.HEAD_DIM, + hidden_size=self.HIDDEN_SIZE, + ) + if __name__ == "__main__": unittest.main()