Add support for generic num_tokens_per_bs in TARGET_VERIFY (#25681)
This commit is contained in:
@@ -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] = {}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user