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
14 changes: 12 additions & 2 deletions tensorrt_llm/commands/serve.py
Original file line number Diff line number Diff line change
Expand Up @@ -84,6 +84,7 @@ def get_llm_args(model: str,
num_postprocess_workers: int = 0,
trust_remote_code: bool = False,
reasoning_parser: Optional[str] = None,
served_model_name: Optional[str] = None,
**llm_args_extra_dict: Any):

if gpus_per_node is None:
Expand Down Expand Up @@ -125,6 +126,7 @@ def get_llm_args(model: str,
"num_postprocess_workers": num_postprocess_workers,
"postprocess_tokenizer_dir": tokenizer or model,
"reasoning_parser": reasoning_parser,
"served_model_name": served_model_name,
}

return llm_args, llm_args_extra_dict
Expand All @@ -139,13 +141,16 @@ def launch_server(host: str,
backend = llm_args["backend"]
model = llm_args["model"]

served_model_name = llm_args.pop("served_model_name", None)

if backend == 'pytorch':
llm = PyTorchLLM(**llm_args)
else:
llm = LLM(**llm_args)

server = OpenAIServer(llm=llm,
model=model,
served_model_name=served_model_name,
server_role=server_role,
metadata_server_cfg=metadata_server_cfg)

Expand Down Expand Up @@ -216,6 +221,10 @@ def launch_server(host: str,
default=0.9,
help="Free GPU memory fraction reserved for KV Cache, "
"after allocating model weights and buffers.")
@click.option("--served-model-name",
type=str,
default=None,
help="Name of the model to be served. If not provided, the model name will be used.")
@click.option(
"--num_postprocess_workers",
type=int,
Expand Down Expand Up @@ -249,7 +258,7 @@ def launch_server(host: str,
default=None,
help="Server role. Specify this value only if running in disaggregated mode."
)
def serve(model: str, tokenizer: Optional[str], host: str, port: int,
def serve(model: str, tokenizer: Optional[str], host: str, port: int, served_model_name: Optional[str],
log_level: str, backend: str, max_beam_width: int,
max_batch_size: int, max_num_tokens: int, max_seq_len: int,
tp_size: int, pp_size: int, ep_size: Optional[int],
Expand Down Expand Up @@ -281,7 +290,8 @@ def serve(model: str, tokenizer: Optional[str], host: str, port: int,
free_gpu_memory_fraction=kv_cache_free_gpu_memory_fraction,
num_postprocess_workers=num_postprocess_workers,
trust_remote_code=trust_remote_code,
reasoning_parser=reasoning_parser)
reasoning_parser=reasoning_parser,
served_model_name=served_model_name)

llm_args_extra_dict = {}
if extra_llm_api_options is not None:
Expand Down
6 changes: 3 additions & 3 deletions tensorrt_llm/executor/executor.py
Original file line number Diff line number Diff line change
Expand Up @@ -250,8 +250,8 @@ def _handle_background_error(self, error: Optional[Exception | str] = None):
print_colored(
f"Got background error: {repr(error)}, will shutdown the LLM instance\n",
"red")
self.shutdown()
raise error
# self.shutdown()
raise RequestError(str(error))
elif isinstance(error, str):
if enable_llm_debug():
print_colored(f"Got per-request error: {repr(error)}\n",
Expand All @@ -265,7 +265,7 @@ def _handle_background_error(self, error: Optional[Exception | str] = None):
if not self._error_queue.empty():
e = self._error_queue.get()
self._error_queue.task_done()
self.shutdown()
# self.shutdown()
# We can catch some exceptions here.
raise e

Expand Down
7 changes: 3 additions & 4 deletions tensorrt_llm/executor/worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -902,14 +902,13 @@ def handle_for_ipc_batched(self, responses: List[tllm.Response]) -> None:
rsp_batch = [] if not self.enable_postprocprocess_parallel else None

for response in responses:

if self.worker._has_background_error():
response = self.worker._create_error_response(response)
elif response.has_error():
# Convert to ErrorResponse, because tllm.Response cannot be
# serialized when it has error.
elif isinstance(response, tllm.Response) and response.has_error():
response = ErrorResponse(response.client_id, response.error_msg,
response.request_id)
elif isinstance(response, ErrorResponse):
pass
else:
logprobs_result = _get_logprobs(self.worker, response,
self.worker._is_pytorch_backend)
Expand Down
16 changes: 0 additions & 16 deletions tensorrt_llm/serve/openai_protocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -263,14 +263,6 @@ def check_logprobs(cls, data):
raise ValueError("logprobs is not supported")
return data

@model_validator(mode="before")
@classmethod
def validate_stream_options(cls, data):
if data.get("stream_options") and not data.get("stream"):
raise ValueError(
"Stream options can only be defined when stream is true.")
return data

@model_validator(mode="before")
@classmethod
def check_suffix(cls, data):
Expand Down Expand Up @@ -547,14 +539,6 @@ def to_sampling_params(self) -> SamplingParams:
)
return sampling_params

@model_validator(mode='before')
@classmethod
def validate_stream_options(cls, values):
if (values.get('stream_options') is not None
and not values.get('stream')):
raise ValueError("stream_options can only be set if stream is true")
return values

@model_validator(mode="before")
@classmethod
def check_tool_choice(cls, data):
Expand Down
55 changes: 50 additions & 5 deletions tensorrt_llm/serve/openai_server.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
#!/usr/bin/env python
import asyncio
import os
import signal
import traceback
from contextlib import asynccontextmanager
Expand Down Expand Up @@ -39,19 +40,21 @@
ChatPostprocArgs, CompletionPostprocArgs, chat_response_post_processor,
chat_stream_post_processor, completion_response_post_processor,
completion_stream_post_processor)
from tensorrt_llm.executor.request import LoRARequest
from tensorrt_llm.version import __version__ as VERSION

from .._utils import nvtx_mark

# yapf: enale
TIMEOUT_KEEP_ALIVE = 5 # seconds.

LORA_DIR = os.getenv("LORA_DIR", "/workspace")

class OpenAIServer:

def __init__(self,
llm: LLM,
model: str,
served_model_name: Optional[str],
server_role: Optional[ServerRole],
metadata_server_cfg: MetadataServerConfig):
self.llm = llm
Expand All @@ -73,10 +76,20 @@ def __init__(self,
self.model_config = None

model_dir = Path(model)
if model_dir.exists() and model_dir.is_dir():
self.model = model_dir.name

if served_model_name is not None:
self.model = served_model_name
else:
self.model = model
if model_dir.exists() and model_dir.is_dir():
self.model = model_dir.name
else:
self.model = model

self.lora_names = set([os.path.basename(d) for d in os.listdir(LORA_DIR) if os.path.isdir(os.path.join(LORA_DIR, d))])
if self.model in self.lora_names:
self.lora_names.remove(self.model)

self.lora_mapper = {k: i for i, k in enumerate(self.lora_names)}

@asynccontextmanager
async def lifespan(app: FastAPI):
Expand Down Expand Up @@ -213,7 +226,21 @@ async def get_kv_cache_events(self) -> JSONResponse:
return JSONResponse(content=events)

async def openai_chat(self, request: ChatCompletionRequest, raw_request: Request) -> Response:


request_for_lora = request.model != self.model and request.model in self.lora_names

if request_for_lora and request.lora_request is None:
lora_dir = os.path.join(LORA_DIR, request.model)
if not os.path.exists(lora_dir):
logger.error(f"LORA_DIR {LORA_DIR} does not exist")
return self.create_error_response("Model not found")

request.lora_request = LoRARequest(
lora_name=request.model,
lora_int_id=self.lora_mapper[request.model],
lora_path=lora_dir
)

def get_role() -> str:
if request.add_generation_prompt:
role = "assistant"
Expand Down Expand Up @@ -319,6 +346,24 @@ async def create_chat_response(

async def openai_completion(self, request: CompletionRequest, raw_request: Request) -> Response:

request_for_lora = request.model != self.model and request.model in self.lora_names

if request_for_lora and request.lora_request is None:
lora_dir = os.path.join(LORA_DIR, request.model)
if not os.path.exists(lora_dir):
logger.error(f"LORA_DIR {LORA_DIR} does not exist")
return self.create_error_response("Model not found")

if request.model not in self.lora_mapper:
self.adapter_id_counter += 1
self.lora_mapper[request.model] = self.adapter_id_counter

request.lora_request = LoRARequest(
lora_name=request.model,
lora_int_id=self.lora_mapper[request.model],
lora_path=lora_dir
)

async def completion_response(promise: RequestOutput,
postproc_params: Optional[PostprocParams]) -> CompletionResponse:
response = await promise
Expand Down