diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index 8277a8d27..9283ebd9e 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -14,6 +14,9 @@ from sglang.srt.layers.attention.triton_ops.metadata import ( normal_decode_set_metadata, prepare_swa_spec_page_table_triton, ) +from sglang.srt.layers.attention.triton_ops.trtllm_mha_page_table import ( + build_trtllm_mha_page_table, +) from sglang.srt.layers.attention.utils import assert_buffer_fits from sglang.srt.layers.cp.base import CPAttentionBackendKind, get_cp_strategy from sglang.srt.layers.cp.utils import is_cp_v2_active @@ -26,7 +29,7 @@ from sglang.srt.mem_cache.memory_pool import KVWriteLoc from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.server_args import get_global_server_args -from sglang.srt.speculative.spec_info import SpecInput +from sglang.srt.speculative.spec_info import SpecInput, SpeculativeAlgorithm from sglang.srt.utils import get_compiler_backend if TYPE_CHECKING: @@ -241,6 +244,16 @@ class FlashAttentionBackend(AttentionBackend): self.kv_cache_dtype = model_runner.kv_cache_dtype self.kv_cache_dtype_str = model_runner.server_args.kv_cache_dtype self.page_size = model_runner.page_size + # Static page-table width (upper bound). The device-side page-table build + # sizes to this constant, so no runtime host max is needed. + self.max_num_pages = ( + self.max_context_len + self.page_size - 1 + ) // self.page_size + # Opt out of the seq_lens_cpu D2H only for dflash (the worker adapted to + # the GPU-only relay); EAGLE/MTP/standalone/non-spec keep the CPU mirror. + self.needs_cpu_seq_lens = not SpeculativeAlgorithm.from_string( + model_runner.server_args.speculative_algorithm + ).is_dflash() self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA self.skip_prefill = skip_prefill self.attn_cp_size = model_runner.attn_cp_size @@ -489,6 +502,13 @@ class FlashAttentionBackend(AttentionBackend): seqlens_in_batch = forward_batch.seq_lens batch_size = forward_batch.batch_size device = seqlens_in_batch.device + # Eager path needs a host int for dynamic page-table sizing: the CPU + # mirror when published, else a local D2H (not the overlap hot path). + seq_lens_cpu = ( + forward_batch.seq_lens_cpu + if forward_batch.seq_lens_cpu is not None + else seqlens_in_batch.cpu() + ) if forward_batch.forward_mode.is_decode_or_idle(): # Draft Decode @@ -497,7 +517,7 @@ class FlashAttentionBackend(AttentionBackend): metadata.cache_seqlens_int32 = ( seqlens_in_batch + (self.speculative_step_id + 1) ).to(torch.int32) - metadata.max_seq_len_k = forward_batch.seq_lens_cpu.max().item() + ( + metadata.max_seq_len_k = seq_lens_cpu.max().item() + ( self.speculative_step_id + 1 ) metadata.cu_seqlens_q = torch.arange( @@ -516,7 +536,7 @@ class FlashAttentionBackend(AttentionBackend): # Draft-extend's idle batch (padded for DP MLP-sync) has no # tree; build plain metadata (padded output is discarded). metadata.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32) - metadata.max_seq_len_k = forward_batch.seq_lens_cpu.max().item() + metadata.max_seq_len_k = seq_lens_cpu.max().item() metadata.cu_seqlens_q = torch.arange( 0, batch_size + 1, dtype=torch.int32, device=device ) @@ -529,7 +549,7 @@ class FlashAttentionBackend(AttentionBackend): else: metadata.cache_seqlens_int32 = (seqlens_in_batch).to(torch.int32) metadata.max_seq_len_q = self.topk - metadata.max_seq_len_k = forward_batch.seq_lens_cpu.max().item() + metadata.max_seq_len_k = seq_lens_cpu.max().item() metadata.cu_seqlens_q = torch.arange( 0, batch_size * self.topk + 1, @@ -579,7 +599,7 @@ class FlashAttentionBackend(AttentionBackend): else: # Normal Decode metadata.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32) - metadata.max_seq_len_k = forward_batch.seq_lens_cpu.max().item() + metadata.max_seq_len_k = seq_lens_cpu.max().item() metadata.cu_seqlens_q = torch.arange( 0, batch_size + 1, dtype=torch.int32, device=device ) @@ -625,8 +645,7 @@ class FlashAttentionBackend(AttentionBackend): ).to(torch.int32) metadata.max_seq_len_q = self.speculative_num_draft_tokens metadata.max_seq_len_k = ( - forward_batch.seq_lens_cpu.max().item() - + self.speculative_num_draft_tokens + seq_lens_cpu.max().item() + self.speculative_num_draft_tokens ) metadata.cu_seqlens_q = torch.arange( 0, @@ -649,7 +668,7 @@ class FlashAttentionBackend(AttentionBackend): else: metadata.cache_seqlens_int32 = forward_batch.seq_lens.to(torch.int32) metadata.max_seq_len_q = self.speculative_num_draft_tokens - metadata.max_seq_len_k = forward_batch.seq_lens_cpu.max().item() + metadata.max_seq_len_k = seq_lens_cpu.max().item() metadata.cu_seqlens_q = torch.arange( 0, batch_size * self.speculative_num_draft_tokens + 1, @@ -749,7 +768,7 @@ class FlashAttentionBackend(AttentionBackend): include_draft_extend_v2=True ): metadata.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32) - metadata.max_seq_len_k = forward_batch.seq_lens_cpu.max().item() + metadata.max_seq_len_k = seq_lens_cpu.max().item() metadata.cu_seqlens_k = torch.nn.functional.pad( torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0) ) @@ -797,7 +816,7 @@ class FlashAttentionBackend(AttentionBackend): self._pa_swa_prefill_lens[ forward_batch.req_pool_indices[:batch_size] ] = forward_batch.seq_lens[:batch_size].to(torch.int32) - max_pf = int(forward_batch.seq_lens_cpu[:batch_size].max().item()) + max_pf = int(seq_lens_cpu[:batch_size].max().item()) if max_pf > self._pa_swa_max_prefill_len: self._pa_swa_max_prefill_len = max_pf @@ -2257,6 +2276,16 @@ class FlashAttentionBackend(AttentionBackend): return metadata, metadata_expand + @staticmethod + def _host_max_seq_len( + seq_lens_cpu: Optional[torch.Tensor], seq_lens: torch.Tensor + ) -> int: + """Host-side max KV length: the CPU mirror when published, else a local + D2H. For cold paths (topk>1, draft-extend, eager) that need a host max -- + not the dflash hot path (topk=1, device-side build).""" + src = seq_lens_cpu if seq_lens_cpu is not None else seq_lens.cpu() + return src.max().item() + def _apply_cuda_graph_metadata( self, bs: int, @@ -2277,7 +2306,10 @@ class FlashAttentionBackend(AttentionBackend): are gone. """ seq_lens = seq_lens[:bs] - seq_lens_cpu = seq_lens_cpu[:bs] + # The GPU-only path passes seq_lens_cpu=None; the topk>1 branches below + # still need a host max, so sync locally in that case (not the dflash + # overlap hot path, which uses topk=1 and the device-side build). + seq_lens_cpu = seq_lens_cpu[:bs] if seq_lens_cpu is not None else None req_pool_indices = req_pool_indices[:bs] device = seq_lens.device metadata = None @@ -2298,17 +2330,9 @@ class FlashAttentionBackend(AttentionBackend): if self.topk <= 1: # When topk = 1, we use the normal decode metadata metadata = self.decode_cuda_graph_metadata[bs] - max_len = seq_lens_cpu.max().item() - metadata.max_seq_len_k = max_len + self.speculative_step_id + 1 - max_seq_pages = ( - metadata.max_seq_len_k + self.page_size - 1 - ) // self.page_size - - assert_buffer_fits( - max_seq_pages, - metadata.page_table.shape[1], - "FA3 draft-decode page_table", - ) + # Page table built on-device (self-guards on cache_seqlens); + # max_seq_len_k left unset -- unread here (scheduler_metadata + # is normal-decode-only). normal_decode_set_metadata( metadata.cache_seqlens_int32, metadata.cu_seqlens_k, @@ -2316,7 +2340,7 @@ class FlashAttentionBackend(AttentionBackend): self.req_to_token, req_pool_indices, self.decode_cuda_graph_metadata["strided_indices"], - max_seq_pages, + self.max_num_pages, seq_lens, self.speculative_step_id + 1, self.page_size, @@ -2341,7 +2365,9 @@ class FlashAttentionBackend(AttentionBackend): # metadata.cu_seqlens_q already set in capture # metadata.cu_seqlens_k is not needed - metadata.max_seq_len_k = seq_lens_cpu.max().item() + metadata.max_seq_len_k = self._host_max_seq_len( + seq_lens_cpu, seq_lens + ) max_seq_pages = ( metadata.max_seq_len_k + self.page_size - 1 ) // self.page_size @@ -2389,16 +2415,11 @@ class FlashAttentionBackend(AttentionBackend): else: # Normal Decode metadata = self.decode_cuda_graph_metadata[bs] - max_len = seq_lens_cpu.max().item() - max_seq_pages = (max_len + self.page_size - 1) // self.page_size - metadata.max_seq_len_k = max_len - - assert_buffer_fits( - max_seq_pages, - metadata.page_table.shape[1], - "FA3 decode page_table", - ) if self.is_prefill_aware_swa: + # Prefill-aware SWA still needs a host max to bound the + # per-batch page table built below. + max_len = self._host_max_seq_len(seq_lens_cpu, seq_lens) + metadata.max_seq_len_k = max_len pa_max_len = min( self._pa_swa_max_prefill_len + self.sliding_window_size, max_len, @@ -2417,6 +2438,14 @@ class FlashAttentionBackend(AttentionBackend): dst_kv_lens=metadata.cache_seqlens_int32, ) else: + # Page table uses the static max_num_pages bound (no D2H). + # max_seq_len_k only feeds scheduler_metadata below, so use + # the free CPU mirror for a tight split heuristic when present. + metadata.max_seq_len_k = ( + seq_lens_cpu.max().item() + if seq_lens_cpu is not None + else self.max_context_len + ) normal_decode_set_metadata( metadata.cache_seqlens_int32, metadata.cu_seqlens_k, @@ -2424,7 +2453,7 @@ class FlashAttentionBackend(AttentionBackend): self.req_to_token, req_pool_indices, self.decode_cuda_graph_metadata["strided_indices"], - max_seq_pages, + self.max_num_pages, seq_lens, 0, self.page_size, @@ -2464,40 +2493,33 @@ class FlashAttentionBackend(AttentionBackend): (seq_lens + self.speculative_num_draft_tokens) ) - metadata.max_seq_len_k = ( - seq_lens_cpu.max().item() + self.speculative_num_draft_tokens - ) + # Page table built on-device (self-guards on cache_seqlens); + # max_seq_len_k left unset -- unread here (scheduler_metadata is + # normal-decode-only). metadata.cu_seqlens_k[1:].copy_( torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32) ) - max_seq_pages = ( - metadata.max_seq_len_k + self.page_size - 1 - ) // self.page_size - page_indices = self.req_to_token[ - req_pool_indices[:, None], - self.decode_cuda_graph_metadata["strided_indices"][:max_seq_pages], - ] - if ( - self.use_sliding_window_kv_pool - and metadata.swa_page_table is not None - ): - swa_page_indices = ( - self.token_to_kv_pool.translate_loc_from_full_to_swa( - page_indices - ) - ) - metadata.swa_page_table[:, :max_seq_pages].copy_( - swa_page_indices // self.page_size - ) - page_indices //= self.page_size - metadata.page_table[:, :max_seq_pages].copy_(page_indices) + has_swa = self.use_sliding_window_kv_pool + build_trtllm_mha_page_table( + req_to_token=self.req_to_token, + req_pool_indices=req_pool_indices, + cache_seqlens=metadata.cache_seqlens_int32, + page_table=metadata.page_table, + page_size=self.page_size, + swa_page_table=metadata.swa_page_table if has_swa else None, + full_to_swa=( + self.token_to_kv_pool.full_to_swa_index_mapping + if has_swa + else None + ), + ) else: # When topk > 1, we need two specific target verify metadata, and then merge states # 1. The first half of metadata for prefix tokens metadata = self.target_verify_metadata_topk_normal[bs] metadata.cache_seqlens_int32.copy_(seq_lens) # metadata.max_seq_len_q = self.speculative_num_draft_tokens, already set in capture - metadata.max_seq_len_k = seq_lens_cpu.max().item() + metadata.max_seq_len_k = self._host_max_seq_len(seq_lens_cpu, seq_lens) # metadata.cu_seqlens_q already set in capture metadata.cu_seqlens_k[1:].copy_( torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32) @@ -2583,7 +2605,7 @@ class FlashAttentionBackend(AttentionBackend): metadata = self.draft_extend_metadata[bs] metadata.cache_seqlens_int32.copy_(seq_lens) - metadata.max_seq_len_k = seq_lens_cpu.max().item() + metadata.max_seq_len_k = self._host_max_seq_len(seq_lens_cpu, seq_lens) metadata.cu_seqlens_k[1:].copy_( torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32) ) diff --git a/test/registered/attention/test_trtllm_mha_page_table.py b/test/registered/attention/test_trtllm_mha_page_table.py index 14d2de417..f10a680b7 100644 --- a/test/registered/attention/test_trtllm_mha_page_table.py +++ b/test/registered/attention/test_trtllm_mha_page_table.py @@ -6,6 +6,11 @@ it never reads a runtime max (no D2H sync). This test checks the device build is bit-identical to the legacy gather for the columns each request uses, for both the full page table and the SWA-translated page table, across context lengths, page sizes, and batch sizes. + +It also pins the invariant that lets the no-host-max (GPU-only) path hand the +kernel a static ``max_num_pages``-wide buffer: every column past a request's +page count must be left untouched, i.e. the kernel bounds its writes by the +device-side ``cache_seqlens`` alone. """ import unittest @@ -20,7 +25,7 @@ from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.test_utils import CustomTestCase # Triton kernel unit test for the trtllm_mha device-side page-table build. -register_cuda_ci(est_time=10, stage="base-b", runner_config="1-gpu-small") +register_cuda_ci(est_time=14, stage="base-b", runner_config="1-gpu-small") def _build_page_table_reference( @@ -143,6 +148,82 @@ class TestTrtllmMhaPageTable(CustomTestCase): max_ctx, page_size, num_reqs=max(64, bs), bs=bs, swa=True ) + def _run_self_guard_case(self, max_context_len, page_size, bs, swa=False): + """Short sequences against a full static buffer -- the GPU-only shape. + + Pre-fill the page table with a sentinel and run the kernel with the + static ``max_num_pages`` width (no host max to tighten it). Used columns + must hold the right block ids; every tail column must keep the sentinel, + proving the kernel never writes past the device-side ``cache_seqlens``. + """ + torch.manual_seed(0) + dev = "cuda" + num_reqs = max(64, bs) + max_num_pages = (max_context_len + page_size - 1) // page_size + n_slots = num_reqs * max_context_len + req_to_token = torch.randint( + 0, n_slots, (num_reqs, max_context_len), dtype=torch.int32, device=dev + ) + req_pool_indices = torch.randperm(num_reqs, device=dev)[:bs].to(torch.int32) + # Cap lengths well below max_context_len so most tail columns stay unused. + hi = max(2, max_context_len // 8) + cache_seqlens = torch.randint(1, hi + 1, (bs,), dtype=torch.int32, device=dev) + full_to_swa = ( + torch.randint(0, n_slots, (n_slots,), dtype=torch.int32, device=dev) + if swa + else None + ) + + SENTINEL = -1 + page_table = torch.full( + (bs, max_num_pages), SENTINEL, dtype=torch.int32, device=dev + ) + swa_page_table = ( + torch.full((bs, max_num_pages), SENTINEL, dtype=torch.int32, device=dev) + if swa + else None + ) + build_trtllm_mha_page_table( + req_to_token=req_to_token, + req_pool_indices=req_pool_indices, + cache_seqlens=cache_seqlens, + page_table=page_table, + page_size=page_size, + swa_page_table=swa_page_table, + full_to_swa=full_to_swa, + ) + pt_ref, swa_ref = _build_page_table_reference( + req_to_token, req_pool_indices, cache_seqlens, page_size, full_to_swa + ) + + tag = f"max_ctx={max_context_len} page_size={page_size} bs={bs} swa={swa}" + for i in range(bs): + npages = (int(cache_seqlens[i].item()) + page_size - 1) // page_size + self.assertTrue( + torch.equal(page_table[i, :npages], pt_ref[i, :npages]), + f"used-column mismatch req={i} {tag}", + ) + self.assertTrue( + torch.all(page_table[i, npages:] == SENTINEL), + f"kernel wrote past cache_seqlens req={i} npages={npages} {tag}", + ) + if swa: + self.assertTrue( + torch.equal(swa_page_table[i, :npages], swa_ref[i, :npages]), + f"swa used-column mismatch req={i} {tag}", + ) + self.assertTrue( + torch.all(swa_page_table[i, npages:] == SENTINEL), + f"swa wrote past cache_seqlens req={i} npages={npages} {tag}", + ) + + def test_writes_bounded_by_cache_seqlens(self): + for max_ctx in (4096, 131072): + for page_size in (1, 64, 256): + for bs in (1, 8): + self._run_self_guard_case(max_ctx, page_size, bs=bs) + self._run_self_guard_case(max_ctx, page_size, bs=bs, swa=True) + if __name__ == "__main__": unittest.main()