diff --git a/python/sglang/srt/layers/attention/trtllm_mha_backend.py b/python/sglang/srt/layers/attention/trtllm_mha_backend.py index 0f99206c2..38721b7f4 100644 --- a/python/sglang/srt/layers/attention/trtllm_mha_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mha_backend.py @@ -377,17 +377,16 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): elif forward_mode.is_target_verify(): # Target Verify # Here we only support topk = 1 for now. + tokens_per_req = num_tokens // bs metadata.cache_seqlens_int32 = self.target_verify_metadata["cache_seqlens"][ :bs ] - metadata.cache_seqlens_int32.copy_( - (seq_lens + self.speculative_num_draft_tokens) - ) + metadata.cache_seqlens_int32.copy_(seq_lens + tokens_per_req) metadata.cu_seqlens_q = torch.arange( 0, - bs * self.speculative_num_draft_tokens + 1, - self.speculative_num_draft_tokens, + bs * tokens_per_req + 1, + tokens_per_req, dtype=torch.int32, device=device, ) @@ -396,10 +395,8 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): : (bs + 1) ] - metadata.max_seq_len_q = self.speculative_num_draft_tokens - metadata.max_seq_len_k = ( - seq_lens.max().item() + self.speculative_num_draft_tokens - ) + metadata.max_seq_len_q = tokens_per_req + metadata.max_seq_len_k = seq_lens.max().item() + tokens_per_req metadata.page_table = self.target_verify_metadata["page_table"][:bs, :] self._bind_swa_page_table( @@ -496,13 +493,9 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): 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 + self.speculative_num_draft_tokens) - ) + metadata.cache_seqlens_int32.copy_(seq_lens + metadata.max_seq_len_q) - metadata.max_seq_len_k = ( - seq_lens_cpu.max().item() + self.speculative_num_draft_tokens - ) + metadata.max_seq_len_k = seq_lens_cpu.max().item() + metadata.max_seq_len_q max_len = seq_lens_cpu.max().item() metadata.cu_seqlens_k[1:].copy_( torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32) @@ -516,7 +509,6 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): ] metadata.page_table[:, :max_seq_pages].copy_(page_indices // self.page_size) self._copy_swa_page_table(metadata, page_indices, max_seq_pages) - metadata.max_seq_len_q = self.speculative_num_draft_tokens elif forward_mode.is_draft_extend(): metadata = self.draft_extend_metadata[bs] metadata.cache_seqlens_int32.copy_(seq_lens) @@ -626,18 +618,18 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend): ] elif forward_batch.forward_mode.is_target_verify(): # Only support topk = 1 for now. - metadata.cache_seqlens_int32 = ( - forward_batch.seq_lens + self.speculative_num_draft_tokens - ).to(torch.int32) - metadata.max_seq_len_q = self.speculative_num_draft_tokens + tokens_per_req = forward_batch.input_ids.shape[0] // batch_size + metadata.cache_seqlens_int32 = (forward_batch.seq_lens + tokens_per_req).to( + torch.int32 + ) + metadata.max_seq_len_q = tokens_per_req metadata.max_seq_len_k = ( - forward_batch.seq_lens_cpu.max().item() - + self.speculative_num_draft_tokens + forward_batch.seq_lens_cpu.max().item() + tokens_per_req ) metadata.cu_seqlens_q = torch.arange( 0, - batch_size * self.speculative_num_draft_tokens + 1, - self.speculative_num_draft_tokens, + batch_size * tokens_per_req + 1, + tokens_per_req, dtype=torch.int32, device=device, ) diff --git a/python/sglang/srt/model_executor/cuda_graph_runner.py b/python/sglang/srt/model_executor/cuda_graph_runner.py index 918850de8..ad13066ec 100644 --- a/python/sglang/srt/model_executor/cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/cuda_graph_runner.py @@ -617,7 +617,11 @@ class CudaGraphRunner: ): raise RuntimeError("This should not happen") self.capture_forward_mode = ForwardMode.TARGET_VERIFY - self.num_tokens_per_bs = self.speculative_num_draft_tokens + self.num_tokens_per_bs = ( + model_runner.spec_algorithm.get_num_tokens_per_bs_for_target_verify( + self.speculative_num_draft_tokens, model_runner.is_draft_worker + ) + ) elif self.is_dllm: self.capture_forward_mode = ForwardMode.DLLM_EXTEND self.num_tokens_per_bs = self.dllm_config.block_size diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 9a04a1976..784ddd94f 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -2423,10 +2423,14 @@ class ModelRunner(ModelRunnerKVCacheMixin): num_tokens_per_bs = 1 if self.spec_algorithm.is_speculative(): if self.is_draft_worker: - if not self.spec_algorithm.is_dflash(): + if not self.spec_algorithm.supports_target_verify_for_draft(): raise RuntimeError("This should not happen") capture_forward_mode = ForwardMode.TARGET_VERIFY - num_tokens_per_bs = self.server_args.speculative_num_draft_tokens + num_tokens_per_bs = ( + self.spec_algorithm.get_num_tokens_per_bs_for_target_verify( + self.server_args.speculative_num_draft_tokens, self.is_draft_worker + ) + ) if self.server_args.enable_return_hidden_states: capture_hidden_mode = CaptureHiddenMode.FULL diff --git a/python/sglang/srt/speculative/spec_info.py b/python/sglang/srt/speculative/spec_info.py index 6715fc122..31814e416 100644 --- a/python/sglang/srt/speculative/spec_info.py +++ b/python/sglang/srt/speculative/spec_info.py @@ -135,6 +135,16 @@ class SpeculativeAlgorithm(Enum): def supports_spec_v2(self) -> bool: return (self.is_eagle() and not self.is_frozen_kv_mtp()) or self.is_standalone() + def get_num_tokens_per_bs_for_target_verify( + self, num_draft_tokens: int, is_draft_worker: bool + ) -> int: + # FIXME: Remove this after the forward mode refactor. Target verify is + # essentially a fixed sequence length prefill/extend with full cuda + # graph support. We can use it for target verify, or we can use it for + # other cases which is not target verify but fixed length prefill. + # Here, we expose this interface to allow the other use cases. + return num_draft_tokens + def create_worker( self, server_args: ServerArgs ) -> Optional[Union[Type[BaseSpecWorker], Type[TpModelWorker], Type[NGRAMWorker]]]: diff --git a/python/sglang/srt/speculative/spec_registry.py b/python/sglang/srt/speculative/spec_registry.py index c7e062596..f15d82acb 100644 --- a/python/sglang/srt/speculative/spec_registry.py +++ b/python/sglang/srt/speculative/spec_registry.py @@ -64,6 +64,9 @@ class CustomSpecAlgo: def is_ngram(self) -> bool: return False + def supports_target_verify_for_draft(self) -> bool: + return False + def supports_spec_v2(self) -> bool: return self.supports_overlap @@ -74,6 +77,16 @@ class CustomSpecAlgo: ) return self.factory(server_args) + def get_num_tokens_per_bs_for_target_verify( + self, num_draft_tokens: int, is_draft_worker: bool + ) -> int: + # FIXME: Remove this after the forward mode refactor. Target verify is + # essentially a fixed sequence length prefill/extend with full cuda + # graph support. We can use it for target verify, or we can use it for + # other cases which is not target verify but fixed length prefill. + # Here, we expose this interface to allow the other use cases. + return num_draft_tokens + _REGISTRY: Dict[str, CustomSpecAlgo] = {}