diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index eb4ffa323..f5a96bf9b 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -256,15 +256,16 @@ class FlashAttentionBackend(AttentionBackend): 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/dspark (their workers - # 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_family() + # Page table is built on-device (build_trtllm_mha_page_table) and the + # tree-mask scratch is preallocated (get_verify_buffers_to_fill_after_draft), + # so no seq_lens_cpu / seq_lens_sum D2H sync is ever needed. + self.needs_cpu_seq_lens = False 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 + # Preallocated FULL_MASK tree-mask scratch; lets build_tree_kernel_efficient + # avoid the seq_lens_sum D2H sync (see get_verify_buffers_to_fill_after_draft). + self.cuda_graph_custom_mask = None self.use_sliding_window_kv_pool = ( isinstance(model_runner.token_to_kv_pool, SWAKVPool) @@ -589,12 +590,16 @@ 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() + # Eager (non-cuda-graph) path: max_seq_len_k only feeds Python-side + # page-table slicing and the scheduler_metadata heuristic -- never the + # kernel. Use the CPU mirror when published; otherwise the static + # max_context_len bound (over-wide page table, slightly suboptimal + # split heuristic, but no D2H sync). + seq_lens_cpu = forward_batch.seq_lens_cpu + eager_max_k = ( + seq_lens_cpu.max().item() + if seq_lens_cpu is not None + else self.max_context_len ) if forward_batch.forward_mode.is_decode_or_idle(): @@ -604,7 +609,7 @@ class FlashAttentionBackend(AttentionBackend): metadata.cache_seqlens_int32 = ( seqlens_in_batch + (self.speculative_step_id + 1) ).to(torch.int32) - metadata.max_seq_len_k = seq_lens_cpu.max().item() + ( + metadata.max_seq_len_k = eager_max_k + ( self.speculative_step_id + 1 ) metadata.cu_seqlens_q = torch.arange( @@ -623,7 +628,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 = seq_lens_cpu.max().item() + metadata.max_seq_len_k = eager_max_k metadata.cu_seqlens_q = torch.arange( 0, batch_size + 1, dtype=torch.int32, device=device ) @@ -636,7 +641,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 = seq_lens_cpu.max().item() + metadata.max_seq_len_k = eager_max_k metadata.cu_seqlens_q = torch.arange( 0, batch_size * self.topk + 1, @@ -686,7 +691,7 @@ class FlashAttentionBackend(AttentionBackend): else: # Normal Decode metadata.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32) - metadata.max_seq_len_k = seq_lens_cpu.max().item() + metadata.max_seq_len_k = eager_max_k metadata.cu_seqlens_q = torch.arange( 0, batch_size + 1, dtype=torch.int32, device=device ) @@ -753,7 +758,7 @@ class FlashAttentionBackend(AttentionBackend): ).to(torch.int32) metadata.max_seq_len_q = self.speculative_num_draft_tokens metadata.max_seq_len_k = ( - seq_lens_cpu.max().item() + self.speculative_num_draft_tokens + eager_max_k + self.speculative_num_draft_tokens ) metadata.cu_seqlens_q = torch.arange( 0, @@ -776,7 +781,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 = seq_lens_cpu.max().item() + metadata.max_seq_len_k = eager_max_k metadata.cu_seqlens_q = torch.arange( 0, batch_size * self.speculative_num_draft_tokens + 1, @@ -876,7 +881,7 @@ class FlashAttentionBackend(AttentionBackend): include_draft_extend_v2=True ): metadata.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32) - metadata.max_seq_len_k = seq_lens_cpu.max().item() + metadata.max_seq_len_k = eager_max_k metadata.cu_seqlens_k = torch.nn.functional.pad( torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0) ) @@ -913,7 +918,9 @@ class FlashAttentionBackend(AttentionBackend): torch.cumsum(extend_seq_lens, dim=0, dtype=torch.int32), (1, 0) ) else: - metadata.max_seq_len_q = metadata.max_seq_len_k + # max_seq_len_q reaches the kernel -- needs the real host max, + # not the static max_seq_len_k bound. + metadata.max_seq_len_q = max(forward_batch.extend_seq_lens_cpu) metadata.cu_seqlens_q = metadata.cu_seqlens_k # Setup local attention if enabled @@ -924,7 +931,12 @@ 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(seq_lens_cpu[:batch_size].max().item()) + if seq_lens_cpu is not None: + max_pf = int(seq_lens_cpu[:batch_size].max().item()) + else: + # Ratchet needs a true upper bound; a local D2H beats + # poisoning it with max_context_len forever. + max_pf = int(forward_batch.seq_lens[:batch_size].max().item()) if max_pf > self._pa_swa_max_prefill_len: self._pa_swa_max_prefill_len = max_pf @@ -2040,6 +2052,18 @@ class FlashAttentionBackend(AttentionBackend): ), } + # Worst-case FULL_MASK tree-mask scratch (bool). build_tree_kernel + # fills it in-place, so the GPU-only path needs no seq_lens_sum. + # Costs max_num_tokens * max_context_len bytes (can reach 100s of + # MB at long context) and is fully memset every verify step. + if not self.skip_prefill: + self.cuda_graph_custom_mask = torch.zeros( + max_num_tokens + * (self.max_context_len + self.speculative_num_draft_tokens), + dtype=torch.bool, + device=self.device, + ) + self.draft_extend_metadata = { "cache_seqlens": torch.zeros( max_bs, dtype=torch.int32, device=self.device @@ -2143,9 +2167,11 @@ class FlashAttentionBackend(AttentionBackend): dtype=torch.int32, device=self.device, ), + # Width covers the static max_seq_len_k bound + draft + # columns (checked by assert_buffer_fits at the merge). "page_table": torch.zeros( max_bs * self.speculative_num_draft_tokens, - self.max_context_len, + self.max_context_len + self.speculative_num_draft_tokens, dtype=torch.int32, device=self.device, ), @@ -2384,13 +2410,19 @@ class FlashAttentionBackend(AttentionBackend): return metadata, metadata_expand + def get_verify_buffers_to_fill_after_draft(self): + # Return the preallocated FULL_MASK tree-mask scratch so that + # build_tree_kernel_efficient fills it in-place and the worker never + # needs seq_lens_sum to size a dynamic allocation (no D2H sync). + return [self.cuda_graph_custom_mask, None] + @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).""" + D2H. Only for paths that accept a host sync (currently prefill-aware + SWA decode replay).""" src = seq_lens_cpu if seq_lens_cpu is not None else seq_lens.cpu() return src.max().item() @@ -2473,8 +2505,11 @@ class FlashAttentionBackend(AttentionBackend): # metadata.cu_seqlens_q already set in capture # metadata.cu_seqlens_k is not needed - metadata.max_seq_len_k = self._host_max_seq_len( - seq_lens_cpu, seq_lens + # Tight page-table bound when the CPU mirror is free. + metadata.max_seq_len_k = ( + seq_lens_cpu.max().item() + if seq_lens_cpu is not None + else self.max_context_len ) max_seq_pages = ( metadata.max_seq_len_k + self.page_size - 1 @@ -2636,7 +2671,12 @@ class FlashAttentionBackend(AttentionBackend): 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 = self._host_max_seq_len(seq_lens_cpu, seq_lens) + # Tight page-table bound when the CPU mirror is free. + metadata.max_seq_len_k = ( + seq_lens_cpu.max().item() + if seq_lens_cpu is not None + else self.max_context_len + ) # 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) @@ -2722,7 +2762,12 @@ class FlashAttentionBackend(AttentionBackend): metadata = self.draft_extend_metadata[bs] metadata.cache_seqlens_int32.copy_(seq_lens) - metadata.max_seq_len_k = self._host_max_seq_len(seq_lens_cpu, seq_lens) + # Tight page-table bound when the CPU mirror is free. + metadata.max_seq_len_k = ( + seq_lens_cpu.max().item() + if seq_lens_cpu is not None + else self.max_context_len + ) metadata.cu_seqlens_k[1:].copy_( torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32) ) 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 2af456a66..ac5e11f59 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 @@ -391,7 +391,13 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): post_warmup_hook=post_warmup_hook, ) - def replay(self, bs: int, seq_lens_sum: int, spec_info: EagleDraftExtendInput): + def replay( + self, + bs: int, + seq_lens_sum: Optional[int], + spec_info: EagleDraftExtendInput, + seq_lens_cpu: Optional[torch.Tensor], + ): """Init this step's attention metadata for the prepared bucket and replay its graph. Buffers must already be populated by the composite runner's ``prepare`` (step 0) or by the previous step's in-graph chain @@ -411,7 +417,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): req_pool_indices=buffers.req_pool_indices, seq_lens=buffers.seq_lens, seq_lens_sum=seq_lens_sum, - seq_lens_cpu=buffers.seq_lens_cpu, + seq_lens_cpu=seq_lens_cpu, encoder_lens=None, # per-step write target; out_cache_loc is frozen at prepare() time. out_cache_loc=buffers.out_cache_loc[:num_tokens], @@ -646,10 +652,15 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner: forward_batch.spec_info.num_accept_tokens ) + # Refresh the host mirror only when published; hand replay None + # otherwise so no consumer reads a stale buffer. if forward_batch.seq_lens_cpu is not None: if bs != raw_bs: buffers.seq_lens_cpu.fill_(self.seq_len_fill_value) buffers.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu) + self.seq_lens_cpu = buffers.seq_lens_cpu + else: + self.seq_lens_cpu = None # select_index[i] = i * window + num_correct_drafts[i]: the flat index # of request i's last accepted token. Used by the in-graph top-k gather @@ -681,9 +692,10 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner: self.raw_bs = raw_bs self.bs = bs self.raw_num_tokens = num_tokens - self.seq_lens_sum = ( - forward_batch.seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value - ) + seq_lens_sum = forward_batch.seq_lens_sum + if seq_lens_sum is not None: + seq_lens_sum = seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value + self.seq_lens_sum = seq_lens_sum self._prepare_extra(forward_batch) @@ -693,7 +705,9 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner: batch size.""" runner = self.runners[step] runner.raw_bs = self.raw_bs - out = runner.replay(self.bs, self.seq_lens_sum, self._replay_spec_info) + out = runner.replay( + self.bs, self.seq_lens_sum, self._replay_spec_info, self.seq_lens_cpu + ) raw_bs = self.raw_bs raw_num_tokens = self.raw_num_tokens logits_output = LogitsProcessorOutput( diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index f0dea5595..aa1cd6184 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -266,6 +266,20 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase): tree_mask_buf, position_buf = ( self.target_worker.model_runner.attn_backend.get_verify_buffers_to_fill_after_draft() ) + + # build_tree_kernel uses seq_lens_sum only to size the (non-preallocated) + # tree mask; over-size is safe. Skip per-iter .sum().item() D2H via UB. + seq_lens_sum = batch.seq_lens_sum + if seq_lens_sum is None: + if tree_mask_buf is None: + max_context_len = ( + self.target_worker.model_runner.attn_backend.max_context_len + ) + seq_lens_sum = batch.seq_lens.shape[0] * max_context_len + else: + # tree_mask_buf preallocated -> kernel ignores seq_lens_sum. + seq_lens_sum = 0 + ( tree_mask, position, @@ -279,7 +293,7 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase): top_scores_index, draft_tokens, batch.seq_lens, - batch.seq_lens_sum, + seq_lens_sum, self.topk, self.speculative_num_steps, self.speculative_num_draft_tokens,