diff --git a/tensorrt_llm/_torch/pyexecutor/executor_request_queue.py b/tensorrt_llm/_torch/pyexecutor/executor_request_queue.py index c03673f34e7b..76d9f858bc8e 100644 --- a/tensorrt_llm/_torch/pyexecutor/executor_request_queue.py +++ b/tensorrt_llm/_torch/pyexecutor/executor_request_queue.py @@ -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 diff --git a/tensorrt_llm/_torch/pyexecutor/llm_request.py b/tensorrt_llm/_torch/pyexecutor/llm_request.py index fb0670e3756c..128c514cf114 100644 --- a/tensorrt_llm/_torch/pyexecutor/llm_request.py +++ b/tensorrt_llm/_torch/pyexecutor/llm_request.py @@ -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: @@ -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 diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 5f1e8ac147d7..14a7edbf5a1d 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -1,3 +1,4 @@ +import asyncio import dataclasses import datetime import functools @@ -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 @@ -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) @@ -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']) @@ -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): @@ -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 diff --git a/tensorrt_llm/_torch/pyexecutor/sampler.py b/tensorrt_llm/_torch/pyexecutor/sampler.py index e6d19a9df46c..8926d478a1c5 100644 --- a/tensorrt_llm/_torch/pyexecutor/sampler.py +++ b/tensorrt_llm/_torch/pyexecutor/sampler.py @@ -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 @@ -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: diff --git a/tensorrt_llm/_torch/speculative/external_api.py b/tensorrt_llm/_torch/speculative/external_api.py new file mode 100755 index 000000000000..653659c765ac --- /dev/null +++ b/tensorrt_llm/_torch/speculative/external_api.py @@ -0,0 +1,188 @@ +import asyncio +import json +from typing import List + +import aiohttp + +from tensorrt_llm.logger import logger + +from ..pyexecutor.llm_request import * +from ..pyexecutor.scheduler import ScheduledRequests +from .drafter import Drafter + + +class APIDrafter(Drafter): + + def __init__( + self, + spec_config: "ExternalAPIConfig", + ): + super().__init__() + self.max_draft_len = spec_config.max_draft_len + self.endpoint = spec_config.endpoint + assert self.endpoint is not None, "API endpoint is required for external API speculative decoding." + self.template = spec_config.template if spec_config.template is not None else {} + self.response_field = spec_config.response_field if spec_config.response_field is not None else "draft_tokens" + + def single_draft_call(self): + return True + + def get_nested_field_from_response(self, response: dict) -> List[int]: + # Allows for nested fields in the response. + # Example: "choices.0.message.content" + # Returns the value of the nested field: response["choices"][0]["message"]["content"] + keys = self.response_field.split(".") + current = response + + for key in keys: + try: + if key.isdigit(): + key = int(key) + if isinstance(current, list) and 0 <= key < len(current): + current = current[key] + else: + logger.warning( + f"Response field {self.response_field} is invalid for response {response}. Index {key} is invalid." + ) + return [] + else: + if isinstance(current, dict) and key in current: + current = current[key] + else: + logger.warning( + f"Response field {self.response_field} is invalid for response {response}. Index {key} is invalid." + ) + return [] + + except (KeyError, ValueError, IndexError): + logger.warning( + f"Response field path is invalid: {self.response_field}") + return [] + + if not isinstance(current, list): + logger.warning( + f"API response '{self.response_field}' must be a list. Got type: {type(current)}" + ) + return [] + return current + + async def get_draft_tokens( + self, + prefix: list[int], + request_id: int, + end_id: int, + max_sequence_length: int, + ) -> List[int]: + try: + request_data = { + "prefix": prefix, + "request_id": request_id, + "end_id": end_id, + "max_sequence_length": max_sequence_length, + } + if self.template: + request_data.update(self.template) + + async with aiohttp.ClientSession() as session: + async with session.post( + url=self.endpoint, + json=request_data, + headers={"Content-Type": "application/json"}, + timeout=aiohttp.ClientTimeout(total=10), + ) as response: + + # check for unsuccessful response + if response.status != 200: + logger.error( + f"Failed to get draft tokens. API call failed for request {request_id} with status code {response.status}" + ) + return [] + + result = await response.json() + draft_tokens = self.get_nested_field_from_response(result) + if len(draft_tokens) > self.max_draft_len: + draft_tokens = draft_tokens[:self.max_draft_len] + logger.debug( + f"Retrieved draft tokens for request {request_id}: {draft_tokens}" + ) + return draft_tokens + + except json.JSONDecodeError as e: + logger.error( + f"Failed to parse JSON response for request {request_id}: {e}") + return [] + + except Exception as e: + logger.error( + f"Failed to get draft tokens. API call failed for request {request_id} with the following error: {e}" + ) + return [] + + async def async_prepare_draft_tokens( + self, + scheduled_requests: ScheduledRequests, + resource_manager: None, + ) -> None: + # Sort by request_id when py_batch_idx is None as a fallback. + # This happens in the disagg case: for a set of new requests, we draft + # before forward_step, so py_batch_idx is not assigned. + sorted_requests = sorted( + scheduled_requests.generation_requests, + key=lambda r: + (r.py_batch_idx is None, r.py_batch_idx or r.request_id), + ) + + tasks = [] + for request in sorted_requests: + # Add new token to a copy of the generated tokens to find new draft tokens + prefix = list(request.get_tokens()[0]) # Get a copy + if request.py_parallel_spec_dec_params: + if not request.py_parallel_spec_dec_params["pre_verify"]: + prefix += list( + request.py_parallel_spec_dec_params["old_draft_tokens"]) + task = self.get_draft_tokens( + prefix, + request.request_id, + request.py_end_id, + request.py_orig_prompt_len + request.py_max_new_tokens, + ) + tasks.append(task) + try: + all_draft_tokens = await asyncio.wait_for(asyncio.gather( + *tasks, return_exceptions=True), + timeout=10.0) + except asyncio.TimeoutError: + logger.error( + f"Timeout occurred while getting draft tokens for batch of requests" + ) + all_draft_tokens = [[] for _ in tasks] + + draft_tokens_result = {} + for request, draft_tokens in zip(sorted_requests, all_draft_tokens): + if isinstance(draft_tokens, Exception): + logger.error( + f"An exception occurred while getting draft tokens for request {request.request_id}. Set TLLM_LOG_LEVEL for more details." + ) + draft_tokens = [] + elif len(draft_tokens) == 0: + logger.error( + f"Returning empty draft tokens for request {request.request_id}. Set TLLM_LOG_LEVEL for more details." + ) + else: + # Pad length to `self.max_draft_len` + if len(draft_tokens) > 0: + pad_length = self.max_draft_len - len(draft_tokens) + draft_tokens.extend([request.py_end_id] * pad_length) + draft_tokens_result[request.py_request_id] = draft_tokens + return draft_tokens_result + + def prepare_draft_tokens( + self, + scheduled_requests: ScheduledRequests, + resource_manager: None, + ) -> None: + draft_tokens = asyncio.run( + self.async_prepare_draft_tokens(scheduled_requests, + resource_manager)) + for request in scheduled_requests.generation_requests: + request.py_draft_tokens = draft_tokens[request.py_request_id] diff --git a/tensorrt_llm/_torch/speculative/interface.py b/tensorrt_llm/_torch/speculative/interface.py index 1d306b902910..9ad3ee9d9684 100644 --- a/tensorrt_llm/_torch/speculative/interface.py +++ b/tensorrt_llm/_torch/speculative/interface.py @@ -17,6 +17,7 @@ class SpeculativeDecodingMode(IntEnum): NGRAM = auto() DRAFT_TARGET = auto() USER_PROVIDED = auto() + EXTERNAL_API = auto() NONE = auto() AUTO = auto() @@ -44,6 +45,9 @@ def is_ngram(self): def is_user_provided(self): return self == SpeculativeDecodingMode.USER_PROVIDED + def is_external_api(self): + return self == SpeculativeDecodingMode.EXTERNAL_API + def is_none(self): return self == SpeculativeDecodingMode.NONE @@ -82,7 +86,7 @@ def has_spec_decoder(self): def has_spec_drafter(self): return self.is_eagle3() or self.is_draft_target() or self.is_ngram( - ) or self.is_user_provided() + ) or self.is_user_provided() or self.is_external_api() def extend_ctx(self, attention_backend: Type[AttentionBackend]): """ diff --git a/tensorrt_llm/_torch/speculative/utils.py b/tensorrt_llm/_torch/speculative/utils.py index 16fef4862b3f..aaf3973fa551 100644 --- a/tensorrt_llm/_torch/speculative/utils.py +++ b/tensorrt_llm/_torch/speculative/utils.py @@ -7,6 +7,7 @@ from .eagle3 import (Eagle3OneModelSampler, Eagle3OneModelSpecMetadata, Eagle3OneModelWorker, Eagle3ResourceManager, Eagle3SpecMetadata) +from .external_api import APIDrafter from .model_drafter import ModelDrafter from .mtp import (MTPEagleWorker, MTPHiddenStatesManager, MTPSampler, MTPSpecMetadata, MTPWorker) @@ -52,7 +53,8 @@ def get_spec_metadata(spec_config, ) if spec_config.spec_dec_mode.is_draft_target() or \ spec_config.spec_dec_mode.is_ngram() or \ - spec_config.spec_dec_mode.is_user_provided(): + spec_config.spec_dec_mode.is_user_provided() or \ + spec_config.spec_dec_mode.is_external_api(): return SpecMetadata( max_draft_len=spec_config.max_draft_len, spec_dec_mode=spec_config.spec_dec_mode, @@ -101,6 +103,8 @@ def get_spec_resource_manager(model_engine, draft_model_engine=None): return NGramPoolManager(spec_config, max_num_requests) if spec_dec_mode.is_user_provided(): return spec_config.resource_manager + if spec_dec_mode.is_external_api(): + return None return None @@ -129,7 +133,6 @@ def get_spec_drafter(model_engine, if spec_config.spec_dec_mode.is_user_provided(): return spec_config.drafter - max_num_requests = model_engine.batch_size if spec_config.spec_dec_mode.is_draft_target( ) or spec_config.spec_dec_mode.is_eagle3(): @@ -143,7 +146,8 @@ def get_spec_drafter(model_engine, if spec_config.spec_dec_mode.is_ngram(): return NGramDrafter(spec_config, spec_resource_manager) - + if spec_config.spec_dec_mode.is_external_api(): + return APIDrafter(spec_config) return None diff --git a/tensorrt_llm/executor/executor.py b/tensorrt_llm/executor/executor.py index d85f94c34267..87c1a6e62bbb 100644 --- a/tensorrt_llm/executor/executor.py +++ b/tensorrt_llm/executor/executor.py @@ -124,6 +124,7 @@ def generate_async( postproc_params: Optional[PostprocParams] = None, multimodal_params: Optional[MultimodalParams] = None, scheduling_params: Optional[SchedulingParams] = None, + parallel_spec_dec_params: Optional[dict] = None, ) -> GenerationResult: """Generate output for the given prompt token ids in the asynchronous mode. Asynchronous generation accepts single prompt only. @@ -147,7 +148,8 @@ def generate_async( kv_cache_retention_config=kv_cache_retention_config, disaggregated_params=disaggregated_params, multimodal_params=multimodal_params, - scheduling_params=scheduling_params) + scheduling_params=scheduling_params, + parallel_spec_dec_params=parallel_spec_dec_params) result = self.submit(request) # release memory in time if hasattr(request, "multimodal_params"): diff --git a/tensorrt_llm/executor/request.py b/tensorrt_llm/executor/request.py index 00b5deb2eed3..322e8b35c9e8 100644 --- a/tensorrt_llm/executor/request.py +++ b/tensorrt_llm/executor/request.py @@ -97,6 +97,7 @@ def __init__( postproc_params: Optional[PostprocParams] = None, multimodal_params: Optional[MultimodalParams] = None, scheduling_params: Optional[SchedulingParams] = None, + parallel_spec_dec_params: Optional[dict] = None, ): if isinstance(prompt_token_ids, list): self.prompt_token_ids = prompt_token_ids @@ -122,6 +123,7 @@ def __init__( self.id: Optional[int] = None self.disaggregated_params = disaggregated_params self.scheduling_params = scheduling_params + self.parallel_spec_dec_params = parallel_spec_dec_params def set_id(self, id): assert self.id is None, f"Request ID is already set: {self.id}" diff --git a/tensorrt_llm/executor/worker.py b/tensorrt_llm/executor/worker.py index 5266c8dea4c3..4244205d7889 100644 --- a/tensorrt_llm/executor/worker.py +++ b/tensorrt_llm/executor/worker.py @@ -558,6 +558,10 @@ def _deduce_max_tokens(request: GenerationRequest, if self._is_pytorch_backend and request.scheduling_params is not None: executor_request.py_scheduling_params = request.scheduling_params + executor_request.py_parallel_spec_dec_params = None + if self._is_pytorch_backend and request.parallel_spec_dec_params is not None: + executor_request.py_parallel_spec_dec_params = request.parallel_spec_dec_params + if request.query_token_ids is not None: # pytorch star attention workflow # a workaround to avoid public interface update diff --git a/tensorrt_llm/llmapi/__init__.py b/tensorrt_llm/llmapi/__init__.py index 4981b1639170..015c37131109 100644 --- a/tensorrt_llm/llmapi/__init__.py +++ b/tensorrt_llm/llmapi/__init__.py @@ -9,11 +9,11 @@ CapacitySchedulerPolicy, ContextChunkingPolicy, CudaGraphConfig, DraftTargetDecodingConfig, DynamicBatchConfig, EagleDecodingConfig, - ExtendedRuntimePerfKnobConfig, KvCacheConfig, LlmArgs, - LookaheadDecodingConfig, MedusaDecodingConfig, MoeConfig, - MTPDecodingConfig, NGramDecodingConfig, SchedulerConfig, - TorchCompileConfig, TorchLlmArgs, TrtLlmArgs, - UserProvidedDecodingConfig) + ExtendedRuntimePerfKnobConfig, ExternalAPIConfig, + KvCacheConfig, LlmArgs, LookaheadDecodingConfig, + MedusaDecodingConfig, MoeConfig, MTPDecodingConfig, + NGramDecodingConfig, SchedulerConfig, TorchCompileConfig, + TorchLlmArgs, TrtLlmArgs, UserProvidedDecodingConfig) from .llm_utils import (BuildConfig, KvCacheRetentionConfig, QuantAlgo, QuantConfig) from .mm_encoder import MultimodalEncoder @@ -51,6 +51,7 @@ 'CacheTransceiverConfig', 'NGramDecodingConfig', 'UserProvidedDecodingConfig', + 'ExternalAPIConfig', 'TorchCompileConfig', 'DraftTargetDecodingConfig', 'LlmArgs', diff --git a/tensorrt_llm/llmapi/llm.py b/tensorrt_llm/llmapi/llm.py index b95e41f57a2e..3fd4330bbdf6 100644 --- a/tensorrt_llm/llmapi/llm.py +++ b/tensorrt_llm/llmapi/llm.py @@ -206,14 +206,10 @@ def __init__(self, suffix="-llm-workspace", dir=self.args.workspace) else: self._workspace = None - self._hf_model_dir: Optional[Path] = None - self.runtime_context: Optional[_ModelRuntimeContext] = None self.llm_build_stats = LlmBuildStats() - self._build_model() - except Exception: if self.mpi_session is not None: self.mpi_session.shutdown() @@ -325,6 +321,7 @@ def generate_async( disaggregated_params: Optional[DisaggregatedParams] = None, _postproc_params: Optional[PostprocParams] = None, scheduling_params: Optional[SchedulingParams] = None, + parallel_spec_dec_params: Optional[dict] = None, ) -> RequestOutput: """Generate output for the given prompt in the asynchronous mode. Asynchronous generation accepts single prompt only. @@ -444,6 +441,7 @@ def generate_async( postproc_params=_postproc_params, multimodal_params=multimodal_params, scheduling_params=scheduling_params, + parallel_spec_dec_params=parallel_spec_dec_params, ) return RequestOutput._from_generation_result(result, prompt, diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 44abee0994b2..e4a9344b13eb 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -369,6 +369,7 @@ def from_dict(cls, data: dict): "DraftTarget": DraftTargetDecodingConfig, "UserProvided": UserProvidedDecodingConfig, "AUTO": AutoDecodingConfig, + "ExternalAPI": ExternalAPIConfig, } config_class = config_classes.get(decoding_type) @@ -471,6 +472,32 @@ def from_dict(cls, data: dict): decoding_type: ClassVar[str] = "User_Provided" +class ExternalAPIConfig(DecodingBaseConfig): + """ + Configuration for custom drafter speculative decoding. + + Arguments: + endpoint: str + The endpoint of the custom drafter from which to get the draft tokens. + template: dict + The template of the request body. Can be used to include additional fields in the request body. + response_field: str + The field of the response body from which to get the draft tokens. + If not provided, the default field "draft_tokens" is used. + The response body is expected to be a list of draft tokens. + Can be nested, e.g. "draft_tokens.0.tokens" to get draft tokens from response["draft_tokens"][0]["tokens"]. + """ + endpoint: str + template: Optional[dict] = None + response_field: Optional[str] = None + + @classmethod + def from_dict(cls, data: dict): + return cls(**data) + + decoding_type: ClassVar[str] = "External_API" + + class NGramDecodingConfig(DecodingBaseConfig): """ Configuration for NGram drafter speculative decoding. @@ -939,6 +966,7 @@ def supports_backend(self, backend: str) -> bool: NGramDecodingConfig, UserProvidedDecodingConfig, AutoDecodingConfig, + ExternalAPIConfig, ]] @@ -1746,6 +1774,10 @@ def validate_speculative_config(self): elif isinstance(self.speculative_config, AutoDecodingConfig): assert self.backend in ['pytorch', '_autodeploy'] self.build_config.speculative_decoding_mode = SpeculativeDecodingMode.AUTO + + elif isinstance(self.speculative_config, ExternalAPIConfig): + assert self.backend in ['pytorch', '_autodeploy'] + self.build_config.speculative_decoding_mode = SpeculativeDecodingMode.EXTERNAL_API self.build_config.max_draft_len = self.speculative_config.max_draft_len else: diff --git a/tensorrt_llm/llmapi/llm_utils.py b/tensorrt_llm/llmapi/llm_utils.py index b2145ac7935b..81f8bfb88774 100644 --- a/tensorrt_llm/llmapi/llm_utils.py +++ b/tensorrt_llm/llmapi/llm_utils.py @@ -30,8 +30,8 @@ from .build_cache import (BuildCache, BuildCacheConfig, CachedStage, get_build_cache_config_from_env) from .llm_args import (CalibConfig, CudaGraphConfig, DraftTargetDecodingConfig, - EagleDecodingConfig, KvCacheConfig, LlmArgs, - LookaheadDecodingConfig, MedusaDecodingConfig, + EagleDecodingConfig, ExternalAPIConfig, KvCacheConfig, + LlmArgs, LookaheadDecodingConfig, MedusaDecodingConfig, MTPDecodingConfig, NGramDecodingConfig, UserProvidedDecodingConfig, _ModelFormatKind, _ModelWrapper, _ParallelConfig, @@ -874,6 +874,7 @@ class LlmBuildStats: 'NGramDecodingConfig', 'DraftTargetDecodingConfig', 'UserProvidedDecodingConfig', + 'ExternalAPIConfig', 'ContextChunkingPolicy', 'CapacitySchedulerPolicy', 'BuildConfig', diff --git a/tensorrt_llm/models/modeling_utils.py b/tensorrt_llm/models/modeling_utils.py index dcc375320e63..4f43612116b3 100644 --- a/tensorrt_llm/models/modeling_utils.py +++ b/tensorrt_llm/models/modeling_utils.py @@ -99,6 +99,7 @@ class SpeculativeDecodingMode(IntFlag): NGRAM = auto() USER_PROVIDED = auto() AUTO = auto() + EXTERNAL_API = auto() @staticmethod def from_arguments(args: argparse.Namespace): @@ -120,6 +121,8 @@ def from_arguments(args: argparse.Namespace): return SpeculativeDecodingMode.USER_PROVIDED elif args.speculative_decoding_mode == "auto": return SpeculativeDecodingMode.AUTO + elif args.speculative_decoding_mode == "external_api": + return SpeculativeDecodingMode.EXTERNAL_API else: assert False, "Unknown speculative_decoding_mode " + args.speculative_decoding_mode diff --git a/tests/unittest/_torch/speculative/test_external_api.py b/tests/unittest/_torch/speculative/test_external_api.py new file mode 100644 index 000000000000..2710d7769565 --- /dev/null +++ b/tests/unittest/_torch/speculative/test_external_api.py @@ -0,0 +1,262 @@ +import asyncio +import multiprocessing +import os +import sys +import time +import unittest + +import httpx +import pytest +import torch +import uvicorn +from fastapi import FastAPI + +from tensorrt_llm import LLM, SamplingParams +from tensorrt_llm._torch.speculative.external_api import APIDrafter +from tensorrt_llm.llmapi import (CudaGraphConfig, ExternalAPIConfig, + KvCacheConfig) + +sys.path.append(os.path.join(os.path.dirname(__file__), '..')) +from utils.llm_data import llm_models_root +from utils.util import similar + +PORT = 8001 +DEFAULT_DRAFT_TOKENS = [0, 1, 2] + + +def create_server(): + app = FastAPI() + + @app.post("/generate") + async def repeat_draft_tokens(request: dict): + draft_tokens = request["prefix"] + if "extra_token" in request: + draft_tokens.append(request["extra_token"]) + return { + "draft_tokens": DEFAULT_DRAFT_TOKENS, + "draft_tokens_2": draft_tokens, + } + + @app.post("/generate_nested") + async def generate_nested_response(request: dict): + return { + "data": { + "predictions": { + "tokens": DEFAULT_DRAFT_TOKENS + } + }, + "nested_list": [ + { + "tokens": DEFAULT_DRAFT_TOKENS + }, + ] + } + + @app.post("/generate_wrong") + async def generate_wrong_tokens(request: dict): + draft_tokens = "Hello world!" + return { + "draft_tokens": draft_tokens, + } + + @app.post("/generate_none") + async def generate_none_tokens(request: dict): + return { + "draft_tokens": [], + } + + uvicorn.run(app, host="0.0.0.0", port=PORT) + + +@pytest.fixture(scope="module") +def setup_server(): + process = multiprocessing.Process(target=create_server, daemon=True) + process.start() + # wait for server to start + count = 0 + while count < 10: + try: + # check if server is running successfully + response = httpx.post(f"http://localhost:{PORT}/generate", + json={"prefix": [1, 2, 3]}) + if response.status_code == 200: + break + except: + pass + time.sleep(0.5) + count += 0.5 + # wait for tests to run + yield + process.terminate() + + +@pytest.mark.parametrize( + "disable_overlap_scheduler,use_cuda_graph,attn_backend", + [[True, False, "TRTLLM"], [True, True, "TRTLLM"], + [True, False, "FLASHINFER"]]) +def test_llama_user_provided(setup_server, disable_overlap_scheduler: bool, + use_cuda_graph: bool, attn_backend: str): + + max_batch_size = 2 + max_draft_len = 4 + + # endpoint is required to be a non-null value + with pytest.raises(Exception): + spec_config = ExternalAPIConfig(max_draft_len=max_draft_len) + + # test that endpoint can be hit successfully + extra_token = 3 + custom_prefix = [4, 5, 6] + get_draft = lambda drafter: asyncio.run( + drafter.get_draft_tokens(prefix=custom_prefix, + request_id=0, + end_id=0, + max_sequence_length=max_draft_len)) + + # no template, no response field + spec_config = ExternalAPIConfig( + max_draft_len=max_draft_len, + endpoint=f"http://localhost:{PORT}/generate") + spec_drafter = APIDrafter(spec_config) + draft_tokens = get_draft(spec_drafter) + assert draft_tokens == DEFAULT_DRAFT_TOKENS + + # with template, no response field + spec_config = ExternalAPIConfig( + max_draft_len=max_draft_len, + endpoint=f"http://localhost:{PORT}/generate", + template={ + "extra_token": extra_token, + }) + spec_drafter = APIDrafter(spec_config) + draft_tokens = get_draft(spec_drafter) + assert draft_tokens == DEFAULT_DRAFT_TOKENS + + # no template, with response field + spec_config = ExternalAPIConfig( + max_draft_len=max_draft_len, + endpoint=f"http://localhost:{PORT}/generate", + response_field="draft_tokens_2") + spec_drafter = APIDrafter(spec_config) + draft_tokens = get_draft(spec_drafter) + assert draft_tokens == custom_prefix + + # with template, with response field + spec_config = ExternalAPIConfig( + max_draft_len=max_draft_len, + endpoint=f"http://localhost:{PORT}/generate", + template={ + "extra_token": extra_token, + }, + response_field="draft_tokens_2") + spec_drafter = APIDrafter(spec_config) + draft_tokens = get_draft(spec_drafter) + assert draft_tokens == custom_prefix + [extra_token] + + # test correct drafting length (max_draft_len = 4) + draft_tokens = asyncio.run( + spec_drafter.get_draft_tokens(prefix=[0, 1, 2, 3, 4, 5, 6], + request_id=0, + end_id=0, + max_sequence_length=max_draft_len)) + assert draft_tokens == [0, 1, 2, 3] + + # test nested response field + spec_config = ExternalAPIConfig( + max_draft_len=max_draft_len, + endpoint=f"http://localhost:{PORT}/generate_nested", + response_field="data.predictions.tokens") + spec_drafter = APIDrafter(spec_config) + draft_tokens = get_draft(spec_drafter) + assert draft_tokens == DEFAULT_DRAFT_TOKENS + + spec_config = ExternalAPIConfig( + max_draft_len=max_draft_len, + endpoint=f"http://localhost:{PORT}/generate_nested", + response_field="nested_list.0.tokens") + spec_drafter = APIDrafter(spec_config) + draft_tokens = get_draft(spec_drafter) + assert draft_tokens == DEFAULT_DRAFT_TOKENS + + spec_config = ExternalAPIConfig( + max_draft_len=max_draft_len, + endpoint=f"http://localhost:{PORT}/generate_nested", + response_field="data.predictions.wrong") + spec_drafter = APIDrafter(spec_config) + draft_tokens = get_draft(spec_drafter) + assert draft_tokens == [] + + spec_config = ExternalAPIConfig( + max_draft_len=max_draft_len, + endpoint=f"http://localhost:{PORT}/generate_nested", + response_field="nested_list.3.tokens") + spec_drafter = APIDrafter(spec_config) + draft_tokens = get_draft(spec_drafter) + assert draft_tokens == [] + + # test wrong response field type + spec_config = ExternalAPIConfig( + max_draft_len=max_draft_len, + endpoint=f"http://localhost:{PORT}/generate_wrong", + ) + spec_drafter = APIDrafter(spec_config) + draft_tokens = get_draft(spec_drafter) + assert draft_tokens == [] + + # test non-existent endpoint + spec_config = ExternalAPIConfig( + max_draft_len=max_draft_len, + endpoint=f"http://localhost:{PORT}/nonexistent", + ) + spec_drafter = APIDrafter(spec_config) + draft_tokens = get_draft(spec_drafter) + assert draft_tokens == [] + + # spec dec correctness test + # no draft tokens generated, so should be identical to target + total_mem_gb = torch.cuda.get_device_properties(0).total_memory / 1e9 + if total_mem_gb < 20: + pytest.skip("Not enough memory to load target model") + + kv_cache_config = KvCacheConfig(enable_block_reuse=False) + cuda_graph_config = CudaGraphConfig( + batch_sizes=[1]) if use_cuda_graph else None + + llm_common_config = dict( \ + model=llm_models_root() / "llama-3.1-model" /"Meta-Llama-3.1-8B", + backend='pytorch', + attn_backend=attn_backend, + disable_overlap_scheduler=disable_overlap_scheduler, + cuda_graph_config=cuda_graph_config, + max_batch_size=max_batch_size, + kv_cache_config=kv_cache_config, + max_num_tokens=2048, + ) + prompts = [ + "The capital of France is", + "The president of the United States is", + ] + + spec_config = ExternalAPIConfig( + max_draft_len=max_draft_len, + endpoint=f"http://localhost:{PORT}/generate_none", + ) + sampling_params = SamplingParams(max_tokens=32) + + llm_spec = LLM(**llm_common_config, speculative_config=spec_config) + results_spec = llm_spec.generate(prompts, sampling_params) + generated_text_spec = [result.outputs[0].text for result in results_spec] + llm_spec.shutdown() + + llm_ref = LLM(**llm_common_config) + results_ref = llm_ref.generate(prompts, sampling_params) + generated_text_ref = [result.outputs[0].text for result in results_ref] + llm_ref.shutdown() + + for text_spec, text_ref in zip(generated_text_spec, generated_text_ref): + # The spec decode algorithm currently guarantees identical results + assert similar(text_spec, text_ref) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/unittest/api_stability/references_committed/llm.yaml b/tests/unittest/api_stability/references_committed/llm.yaml index a722da549586..426fccd61c0b 100644 --- a/tests/unittest/api_stability/references_committed/llm.yaml +++ b/tests/unittest/api_stability/references_committed/llm.yaml @@ -59,7 +59,7 @@ methods: default: null # Speculative decoding speculative_config: - annotation: Union[tensorrt_llm.llmapi.llm_args.DraftTargetDecodingConfig, tensorrt_llm.llmapi.llm_args.EagleDecodingConfig,tensorrt_llm.llmapi.llm_args.LookaheadDecodingConfig, tensorrt_llm.llmapi.llm_args.MedusaDecodingConfig, tensorrt_llm.llmapi.llm_args.MTPDecodingConfig, tensorrt_llm.llmapi.llm_args.NGramDecodingConfig, tensorrt_llm.llmapi.llm_args.UserProvidedDecodingConfig, tensorrt_llm.llmapi.llm_args.AutoDecodingConfig, NoneType] + annotation: Union[tensorrt_llm.llmapi.llm_args.DraftTargetDecodingConfig, tensorrt_llm.llmapi.llm_args.EagleDecodingConfig,tensorrt_llm.llmapi.llm_args.LookaheadDecodingConfig, tensorrt_llm.llmapi.llm_args.MedusaDecodingConfig, tensorrt_llm.llmapi.llm_args.MTPDecodingConfig, tensorrt_llm.llmapi.llm_args.NGramDecodingConfig, tensorrt_llm.llmapi.llm_args.UserProvidedDecodingConfig, tensorrt_llm.llmapi.llm_args.ExternalAPIConfig, tensorrt_llm.llmapi.llm_args.AutoDecodingConfig, NoneType] default: null # generation constraints max_batch_size: