Skip to content
Merged
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
20 changes: 19 additions & 1 deletion python/openai/openai_frontend/engine/triton_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -105,6 +105,8 @@ class TritonModelMetadata:
tokenizer: Optional[Any]
# LoRA names supported by the backend
lora_names: Optional[List[str]]
# Name of the input tensor enabling "echo" parameter in /v1/completions endpoint
echo_tensor_name: Optional[str]
# Time that model was loaded by Triton
create_time: int
# Conversion format between OpenAI and Triton requests
Expand Down Expand Up @@ -202,7 +204,12 @@ async def chat(
# Convert to Triton request format and perform inference
responses = metadata.model.async_infer(
metadata.inference_request_converter(
metadata.model, prompt, request, lora_name, self.default_max_tokens
metadata.model,
prompt,
request,
lora_name,
metadata.echo_tensor_name,
self.default_max_tokens,
)
)

Expand Down Expand Up @@ -330,6 +337,7 @@ async def completion(
request.prompt,
request,
lora_name,
metadata.echo_tensor_name,
self.default_max_tokens,
)
)
Expand Down Expand Up @@ -484,12 +492,22 @@ def _get_model_metadata(self) -> Dict[str, TritonModelMetadata]:
self.server.options.model_repository, name, model.version
)

echo_tensor_name = None
for input in model.config()["input"]:
if input["name"] in [
"exclude_input_in_output",
"sampling_param_exclude_input_from_output",
]:
echo_tensor_name = input["name"]
break

metadata = TritonModelMetadata(
name=name,
backend=backend,
model=model,
tokenizer=self.tokenizer,
lora_names=lora_names,
echo_tensor_name=echo_tensor_name,
create_time=self.create_time,
inference_request_converter=self._determine_request_converter(
backend, RequestKind.GENERATION
Expand Down
11 changes: 7 additions & 4 deletions python/openai/openai_frontend/engine/utils/triton.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,7 @@ def _create_vllm_generate_request(
prompt,
request: CreateChatCompletionRequest | CreateCompletionRequest,
lora_name: str | None,
echo_tensor_name: str | None,
default_max_tokens: int,
):
inputs = {}
Expand Down Expand Up @@ -128,7 +129,7 @@ def _create_vllm_generate_request(

inputs["text_input"] = [prompt]
inputs["stream"] = np.bool_([request.stream])
inputs["exclude_input_in_output"] = np.bool_([exclude_input_in_output])
inputs[echo_tensor_name] = np.bool_([exclude_input_in_output])
# Pass sampling_parameters as serialized JSON string input to support List
# fields like 'stop' that aren't supported by TRITONSERVER_Parameters yet.
inputs["sampling_parameters"] = [sampling_parameters]
Expand All @@ -142,6 +143,7 @@ def _create_trtllm_generate_request(
prompt,
request: CreateChatCompletionRequest | CreateCompletionRequest,
lora_name: str | None,
echo_tensor_name: str | None,
default_max_tokens: int,
):
if lora_name is not None:
Expand Down Expand Up @@ -184,6 +186,10 @@ def _create_trtllm_generate_request(
inputs["seed"] = np.uint64([[request.seed]])
if request.temperature is not None:
inputs["temperature"] = np.float32([[request.temperature]])
# Only limited TRT-LLM models support "echo" (inflight_batcher_llm, disaggregated_serving, llmapi)
echo = getattr(request, "echo", None)
if echo is not None and echo_tensor_name is not None:
inputs[echo_tensor_name] = np.bool_([[not echo]])

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

inputs["return_num_input_tokens"] = np.bool_([[True]])
inputs["return_num_output_tokens"] = np.bool_([[True]])

# FIXME: TRT-LLM doesn't currently support runtime changes of 'echo' and it
# is configured at model load time, so we don't handle it here for now.
return model.create_request(inputs=inputs)


Expand Down
13 changes: 13 additions & 0 deletions python/openai/tests/test_completions.py
Original file line number Diff line number Diff line change
Expand Up @@ -367,6 +367,19 @@ def test_lora(self):
def test_multi_lora(self):
pass

@pytest.mark.parametrize("echo", [False, True])
def test_echo(self, client, model: str, prompt: str, echo: bool):
response = client.post(
"/v1/completions", json={"model": model, "prompt": prompt, "echo": echo}
)

response_text = response.json()["choices"][0]["text"].strip()
if echo:
assert response_text.startswith(prompt)
else:
# 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".
assert prompt not in response_text

def test_usage_response(self, client, model: str, prompt: str):
response = client.post(
"/v1/completions",
Expand Down
11 changes: 3 additions & 8 deletions python/openai/tests/test_openai_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -92,20 +92,15 @@ def test_openai_client_chat_completion(

@pytest.mark.parametrize("echo", [False, True])
def test_openai_client_completion_echo(
self, client: openai.OpenAI, echo: bool, backend: str, model: str, prompt: str
self, client: openai.OpenAI, echo: bool, model: str, prompt: str
):
if backend == "tensorrtllm":
pytest.skip(
reason="TRT-LLM backend currently only supports setting this parameter at model load time",
)

completion = client.completions.create(prompt=prompt, model=model, echo=echo)

print(f"Completion results: {completion}")
response = completion.choices[0].text
if echo:
assert prompt in response
assert response.startswith(prompt)
else:
# 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".
assert prompt not in response

@pytest.mark.skip(reason="Not Implemented Yet")
Expand Down
Loading