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():
|
elif forward_mode.is_target_verify():
|
||||||
# Target Verify
|
# Target Verify
|
||||||
# Here we only support topk = 1 for now.
|
# Here we only support topk = 1 for now.
|
||||||
|
tokens_per_req = num_tokens // bs
|
||||||
metadata.cache_seqlens_int32 = self.target_verify_metadata["cache_seqlens"][
|
metadata.cache_seqlens_int32 = self.target_verify_metadata["cache_seqlens"][
|
||||||
:bs
|
:bs
|
||||||
]
|
]
|
||||||
metadata.cache_seqlens_int32.copy_(
|
metadata.cache_seqlens_int32.copy_(seq_lens + tokens_per_req)
|
||||||
(seq_lens + self.speculative_num_draft_tokens)
|
|
||||||
)
|
|
||||||
|
|
||||||
metadata.cu_seqlens_q = torch.arange(
|
metadata.cu_seqlens_q = torch.arange(
|
||||||
0,
|
0,
|
||||||
bs * self.speculative_num_draft_tokens + 1,
|
bs * tokens_per_req + 1,
|
||||||
self.speculative_num_draft_tokens,
|
tokens_per_req,
|
||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
device=device,
|
device=device,
|
||||||
)
|
)
|
||||||
@@ -396,10 +395,8 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
|||||||
: (bs + 1)
|
: (bs + 1)
|
||||||
]
|
]
|
||||||
|
|
||||||
metadata.max_seq_len_q = self.speculative_num_draft_tokens
|
metadata.max_seq_len_q = tokens_per_req
|
||||||
metadata.max_seq_len_k = (
|
metadata.max_seq_len_k = seq_lens.max().item() + tokens_per_req
|
||||||
seq_lens.max().item() + self.speculative_num_draft_tokens
|
|
||||||
)
|
|
||||||
|
|
||||||
metadata.page_table = self.target_verify_metadata["page_table"][:bs, :]
|
metadata.page_table = self.target_verify_metadata["page_table"][:bs, :]
|
||||||
self._bind_swa_page_table(
|
self._bind_swa_page_table(
|
||||||
@@ -496,13 +493,9 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
|||||||
elif forward_mode.is_target_verify():
|
elif forward_mode.is_target_verify():
|
||||||
# Here we only support topk = 1 for now.
|
# Here we only support topk = 1 for now.
|
||||||
metadata = self.target_verify_metadata[bs]
|
metadata = self.target_verify_metadata[bs]
|
||||||
metadata.cache_seqlens_int32.copy_(
|
metadata.cache_seqlens_int32.copy_(seq_lens + metadata.max_seq_len_q)
|
||||||
(seq_lens + self.speculative_num_draft_tokens)
|
|
||||||
)
|
|
||||||
|
|
||||||
metadata.max_seq_len_k = (
|
metadata.max_seq_len_k = seq_lens_cpu.max().item() + metadata.max_seq_len_q
|
||||||
seq_lens_cpu.max().item() + self.speculative_num_draft_tokens
|
|
||||||
)
|
|
||||||
max_len = seq_lens_cpu.max().item()
|
max_len = seq_lens_cpu.max().item()
|
||||||
metadata.cu_seqlens_k[1:].copy_(
|
metadata.cu_seqlens_k[1:].copy_(
|
||||||
torch.cumsum(metadata.cache_seqlens_int32, dim=0, dtype=torch.int32)
|
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)
|
metadata.page_table[:, :max_seq_pages].copy_(page_indices // self.page_size)
|
||||||
self._copy_swa_page_table(metadata, page_indices, max_seq_pages)
|
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():
|
elif forward_mode.is_draft_extend():
|
||||||
metadata = self.draft_extend_metadata[bs]
|
metadata = self.draft_extend_metadata[bs]
|
||||||
metadata.cache_seqlens_int32.copy_(seq_lens)
|
metadata.cache_seqlens_int32.copy_(seq_lens)
|
||||||
@@ -626,18 +618,18 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
|||||||
]
|
]
|
||||||
elif forward_batch.forward_mode.is_target_verify():
|
elif forward_batch.forward_mode.is_target_verify():
|
||||||
# Only support topk = 1 for now.
|
# Only support topk = 1 for now.
|
||||||
metadata.cache_seqlens_int32 = (
|
tokens_per_req = forward_batch.input_ids.shape[0] // batch_size
|
||||||
forward_batch.seq_lens + self.speculative_num_draft_tokens
|
metadata.cache_seqlens_int32 = (forward_batch.seq_lens + tokens_per_req).to(
|
||||||
).to(torch.int32)
|
torch.int32
|
||||||
metadata.max_seq_len_q = self.speculative_num_draft_tokens
|
)
|
||||||
|
metadata.max_seq_len_q = tokens_per_req
|
||||||
metadata.max_seq_len_k = (
|
metadata.max_seq_len_k = (
|
||||||
forward_batch.seq_lens_cpu.max().item()
|
forward_batch.seq_lens_cpu.max().item() + tokens_per_req
|
||||||
+ self.speculative_num_draft_tokens
|
|
||||||
)
|
)
|
||||||
metadata.cu_seqlens_q = torch.arange(
|
metadata.cu_seqlens_q = torch.arange(
|
||||||
0,
|
0,
|
||||||
batch_size * self.speculative_num_draft_tokens + 1,
|
batch_size * tokens_per_req + 1,
|
||||||
self.speculative_num_draft_tokens,
|
tokens_per_req,
|
||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
device=device,
|
device=device,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -617,7 +617,11 @@ class CudaGraphRunner:
|
|||||||
):
|
):
|
||||||
raise RuntimeError("This should not happen")
|
raise RuntimeError("This should not happen")
|
||||||
self.capture_forward_mode = ForwardMode.TARGET_VERIFY
|
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:
|
elif self.is_dllm:
|
||||||
self.capture_forward_mode = ForwardMode.DLLM_EXTEND
|
self.capture_forward_mode = ForwardMode.DLLM_EXTEND
|
||||||
self.num_tokens_per_bs = self.dllm_config.block_size
|
self.num_tokens_per_bs = self.dllm_config.block_size
|
||||||
|
|||||||
@@ -2423,10 +2423,14 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
num_tokens_per_bs = 1
|
num_tokens_per_bs = 1
|
||||||
if self.spec_algorithm.is_speculative():
|
if self.spec_algorithm.is_speculative():
|
||||||
if self.is_draft_worker:
|
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")
|
raise RuntimeError("This should not happen")
|
||||||
capture_forward_mode = ForwardMode.TARGET_VERIFY
|
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:
|
if self.server_args.enable_return_hidden_states:
|
||||||
capture_hidden_mode = CaptureHiddenMode.FULL
|
capture_hidden_mode = CaptureHiddenMode.FULL
|
||||||
|
|||||||
@@ -135,6 +135,16 @@ class SpeculativeAlgorithm(Enum):
|
|||||||
def supports_spec_v2(self) -> bool:
|
def supports_spec_v2(self) -> bool:
|
||||||
return (self.is_eagle() and not self.is_frozen_kv_mtp()) or self.is_standalone()
|
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(
|
def create_worker(
|
||||||
self, server_args: ServerArgs
|
self, server_args: ServerArgs
|
||||||
) -> Optional[Union[Type[BaseSpecWorker], Type[TpModelWorker], Type[NGRAMWorker]]]:
|
) -> Optional[Union[Type[BaseSpecWorker], Type[TpModelWorker], Type[NGRAMWorker]]]:
|
||||||
|
|||||||
@@ -64,6 +64,9 @@ class CustomSpecAlgo:
|
|||||||
def is_ngram(self) -> bool:
|
def is_ngram(self) -> bool:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
def supports_target_verify_for_draft(self) -> bool:
|
||||||
|
return False
|
||||||
|
|
||||||
def supports_spec_v2(self) -> bool:
|
def supports_spec_v2(self) -> bool:
|
||||||
return self.supports_overlap
|
return self.supports_overlap
|
||||||
|
|
||||||
@@ -74,6 +77,16 @@ class CustomSpecAlgo:
|
|||||||
)
|
)
|
||||||
return self.factory(server_args)
|
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] = {}
|
_REGISTRY: Dict[str, CustomSpecAlgo] = {}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user