diff --git a/.github/actions/setup-trtllm/action.yml b/.github/actions/setup-trtllm/action.yml index 52b0c01c13..4c4c183dd2 100644 --- a/.github/actions/setup-trtllm/action.yml +++ b/.github/actions/setup-trtllm/action.yml @@ -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 diff --git a/e2e_test/chat_completions/test_openai_server.py b/e2e_test/chat_completions/test_openai_server.py index 699f634026..cd446b6109 100644 --- a/e2e_test/chat_completions/test_openai_server.py +++ b/e2e_test/chat_completions/test_openai_server.py @@ -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 # ------------------------------------------------------------------------- diff --git a/grpc_client/src/trtllm_service.rs b/grpc_client/src/trtllm_service.rs index 512b1df887..1bfd8c592b 100644 --- a/grpc_client/src/trtllm_service.rs +++ b/grpc_client/src/trtllm_service.rs @@ -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). + let stop_words = vec![]; let max_tokens = body.max_completion_tokens.unwrap_or(2048); @@ -478,14 +479,6 @@ impl TrtllmServiceClient { } } - /// Extract stop words from request - fn extract_stop_words(&self, request: &ChatCompletionRequest) -> Vec { - // 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, diff --git a/model_gateway/src/routers/grpc/regular/stages/chat/request_building.rs b/model_gateway/src/routers/grpc/regular/stages/chat/request_building.rs index e12256455c..f5b57580ee 100644 --- a/model_gateway/src/routers/grpc/regular/stages/chat/request_building.rs +++ b/model_gateway/src/routers/grpc/regular/stages/chat/request_building.rs @@ -10,6 +10,8 @@ use crate::routers::{ grpc::{ common::stages::{helpers, PipelineStage}, context::{ClientSelection, RequestContext}, + proto_wrapper::ProtoGenerateRequest, + utils, }, }; @@ -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); diff --git a/model_gateway/src/routers/grpc/regular/stages/generate/request_building.rs b/model_gateway/src/routers/grpc/regular/stages/generate/request_building.rs index e071fc8d1d..9060c9fec9 100644 --- a/model_gateway/src/routers/grpc/regular/stages/generate/request_building.rs +++ b/model_gateway/src/routers/grpc/regular/stages/generate/request_building.rs @@ -10,6 +10,8 @@ use crate::routers::{ grpc::{ common::stages::{helpers, PipelineStage}, context::{ClientSelection, RequestContext}, + proto_wrapper::ProtoGenerateRequest, + utils, }, }; @@ -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); + } + } + } + if self.inject_pd_metadata { if let Some(workers) = ctx.state.workers.as_ref() { helpers::maybe_inject_pd_metadata(&mut proto_request, workers); diff --git a/model_gateway/src/routers/grpc/utils.rs b/model_gateway/src/routers/grpc/utils.rs index d96ee36f1a..216657fa3a 100644 --- a/model_gateway/src/routers/grpc/utils.rs +++ b/model_gateway/src/routers/grpc/utils.rs @@ -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}, @@ -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, diff --git a/scripts/ci_install_trtllm.sh b/scripts/ci_install_trtllm.sh index 9fdc49c484..e4755a6865 100755 --- a/scripts/ci_install_trtllm.sh +++ b/scripts/ci_install_trtllm.sh @@ -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