Add support for generic num_tokens_per_bs in TARGET_VERIFY (#25681)

This commit is contained in:
jasonjk-park
2026-05-20 11:34:45 -07:00
committed by GitHub
parent 61ac6792e6
commit 1f209b4433
5 changed files with 50 additions and 27 deletions
@@ -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,
)
@@ -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
@@ -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
@@ -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]]]:
@@ -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] = {}