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
2 changes: 0 additions & 2 deletions .github/actions/setup-trtllm/action.yml
Original file line number Diff line number Diff line change
Expand Up @@ -14,8 +14,6 @@ runs:
with:
path: /tmp/trtllm-wheel
key: trtllm-wheel-${{ runner.os }}-${{ hashFiles('scripts/ci_install_trtllm.sh') }}
restore-keys: |
trtllm-wheel-${{ runner.os }}-

- name: Install TRT-LLM
shell: bash
Expand Down
50 changes: 50 additions & 0 deletions e2e_test/chat_completions/test_openai_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -201,6 +201,56 @@ def test_retrieve_model(self, setup_backend):
with pytest.raises(openai.NotFoundError):
client.models.retrieve("non-existent-model")

def test_stop_sequences(self, setup_backend):
"""Test that stop sequences cause the model to stop generating."""
_, model, client, gateway = setup_backend

response = client.chat.completions.create(
model=model,
messages=[
{"role": "user", "content": "Count from 1 to 10: 1, 2, 3, 4, 5, 6, 7, 8, 9, 10"},
],
temperature=0,
max_tokens=50,
stop=[","],
)

assert response.choices[0].finish_reason == "stop"
content = response.choices[0].message.content
assert "," not in content, f"Stop sequence ',' should not appear in output: {content}"

def test_stop_sequences_stream(self, setup_backend):
"""Test that stop sequences work in streaming mode."""
_, model, client, gateway = setup_backend

chunks = list(
client.chat.completions.create(
model=model,
messages=[
{
"role": "user",
"content": "Count from 1 to 10: 1, 2, 3, 4, 5, 6, 7, 8, 9, 10",
},
],
temperature=0,
max_tokens=50,
stop=[","],
stream=True,
)
)

# Find the chunk with finish_reason
finish_reasons = [
c.choices[0].finish_reason for c in chunks if c.choices and c.choices[0].finish_reason
]
assert "stop" in finish_reasons

# Collect all content
content = "".join(
c.choices[0].delta.content for c in chunks if c.choices and c.choices[0].delta.content
)
assert "," not in content, f"Stop sequence ',' should not appear in output: {content}"

# -------------------------------------------------------------------------
# Helper methods
# -------------------------------------------------------------------------
Expand Down
13 changes: 3 additions & 10 deletions grpc_client/src/trtllm_service.rs
Original file line number Diff line number Diff line change
Expand Up @@ -267,8 +267,9 @@ impl TrtllmServiceClient {
// Build guided decoding params if needed
let guided_decoding = self.build_guided_decoding_from_chat(body, tool_call_constraint)?;

// Extract stop words
let stop_words = self.extract_stop_words(body);
// Stop words are injected by the router after building (via tokenization),
// since TRT-LLM requires tokenized stop sequences (Vec<TokenSequence>).
let stop_words = vec![];

let max_tokens = body.max_completion_tokens.unwrap_or(2048);

Expand Down Expand Up @@ -478,14 +479,6 @@ impl TrtllmServiceClient {
}
}

/// Extract stop words from request
fn extract_stop_words(&self, request: &ChatCompletionRequest) -> Vec<proto::TokenSequence> {
// Note: This returns empty because stop words need to be tokenized
// The router should handle tokenization of stop strings
let _ = request; // suppress unused warning
vec![]
}

/// Build GuidedDecodingParams from ChatCompletionRequest
fn build_guided_decoding_from_chat(
&self,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,8 @@ use crate::routers::{
grpc::{
common::stages::{helpers, PipelineStage},
context::{ClientSelection, RequestContext},
proto_wrapper::ProtoGenerateRequest,
utils,
},
};

Expand Down Expand Up @@ -76,6 +78,15 @@ impl PipelineStage for ChatRequestBuildingStage {
error::bad_request("invalid_request_parameters", format!("Invalid request parameters: {}", e))
})?;

// Inject tokenized stop sequences for TRT-LLM requests
if let ProtoGenerateRequest::Trtllm(ref mut req) = proto_request {
if let Some(stop) = &body_ref.stop {
if let Some(tokenizer) = ctx.state.tokenizer.as_ref() {
utils::inject_trtllm_stop_words(req, tokenizer.as_ref(), stop);
}
}
}

if self.inject_pd_metadata {
if let Some(workers) = ctx.state.workers.as_ref() {
helpers::maybe_inject_pd_metadata(&mut proto_request, workers);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,8 @@ use crate::routers::{
grpc::{
common::stages::{helpers, PipelineStage},
context::{ClientSelection, RequestContext},
proto_wrapper::ProtoGenerateRequest,
utils,
},
};

Expand Down Expand Up @@ -75,6 +77,19 @@ impl PipelineStage for GenerateRequestBuildingStage {
error::bad_request("build_request_failed", e)
})?;

// Inject tokenized stop sequences for TRT-LLM requests
if let ProtoGenerateRequest::Trtllm(ref mut req) = proto_request {
let stop = generate_request
.sampling_params
.as_ref()
.and_then(|p| p.stop.as_ref());
if let Some(stop) = stop {
if let Some(tokenizer) = ctx.state.tokenizer.as_ref() {
utils::inject_trtllm_stop_words(req, tokenizer.as_ref(), stop);
}
Comment thread
ppraneth marked this conversation as resolved.
}
}

if self.inject_pd_metadata {
if let Some(workers) = ctx.state.workers.as_ref() {
helpers::maybe_inject_pd_metadata(&mut proto_request, workers);
Expand Down
29 changes: 29 additions & 0 deletions model_gateway/src/routers/grpc/utils.rs
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ use super::{
};
use crate::{
core::Worker,
grpc_client::trtllm_proto::{GenerateRequest as TrtllmGenerateRequest, TokenSequence},
observability::metrics::metrics_labels,
protocols::{
chat::{ChatCompletionRequest, ChatMessage},
Expand Down Expand Up @@ -793,6 +794,34 @@ pub(crate) fn get_reasoning_parser(
}
}

/// Inject tokenized stop sequences into a TRT-LLM request.
///
/// TRT-LLM requires stop words as tokenized `TokenSequence` entries, unlike
/// SGLang/vLLM which handle stop strings at the client level.
pub(crate) fn inject_trtllm_stop_words(
req: &mut TrtllmGenerateRequest,
tokenizer: &dyn Tokenizer,
stop: &StringOrArray,
) {
for stop_str in stop.iter().filter(|s| !s.is_empty()) {
match tokenizer.encode(stop_str, false) {
Ok(encoding) => {
let token_ids = encoding.token_ids().to_vec();
if !token_ids.is_empty() {
req.stop_words.push(TokenSequence { token_ids });
}
}
Err(e) => {
warn!(
stop_string = %stop_str,
error = %e,
"Failed to tokenize stop sequence, skipping"
);
}
}
}
}

/// Create a fresh reasoning parser instance (for streaming where state isolation is needed)
pub(crate) fn create_reasoning_parser(
reasoning_parser_factory: &ReasoningParserFactory,
Expand Down
2 changes: 2 additions & 0 deletions scripts/ci_install_trtllm.sh
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,8 @@
# so we build from source (main branch) which compiles the C++
# extensions properly and includes the gRPC serve command.
#
# Cache version: 2 — rebuild to pick up gRPC stop_words fix (#11292)
#
# Prerequisites (expected on k8s-runner-gpu nodes):
# - NVIDIA driver 580+ (CUDA 13)
# - CUDA 13.0 toolkit at /usr/local/cuda-13.0
Expand Down
Loading