Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 3 additions & 1 deletion tensorrt_llm/_torch/pyexecutor/executor_request_queue.py
Original file line number Diff line number Diff line change
Expand Up @@ -433,10 +433,12 @@ def _handle_request_broadcasting(self,
new_requests, "py_multimodal_data")
py_scheduling_params = self._collect_py_objects_from_requests(
new_requests, "py_scheduling_params")
py_parallel_spec_dec_params = self._collect_py_objects_from_requests(
new_requests, "py_parallel_spec_dec_params")
py_request_objects = tuple(
filter(None, [
py_logits_post_processors, py_multimodal_data,
py_scheduling_params
py_scheduling_params, py_parallel_spec_dec_params
]))
else:
py_request_objects = None
Expand Down
8 changes: 6 additions & 2 deletions tensorrt_llm/_torch/pyexecutor/llm_request.py
Original file line number Diff line number Diff line change
Expand Up @@ -307,6 +307,8 @@ def __init__(
self.py_lora_path: str | None = kwargs.pop("py_lora_path", None)
# Multimodal data
self.py_multimodal_data = kwargs.pop("py_multimodal_data", None)
self.py_parallel_spec_dec_params = kwargs.pop(
"py_parallel_spec_dec_params", None)
if llm_request is not None:
super().__init__(llm_request)
else:
Expand Down Expand Up @@ -547,11 +549,13 @@ def executor_request_to_llm_request(
llm_request_type=llm_request_type,
context_phase_params=executor_request.context_phase_params,
py_multimodal_data=getattr(executor_request, "py_multimodal_data",
None))
None),
py_parallel_spec_dec_params=getattr(executor_request,
"py_parallel_spec_dec_params",
None))
if child_req_ids:
for child_id in child_req_ids:
llm_request.create_child_request(child_id)

return llm_request


Expand Down
125 changes: 110 additions & 15 deletions tensorrt_llm/_torch/pyexecutor/py_executor.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import asyncio
import dataclasses
import datetime
import functools
Expand Down Expand Up @@ -42,7 +43,7 @@
from .handle_logits import HandleLogits
from .kv_cache_transceiver import KvCacheTransceiver
from .llm_request import (ExecutorRequest, LlmRequest, LlmRequestState,
LlmResponse, get_draft_token_length)
LlmResponse)
from .model_engine import ModelEngine
from .sampler import Sampler, SampleState, SampleStateTensors
from .scheduler import RequestScheduler, ScheduledRequests
Expand Down Expand Up @@ -952,7 +953,6 @@ def _executor_loop(self):
self._handle_first_token_response(scheduled_batch)

self.resource_manager.prepare_resources(scheduled_batch)

if self.kv_cache_transceiver and self.guided_decoder:
self.guided_decoder.init_disagg_gen_requests(
scheduled_batch)
Expand All @@ -963,19 +963,18 @@ def _executor_loop(self):
if self.guided_decoder is not None:
self.guided_decoder.rollback_rejected_tokens(
scheduled_batch)
self.drafter.prepare_draft_tokens(
scheduled_batch, self.resource_manager)

# Pad draft tokens to the max draft length. This is for CUDA
# graph compatibility.
for req in scheduled_batch.generation_requests:
max_draft_tokens = self.max_draft_len
num_draft_tokens = get_draft_token_length(req)
req.py_draft_tokens.extend(
0 for _ in range(max_draft_tokens -
num_draft_tokens))

batch_outputs = self._forward_step(scheduled_batch)

run_parallel_spec_dec = self._get_parallel_spec_dec_mode(
scheduled_batch.generation_requests)
if run_parallel_spec_dec:
batch_outputs = asyncio.run(
self._parallel_spec_dec_forward_step(
scheduled_batch))
else:
if self.drafter is not None and self.use_spec_decode:
self._prepare_draft_tokens(scheduled_batch)
batch_outputs = self._forward_step(scheduled_batch)

self._execute_guided_decoder(scheduled_batch,
batch_outputs['logits'])

Expand Down Expand Up @@ -1432,6 +1431,49 @@ def _send_disagg_ctx_cache(self, scheduled_ctx_requests):

return ctx_transmission_reqs

def _prepare_draft_tokens(self, scheduled_batch):
# External drafters can be configured to only make a single draft call when running with TP > 1
single_draft_call = getattr(self.drafter, 'single_draft_call',
lambda: False)()
if single_draft_call:
# Only rank 0 prepares draft tokens
if self.dist.rank == 0:
self.drafter.prepare_draft_tokens(scheduled_batch,
self.resource_manager)
# Broadcast to other ranks
if self.dist.tp_size > 1:
if self.dist.rank == 0:
draft_data = {}
for req in scheduled_batch.generation_requests:
draft_data[req.py_request_id] = req.py_draft_tokens
self.dist.tp_broadcast(draft_data, root=0)
else:
draft_data = self.dist.tp_broadcast(None, root=0)
for req in scheduled_batch.generation_requests:
req.py_draft_tokens = draft_data[req.py_request_id]
else:
self.drafter.prepare_draft_tokens(scheduled_batch,
self.resource_manager)

async def _parallel_prepare_draft_tokens(self, scheduled_batch):
single_draft_call = getattr(self.drafter, 'single_draft_call',
lambda: False)()
if not single_draft_call or not hasattr(self.drafter,
'async_prepare_draft_tokens'):
raise ValueError(
"Parallel speculation must use an external drafter with async_prepare_draft_tokens implemented"
)
if self.dist.rank == 0:
draft_tokens = await self.drafter.async_prepare_draft_tokens(
scheduled_batch, self.resource_manager)
# Broadcast to other ranks
if self.dist.tp_size > 1:
if self.dist.rank == 0:
self.dist.tp_broadcast(draft_tokens, root=0)
else:
draft_tokens = self.dist.tp_broadcast(None, root=0)
return draft_tokens

def _forward_step(self,
scheduled_requests,
new_tensors_device: Optional[SampleStateTensors] = None):
Expand Down Expand Up @@ -1789,3 +1831,56 @@ def _remove_inflight_ids(self, scheduled_requests):
"""Remove reqids of current requests from self.inflight_req_ids."""
for req in scheduled_requests.all_requests():
self.inflight_req_ids.erase(req.request_id)

def _get_parallel_spec_dec_mode(self, scheduled_requests):
run_parallel = False
if len(scheduled_requests):
run_parallel = scheduled_requests[
0].py_parallel_spec_dec_params is not None
# check that all requests in the batch have the same parallel spec dec mode
for req in scheduled_requests:
if (req.py_parallel_spec_dec_params is not None) != run_parallel:
raise ValueError(
"Parallel spec dec must be enabled for all or none of the requests in the batch"
)
return run_parallel

async def _parallel_forward_step(self, scheduled_batch):
loop = asyncio.get_event_loop()
return await loop.run_in_executor(None, self._forward_step,
scheduled_batch)

async def _parallel_spec_dec_forward_step(self, scheduled_batch):
# Append previous draft tokens to request prefixes
for req in scheduled_batch.generation_requests:
# if running with post-verify, load previous draft tokens
if not req.py_parallel_spec_dec_params["pre_verify"]:
if len(req.py_parallel_spec_dec_params["old_draft_tokens"]
) == 0:
logger.error(
f"Cannot run post-verify as previous draft tokens for request {req.py_request_id} are empty. Switching to pre-verify."
)
req.py_parallel_spec_dec_params["pre_verify"] = True
else:
req.py_draft_tokens = list(
req.py_parallel_spec_dec_params["old_draft_tokens"])
pad_length = self.max_draft_len - len(req.py_draft_tokens)
req.py_draft_tokens.extend([req.py_end_id] * pad_length)

draft_task = self._parallel_prepare_draft_tokens(scheduled_batch)
target_task = self._parallel_forward_step(scheduled_batch)

draft_tokens, batch_outputs = await asyncio.gather(
draft_task, target_task)
for req in scheduled_batch.generation_requests:
req.py_draft_tokens = draft_tokens[req.py_request_id]

# Restore old draft for verification
for req in scheduled_batch.generation_requests:
old_draft_tokens = req.py_parallel_spec_dec_params[
"old_draft_tokens"]
req.py_parallel_spec_dec_params[
"old_draft_tokens"] = req.py_draft_tokens
req.py_draft_tokens = old_draft_tokens

return batch_outputs
26 changes: 26 additions & 0 deletions tensorrt_llm/_torch/pyexecutor/sampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -560,6 +560,8 @@ def update_requests(self, state: SampleState) -> None:
continue
processed = 1
num_accepted = self.process_draft_tokens(req, new_tokens)
if req.py_parallel_spec_dec_params is not None:
self._pre_verify(req, new_tokens)
if get_draft_token_length(req) > 0:
req.py_num_accepted_draft_tokens = num_accepted
req.py_rewind_len = req.py_draft_pages_allocated - num_accepted
Expand Down Expand Up @@ -750,6 +752,30 @@ def _process_requests(self,
non_blocking=True)
offset += steps

def _pre_verify(self, req: LlmRequest, new_tokens: torch.Tensor):
try:
next_draft_token = req.py_parallel_spec_dec_params[
"old_draft_tokens"][0]
pre_verified_token = req.get_tokens(0)[-1]
if next_draft_token == pre_verified_token:
old_draft_tokens = req.py_parallel_spec_dec_params[
"old_draft_tokens"][1:]
# remove draft sequence padding
while old_draft_tokens and old_draft_tokens[-1] == req.py_end_id:
old_draft_tokens.pop()
# switch to post-verify
req.py_parallel_spec_dec_params[
"old_draft_tokens"] = old_draft_tokens
req.py_parallel_spec_dec_params["pre_verify"] = False
else:
# switch to pre-verify
req.py_parallel_spec_dec_params["old_draft_tokens"] = []
req.py_parallel_spec_dec_params["pre_verify"] = True
except Exception:
# fall back to pre-verify
req.py_parallel_spec_dec_params["old_draft_tokens"] = []
req.py_parallel_spec_dec_params["pre_verify"] = True


class Algorithms:

Expand Down
Loading