Skip to content

Commit 116c3b4

Browse files
authored
feat: Support "echo" for TensorRT-LLM models in OpenAI API frontend v1/completions (#8523)
1 parent dd99592 commit 116c3b4

4 files changed

Lines changed: 42 additions & 13 deletions

File tree

python/openai/openai_frontend/engine/triton_engine.py

Lines changed: 19 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -105,6 +105,8 @@ class TritonModelMetadata:
105105
tokenizer: Optional[Any]
106106
# LoRA names supported by the backend
107107
lora_names: Optional[List[str]]
108+
# Name of the input tensor enabling "echo" parameter in /v1/completions endpoint
109+
echo_tensor_name: Optional[str]
108110
# Time that model was loaded by Triton
109111
create_time: int
110112
# Conversion format between OpenAI and Triton requests
@@ -202,7 +204,12 @@ async def chat(
202204
# Convert to Triton request format and perform inference
203205
responses = metadata.model.async_infer(
204206
metadata.inference_request_converter(
205-
metadata.model, prompt, request, lora_name, self.default_max_tokens
207+
metadata.model,
208+
prompt,
209+
request,
210+
lora_name,
211+
metadata.echo_tensor_name,
212+
self.default_max_tokens,
206213
)
207214
)
208215

@@ -330,6 +337,7 @@ async def completion(
330337
request.prompt,
331338
request,
332339
lora_name,
340+
metadata.echo_tensor_name,
333341
self.default_max_tokens,
334342
)
335343
)
@@ -484,12 +492,22 @@ def _get_model_metadata(self) -> Dict[str, TritonModelMetadata]:
484492
self.server.options.model_repository, name, model.version
485493
)
486494

495+
echo_tensor_name = None
496+
for input in model.config()["input"]:
497+
if input["name"] in [
498+
"exclude_input_in_output",
499+
"sampling_param_exclude_input_from_output",
500+
]:
501+
echo_tensor_name = input["name"]
502+
break
503+
487504
metadata = TritonModelMetadata(
488505
name=name,
489506
backend=backend,
490507
model=model,
491508
tokenizer=self.tokenizer,
492509
lora_names=lora_names,
510+
echo_tensor_name=echo_tensor_name,
493511
create_time=self.create_time,
494512
inference_request_converter=self._determine_request_converter(
495513
backend, RequestKind.GENERATION

python/openai/openai_frontend/engine/utils/triton.py

Lines changed: 7 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -57,6 +57,7 @@ def _create_vllm_generate_request(
5757
prompt,
5858
request: CreateChatCompletionRequest | CreateCompletionRequest,
5959
lora_name: str | None,
60+
echo_tensor_name: str | None,
6061
default_max_tokens: int,
6162
):
6263
inputs = {}
@@ -128,7 +129,7 @@ def _create_vllm_generate_request(
128129

129130
inputs["text_input"] = [prompt]
130131
inputs["stream"] = np.bool_([request.stream])
131-
inputs["exclude_input_in_output"] = np.bool_([exclude_input_in_output])
132+
inputs[echo_tensor_name] = np.bool_([exclude_input_in_output])
132133
# Pass sampling_parameters as serialized JSON string input to support List
133134
# fields like 'stop' that aren't supported by TRITONSERVER_Parameters yet.
134135
inputs["sampling_parameters"] = [sampling_parameters]
@@ -142,6 +143,7 @@ def _create_trtllm_generate_request(
142143
prompt,
143144
request: CreateChatCompletionRequest | CreateCompletionRequest,
144145
lora_name: str | None,
146+
echo_tensor_name: str | None,
145147
default_max_tokens: int,
146148
):
147149
if lora_name is not None:
@@ -184,6 +186,10 @@ def _create_trtllm_generate_request(
184186
inputs["seed"] = np.uint64([[request.seed]])
185187
if request.temperature is not None:
186188
inputs["temperature"] = np.float32([[request.temperature]])
189+
# Only limited TRT-LLM models support "echo" (inflight_batcher_llm, disaggregated_serving, llmapi)
190+
echo = getattr(request, "echo", None)
191+
if echo is not None and echo_tensor_name is not None:
192+
inputs[echo_tensor_name] = np.bool_([[not echo]])
187193

188194
guided_json = _get_guided_json_from_tool(request)
189195
if guided_json is not None:
@@ -192,9 +198,6 @@ def _create_trtllm_generate_request(
192198

193199
inputs["return_num_input_tokens"] = np.bool_([[True]])
194200
inputs["return_num_output_tokens"] = np.bool_([[True]])
195-
196-
# FIXME: TRT-LLM doesn't currently support runtime changes of 'echo' and it
197-
# is configured at model load time, so we don't handle it here for now.
198201
return model.create_request(inputs=inputs)
199202

200203

python/openai/tests/test_completions.py

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -367,6 +367,19 @@ def test_lora(self):
367367
def test_multi_lora(self):
368368
pass
369369

370+
@pytest.mark.parametrize("echo", [False, True])
371+
def test_echo(self, client, model: str, prompt: str, echo: bool):
372+
response = client.post(
373+
"/v1/completions", json={"model": model, "prompt": prompt, "echo": echo}
374+
)
375+
376+
response_text = response.json()["choices"][0]["text"].strip()
377+
if echo:
378+
assert response_text.startswith(prompt)
379+
else:
380+
# TODO: Consider using a different prompt. In TRT-LLM model, the second response may contain the prompt in the middle of the response even if echo is False, e.g. " Briefly explained.\nWhat is machine learning? She learns from data\nmachine learning".
381+
assert prompt not in response_text
382+
370383
def test_usage_response(self, client, model: str, prompt: str):
371384
response = client.post(
372385
"/v1/completions",

python/openai/tests/test_openai_client.py

Lines changed: 3 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -92,20 +92,15 @@ def test_openai_client_chat_completion(
9292

9393
@pytest.mark.parametrize("echo", [False, True])
9494
def test_openai_client_completion_echo(
95-
self, client: openai.OpenAI, echo: bool, backend: str, model: str, prompt: str
95+
self, client: openai.OpenAI, echo: bool, model: str, prompt: str
9696
):
97-
if backend == "tensorrtllm":
98-
pytest.skip(
99-
reason="TRT-LLM backend currently only supports setting this parameter at model load time",
100-
)
101-
10297
completion = client.completions.create(prompt=prompt, model=model, echo=echo)
10398

104-
print(f"Completion results: {completion}")
10599
response = completion.choices[0].text
106100
if echo:
107-
assert prompt in response
101+
assert response.startswith(prompt)
108102
else:
103+
# TODO: Consider using a different prompt. In TRT-LLM model, the second response may contain the prompt in the middle of the response even if echo is False, e.g. " Briefly explained.\nWhat is machine learning? She learns from data\nmachine learning".
109104
assert prompt not in response
110105

111106
@pytest.mark.skip(reason="Not Implemented Yet")

0 commit comments

Comments
 (0)