[model-gateway] add streaming metrics (TTFT, TPOT, tokens, duration) for gRPC router (#15125)
This commit is contained in:
@@ -836,6 +836,26 @@ pub mod smg_labels {
|
|||||||
/// SMG Metrics helper struct for the new layered metrics architecture
|
/// SMG Metrics helper struct for the new layered metrics architecture
|
||||||
pub struct SmgMetrics;
|
pub struct SmgMetrics;
|
||||||
|
|
||||||
|
/// Parameters for recording streaming metrics.
|
||||||
|
pub struct StreamingMetricsParams<'a> {
|
||||||
|
/// Router type label (e.g., "grpc", "http")
|
||||||
|
pub router_type: &'static str,
|
||||||
|
/// Backend type label (e.g., "regular", "pd")
|
||||||
|
pub backend_type: &'static str,
|
||||||
|
/// Model identifier (will be converted to owned String for metrics)
|
||||||
|
pub model_id: &'a str,
|
||||||
|
/// Endpoint label (e.g., "chat", "generate")
|
||||||
|
pub endpoint: &'static str,
|
||||||
|
/// Time to first token (None if no tokens were generated)
|
||||||
|
pub ttft: Option<Duration>,
|
||||||
|
/// Total generation time
|
||||||
|
pub generation_duration: Duration,
|
||||||
|
/// Input token count (None for endpoints that don't track this)
|
||||||
|
pub input_tokens: Option<u64>,
|
||||||
|
/// Output token count
|
||||||
|
pub output_tokens: u64,
|
||||||
|
}
|
||||||
|
|
||||||
impl SmgMetrics {
|
impl SmgMetrics {
|
||||||
/// Record an HTTP request
|
/// Record an HTTP request
|
||||||
pub fn record_http_request(method: &str, path: &str, status_class: &str) {
|
pub fn record_http_request(method: &str, path: &str, status_class: &str) {
|
||||||
@@ -1055,6 +1075,86 @@ impl SmgMetrics {
|
|||||||
.record(duration.as_secs_f64());
|
.record(duration.as_secs_f64());
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Record all streaming metrics in a single batch call.
|
||||||
|
///
|
||||||
|
/// This consolidates TTFT, TPOT, generation duration, and token metrics
|
||||||
|
/// into one function, handling TPOT calculation internally.
|
||||||
|
pub fn record_streaming_metrics(params: StreamingMetricsParams<'_>) {
|
||||||
|
let StreamingMetricsParams {
|
||||||
|
router_type,
|
||||||
|
backend_type,
|
||||||
|
model_id,
|
||||||
|
endpoint,
|
||||||
|
ttft,
|
||||||
|
generation_duration,
|
||||||
|
input_tokens,
|
||||||
|
output_tokens,
|
||||||
|
} = params;
|
||||||
|
// metrics-rs requires owned strings for dynamic labels (uses Cow<'static, str>).
|
||||||
|
// We allocate once and clone for each metric - unavoidable with this API.
|
||||||
|
let model = model_id.to_string();
|
||||||
|
|
||||||
|
// TTFT and TPOT (only if we have a first token time)
|
||||||
|
if let Some(ttft_duration) = ttft {
|
||||||
|
histogram!(
|
||||||
|
"smg_router_ttft_seconds",
|
||||||
|
"router_type" => router_type,
|
||||||
|
"backend_type" => backend_type,
|
||||||
|
"model" => model.clone(),
|
||||||
|
"endpoint" => endpoint
|
||||||
|
)
|
||||||
|
.record(ttft_duration.as_secs_f64());
|
||||||
|
|
||||||
|
// TPOT - only meaningful with >1 output token
|
||||||
|
if output_tokens > 1 {
|
||||||
|
let time_after_first = generation_duration.saturating_sub(ttft_duration);
|
||||||
|
let tpot = time_after_first / (output_tokens as u32 - 1);
|
||||||
|
histogram!(
|
||||||
|
"smg_router_tpot_seconds",
|
||||||
|
"router_type" => router_type,
|
||||||
|
"backend_type" => backend_type,
|
||||||
|
"model" => model.clone(),
|
||||||
|
"endpoint" => endpoint
|
||||||
|
)
|
||||||
|
.record(tpot.as_secs_f64());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Generation duration
|
||||||
|
histogram!(
|
||||||
|
"smg_router_generation_duration_seconds",
|
||||||
|
"router_type" => router_type,
|
||||||
|
"backend_type" => backend_type,
|
||||||
|
"model" => model.clone(),
|
||||||
|
"endpoint" => endpoint
|
||||||
|
)
|
||||||
|
.record(generation_duration.as_secs_f64());
|
||||||
|
|
||||||
|
// Input tokens (if available)
|
||||||
|
if let Some(input) = input_tokens {
|
||||||
|
counter!(
|
||||||
|
"smg_router_tokens_total",
|
||||||
|
"router_type" => router_type,
|
||||||
|
"backend_type" => backend_type,
|
||||||
|
"model" => model.clone(),
|
||||||
|
"endpoint" => endpoint,
|
||||||
|
"token_type" => smg_labels::TOKEN_INPUT
|
||||||
|
)
|
||||||
|
.increment(input);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Output tokens
|
||||||
|
counter!(
|
||||||
|
"smg_router_tokens_total",
|
||||||
|
"router_type" => router_type,
|
||||||
|
"backend_type" => backend_type,
|
||||||
|
"model" => model,
|
||||||
|
"endpoint" => endpoint,
|
||||||
|
"token_type" => smg_labels::TOKEN_OUTPUT
|
||||||
|
)
|
||||||
|
.increment(output_tokens);
|
||||||
|
}
|
||||||
|
|
||||||
// ========================================================================
|
// ========================================================================
|
||||||
// Layer 3: Worker metrics
|
// Layer 3: Worker metrics
|
||||||
// ========================================================================
|
// ========================================================================
|
||||||
|
|||||||
@@ -65,6 +65,7 @@ impl RequestPipeline {
|
|||||||
reasoning_parser_factory,
|
reasoning_parser_factory,
|
||||||
configured_tool_parser,
|
configured_tool_parser,
|
||||||
configured_reasoning_parser,
|
configured_reasoning_parser,
|
||||||
|
smg_labels::BACKEND_REGULAR,
|
||||||
));
|
));
|
||||||
|
|
||||||
let stages: Vec<Box<dyn PipelineStage>> = vec![
|
let stages: Vec<Box<dyn PipelineStage>> = vec![
|
||||||
@@ -171,6 +172,7 @@ impl RequestPipeline {
|
|||||||
reasoning_parser_factory,
|
reasoning_parser_factory,
|
||||||
configured_tool_parser,
|
configured_tool_parser,
|
||||||
configured_reasoning_parser,
|
configured_reasoning_parser,
|
||||||
|
smg_labels::BACKEND_PD,
|
||||||
));
|
));
|
||||||
|
|
||||||
let stages: Vec<Box<dyn PipelineStage>> = vec![
|
let stages: Vec<Box<dyn PipelineStage>> = vec![
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ use tracing::{debug, error, warn};
|
|||||||
|
|
||||||
use crate::{
|
use crate::{
|
||||||
grpc_client::sglang_proto::generate_complete::MatchedStop::{MatchedStopStr, MatchedTokenId},
|
grpc_client::sglang_proto::generate_complete::MatchedStop::{MatchedStopStr, MatchedTokenId},
|
||||||
|
observability::metrics::{smg_labels, SmgMetrics, StreamingMetricsParams},
|
||||||
protocols::{
|
protocols::{
|
||||||
chat::{ChatCompletionRequest, ChatCompletionStreamResponse},
|
chat::{ChatCompletionRequest, ChatCompletionStreamResponse},
|
||||||
common::{
|
common::{
|
||||||
@@ -43,6 +44,16 @@ pub struct StreamingProcessor {
|
|||||||
reasoning_parser_factory: ReasoningParserFactory,
|
reasoning_parser_factory: ReasoningParserFactory,
|
||||||
configured_tool_parser: Option<String>,
|
configured_tool_parser: Option<String>,
|
||||||
configured_reasoning_parser: Option<String>,
|
configured_reasoning_parser: Option<String>,
|
||||||
|
backend_type: &'static str,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Context for generate endpoint streaming - groups config params to reduce function arguments
|
||||||
|
struct GenerateStreamContext {
|
||||||
|
request_id: String,
|
||||||
|
weight_version: String,
|
||||||
|
return_logprob: bool,
|
||||||
|
backend_type: &'static str,
|
||||||
|
model: String,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl StreamingProcessor {
|
impl StreamingProcessor {
|
||||||
@@ -52,6 +63,7 @@ impl StreamingProcessor {
|
|||||||
reasoning_parser_factory: ReasoningParserFactory,
|
reasoning_parser_factory: ReasoningParserFactory,
|
||||||
configured_tool_parser: Option<String>,
|
configured_tool_parser: Option<String>,
|
||||||
configured_reasoning_parser: Option<String>,
|
configured_reasoning_parser: Option<String>,
|
||||||
|
backend_type: &'static str,
|
||||||
) -> Self {
|
) -> Self {
|
||||||
Self {
|
Self {
|
||||||
tokenizer,
|
tokenizer,
|
||||||
@@ -59,6 +71,7 @@ impl StreamingProcessor {
|
|||||||
reasoning_parser_factory,
|
reasoning_parser_factory,
|
||||||
configured_tool_parser,
|
configured_tool_parser,
|
||||||
configured_reasoning_parser,
|
configured_reasoning_parser,
|
||||||
|
backend_type,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -164,6 +177,10 @@ impl StreamingProcessor {
|
|||||||
original_request: Arc<ChatCompletionRequest>,
|
original_request: Arc<ChatCompletionRequest>,
|
||||||
tx: &UnboundedSender<Result<Bytes, io::Error>>,
|
tx: &UnboundedSender<Result<Bytes, io::Error>>,
|
||||||
) -> Result<(), String> {
|
) -> Result<(), String> {
|
||||||
|
// Metrics timing
|
||||||
|
let start_time = Instant::now();
|
||||||
|
let mut first_token_time: Option<Instant> = None;
|
||||||
|
|
||||||
// Extract request parameters
|
// Extract request parameters
|
||||||
let separate_reasoning = original_request.separate_reasoning;
|
let separate_reasoning = original_request.separate_reasoning;
|
||||||
let tool_choice = &original_request.tool_choice;
|
let tool_choice = &original_request.tool_choice;
|
||||||
@@ -246,6 +263,11 @@ impl StreamingProcessor {
|
|||||||
|
|
||||||
match gen_response.into_response() {
|
match gen_response.into_response() {
|
||||||
ProtoResponseVariant::Chunk(chunk) => {
|
ProtoResponseVariant::Chunk(chunk) => {
|
||||||
|
// Track TTFT immediately on first chunk received from backend
|
||||||
|
if first_token_time.is_none() {
|
||||||
|
first_token_time = Some(Instant::now());
|
||||||
|
}
|
||||||
|
|
||||||
let index = chunk.index();
|
let index = chunk.index();
|
||||||
|
|
||||||
// For vLLM, accumulate completion tokens (vLLM sends deltas)
|
// For vLLM, accumulate completion tokens (vLLM sends deltas)
|
||||||
@@ -548,6 +570,20 @@ impl StreamingProcessor {
|
|||||||
// Mark stream as completed successfully to prevent abort on drop
|
// Mark stream as completed successfully to prevent abort on drop
|
||||||
grpc_stream.mark_completed();
|
grpc_stream.mark_completed();
|
||||||
|
|
||||||
|
// Record streaming metrics
|
||||||
|
let total_prompt: u32 = prompt_tokens.values().sum();
|
||||||
|
let total_completion: u32 = completion_tokens.values().sum();
|
||||||
|
SmgMetrics::record_streaming_metrics(StreamingMetricsParams {
|
||||||
|
router_type: smg_labels::ROUTER_GRPC,
|
||||||
|
backend_type: self.backend_type,
|
||||||
|
model_id: model,
|
||||||
|
endpoint: smg_labels::ENDPOINT_CHAT,
|
||||||
|
ttft: first_token_time.map(|t| t.duration_since(start_time)),
|
||||||
|
generation_duration: start_time.elapsed(),
|
||||||
|
input_tokens: Some(total_prompt as u64),
|
||||||
|
output_tokens: total_completion as u64,
|
||||||
|
});
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -603,30 +639,28 @@ impl StreamingProcessor {
|
|||||||
generate_request: Arc<GenerateRequest>,
|
generate_request: Arc<GenerateRequest>,
|
||||||
dispatch: context::DispatchMetadata,
|
dispatch: context::DispatchMetadata,
|
||||||
) -> Response {
|
) -> Response {
|
||||||
let return_logprob = generate_request.return_logprob.unwrap_or(false);
|
|
||||||
|
|
||||||
// Create SSE channel
|
// Create SSE channel
|
||||||
let (tx, rx) = mpsc::unbounded_channel::<Result<Bytes, io::Error>>();
|
let (tx, rx) = mpsc::unbounded_channel::<Result<Bytes, io::Error>>();
|
||||||
|
|
||||||
|
// Build context once, clone for spawned task
|
||||||
|
let ctx = GenerateStreamContext {
|
||||||
|
request_id: dispatch.request_id.clone(),
|
||||||
|
weight_version: dispatch
|
||||||
|
.weight_version
|
||||||
|
.clone()
|
||||||
|
.unwrap_or_else(|| "default".to_string()),
|
||||||
|
return_logprob: generate_request.return_logprob.unwrap_or(false),
|
||||||
|
backend_type: self.backend_type,
|
||||||
|
model: dispatch.model.clone(),
|
||||||
|
};
|
||||||
|
|
||||||
// Spawn background task based on execution mode
|
// Spawn background task based on execution mode
|
||||||
match execution_result {
|
match execution_result {
|
||||||
context::ExecutionResult::Single { stream } => {
|
context::ExecutionResult::Single { stream } => {
|
||||||
let tokenizer = self.tokenizer.clone();
|
let tokenizer = self.tokenizer.clone();
|
||||||
let request_id = dispatch.request_id.clone();
|
|
||||||
let weight_version = dispatch
|
|
||||||
.weight_version
|
|
||||||
.clone()
|
|
||||||
.unwrap_or_else(|| "default".to_string());
|
|
||||||
tokio::spawn(async move {
|
tokio::spawn(async move {
|
||||||
let result = Self::process_generate_streaming(
|
let result =
|
||||||
tokenizer,
|
Self::process_generate_streaming(tokenizer, stream, ctx, &tx).await;
|
||||||
stream,
|
|
||||||
request_id,
|
|
||||||
weight_version,
|
|
||||||
return_logprob,
|
|
||||||
&tx,
|
|
||||||
)
|
|
||||||
.await;
|
|
||||||
|
|
||||||
if let Err(e) = result {
|
if let Err(e) = result {
|
||||||
let error_chunk = format!("data: {{\"error\": \"{}\"}}\n\n", e);
|
let error_chunk = format!("data: {{\"error\": \"{}\"}}\n\n", e);
|
||||||
@@ -637,22 +671,10 @@ impl StreamingProcessor {
|
|||||||
});
|
});
|
||||||
}
|
}
|
||||||
context::ExecutionResult::Dual { prefill, decode } => {
|
context::ExecutionResult::Dual { prefill, decode } => {
|
||||||
// For PD mode, need to handle prefill stream for input_logprobs
|
|
||||||
let tokenizer = self.tokenizer.clone();
|
let tokenizer = self.tokenizer.clone();
|
||||||
let request_id = dispatch.request_id.clone();
|
|
||||||
let weight_version = dispatch
|
|
||||||
.weight_version
|
|
||||||
.clone()
|
|
||||||
.unwrap_or_else(|| "default".to_string());
|
|
||||||
tokio::spawn(async move {
|
tokio::spawn(async move {
|
||||||
let result = Self::process_generate_streaming_dual(
|
let result = Self::process_generate_streaming_dual(
|
||||||
tokenizer,
|
tokenizer, prefill, *decode, ctx, &tx,
|
||||||
prefill,
|
|
||||||
*decode,
|
|
||||||
request_id,
|
|
||||||
weight_version,
|
|
||||||
return_logprob,
|
|
||||||
&tx,
|
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
|
|
||||||
@@ -675,12 +697,11 @@ impl StreamingProcessor {
|
|||||||
async fn process_generate_streaming(
|
async fn process_generate_streaming(
|
||||||
tokenizer: Arc<dyn Tokenizer>,
|
tokenizer: Arc<dyn Tokenizer>,
|
||||||
mut stream: ProtoStream,
|
mut stream: ProtoStream,
|
||||||
request_id: String,
|
ctx: GenerateStreamContext,
|
||||||
weight_version: String,
|
|
||||||
_include_logprobs: bool,
|
|
||||||
tx: &UnboundedSender<Result<Bytes, io::Error>>,
|
tx: &UnboundedSender<Result<Bytes, io::Error>>,
|
||||||
) -> Result<(), String> {
|
) -> Result<(), String> {
|
||||||
let start_time = Instant::now();
|
let start_time = Instant::now();
|
||||||
|
let mut first_token_time: Option<Instant> = None;
|
||||||
|
|
||||||
// Track state per index for n>1 case
|
// Track state per index for n>1 case
|
||||||
let mut accumulated_texts: HashMap<u32, String> = HashMap::new();
|
let mut accumulated_texts: HashMap<u32, String> = HashMap::new();
|
||||||
@@ -691,6 +712,11 @@ impl StreamingProcessor {
|
|||||||
|
|
||||||
match gen_response.into_response() {
|
match gen_response.into_response() {
|
||||||
ProtoResponseVariant::Chunk(chunk) => {
|
ProtoResponseVariant::Chunk(chunk) => {
|
||||||
|
// Track TTFT immediately on first chunk received from backend
|
||||||
|
if first_token_time.is_none() {
|
||||||
|
first_token_time = Some(Instant::now());
|
||||||
|
}
|
||||||
|
|
||||||
let index = chunk.index();
|
let index = chunk.index();
|
||||||
|
|
||||||
// Both backends send delta token_ids, so accumulate for both
|
// Both backends send delta token_ids, so accumulate for both
|
||||||
@@ -708,7 +734,7 @@ impl StreamingProcessor {
|
|||||||
accumulated_text.push_str(&chunk_text);
|
accumulated_text.push_str(&chunk_text);
|
||||||
|
|
||||||
// Generate unique ID per index
|
// Generate unique ID per index
|
||||||
let index_id = format!("{}-{}", request_id, index);
|
let index_id = format!("{}-{}", ctx.request_id, index);
|
||||||
|
|
||||||
// Build streaming response chunk (SGLang format)
|
// Build streaming response chunk (SGLang format)
|
||||||
let chunk_response = serde_json::json!({
|
let chunk_response = serde_json::json!({
|
||||||
@@ -718,7 +744,7 @@ impl StreamingProcessor {
|
|||||||
"id": index_id,
|
"id": index_id,
|
||||||
"finish_reason": null,
|
"finish_reason": null,
|
||||||
"prompt_tokens": chunk.prompt_tokens(),
|
"prompt_tokens": chunk.prompt_tokens(),
|
||||||
"weight_version": &weight_version,
|
"weight_version": &ctx.weight_version,
|
||||||
"completion_tokens": current_completion_tokens,
|
"completion_tokens": current_completion_tokens,
|
||||||
"cached_tokens": chunk.cached_tokens()
|
"cached_tokens": chunk.cached_tokens()
|
||||||
},
|
},
|
||||||
@@ -737,7 +763,7 @@ impl StreamingProcessor {
|
|||||||
let accumulated_text =
|
let accumulated_text =
|
||||||
accumulated_texts.get(&index).cloned().unwrap_or_default();
|
accumulated_texts.get(&index).cloned().unwrap_or_default();
|
||||||
let completion_tokens = *completion_tokens_map.get(&index).unwrap_or(&0);
|
let completion_tokens = *completion_tokens_map.get(&index).unwrap_or(&0);
|
||||||
let index_id = format!("{}-{}", request_id, index);
|
let index_id = format!("{}-{}", ctx.request_id, index);
|
||||||
let e2e_latency = start_time.elapsed().as_secs_f64();
|
let e2e_latency = start_time.elapsed().as_secs_f64();
|
||||||
|
|
||||||
// Send final chunk with finish_reason
|
// Send final chunk with finish_reason
|
||||||
@@ -748,7 +774,7 @@ impl StreamingProcessor {
|
|||||||
"id": index_id,
|
"id": index_id,
|
||||||
"finish_reason": complete.finish_reason(),
|
"finish_reason": complete.finish_reason(),
|
||||||
"prompt_tokens": complete.prompt_tokens(),
|
"prompt_tokens": complete.prompt_tokens(),
|
||||||
"weight_version": &weight_version,
|
"weight_version": &ctx.weight_version,
|
||||||
"completion_tokens": completion_tokens,
|
"completion_tokens": completion_tokens,
|
||||||
"cached_tokens": complete.cached_tokens(),
|
"cached_tokens": complete.cached_tokens(),
|
||||||
"e2e_latency": e2e_latency
|
"e2e_latency": e2e_latency
|
||||||
@@ -775,6 +801,10 @@ impl StreamingProcessor {
|
|||||||
// Mark stream as completed successfully to prevent abort on drop
|
// Mark stream as completed successfully to prevent abort on drop
|
||||||
stream.mark_completed();
|
stream.mark_completed();
|
||||||
|
|
||||||
|
// Record streaming metrics
|
||||||
|
let total_completion: u32 = completion_tokens_map.values().sum();
|
||||||
|
Self::record_generate_metrics(start_time, first_token_time, total_completion, &ctx);
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -783,13 +813,11 @@ impl StreamingProcessor {
|
|||||||
tokenizer: Arc<dyn Tokenizer>,
|
tokenizer: Arc<dyn Tokenizer>,
|
||||||
mut prefill_stream: ProtoStream,
|
mut prefill_stream: ProtoStream,
|
||||||
decode_stream: ProtoStream,
|
decode_stream: ProtoStream,
|
||||||
request_id: String,
|
ctx: GenerateStreamContext,
|
||||||
weight_version: String,
|
|
||||||
return_logprob: bool,
|
|
||||||
tx: &UnboundedSender<Result<Bytes, io::Error>>,
|
tx: &UnboundedSender<Result<Bytes, io::Error>>,
|
||||||
) -> Result<(), String> {
|
) -> Result<(), String> {
|
||||||
// Collect input_logprobs from prefill stream if requested
|
// Collect input_logprobs from prefill stream if requested
|
||||||
let input_token_logprobs = if return_logprob {
|
let input_token_logprobs = if ctx.return_logprob {
|
||||||
let mut input_logprobs = None;
|
let mut input_logprobs = None;
|
||||||
while let Some(response) = prefill_stream.next().await {
|
while let Some(response) = prefill_stream.next().await {
|
||||||
let gen_response = response.map_err(|e| format!("Prefill stream error: {}", e))?;
|
let gen_response = response.map_err(|e| format!("Prefill stream error: {}", e))?;
|
||||||
@@ -817,9 +845,7 @@ impl StreamingProcessor {
|
|||||||
let result = Self::process_generate_streaming_with_input_logprobs(
|
let result = Self::process_generate_streaming_with_input_logprobs(
|
||||||
tokenizer,
|
tokenizer,
|
||||||
decode_stream,
|
decode_stream,
|
||||||
request_id,
|
ctx,
|
||||||
weight_version,
|
|
||||||
return_logprob,
|
|
||||||
input_token_logprobs,
|
input_token_logprobs,
|
||||||
tx,
|
tx,
|
||||||
)
|
)
|
||||||
@@ -838,13 +864,12 @@ impl StreamingProcessor {
|
|||||||
async fn process_generate_streaming_with_input_logprobs(
|
async fn process_generate_streaming_with_input_logprobs(
|
||||||
tokenizer: Arc<dyn Tokenizer>,
|
tokenizer: Arc<dyn Tokenizer>,
|
||||||
mut stream: ProtoStream,
|
mut stream: ProtoStream,
|
||||||
request_id: String,
|
ctx: GenerateStreamContext,
|
||||||
weight_version: String,
|
|
||||||
_include_logprobs: bool,
|
|
||||||
input_token_logprobs: Option<Vec<Vec<Option<f64>>>>,
|
input_token_logprobs: Option<Vec<Vec<Option<f64>>>>,
|
||||||
tx: &UnboundedSender<Result<Bytes, io::Error>>,
|
tx: &UnboundedSender<Result<Bytes, io::Error>>,
|
||||||
) -> Result<(), String> {
|
) -> Result<(), String> {
|
||||||
let start_time = Instant::now();
|
let start_time = Instant::now();
|
||||||
|
let mut first_token_time: Option<Instant> = None;
|
||||||
|
|
||||||
// Track state per index for n>1 case
|
// Track state per index for n>1 case
|
||||||
let mut accumulated_texts: HashMap<u32, String> = HashMap::new();
|
let mut accumulated_texts: HashMap<u32, String> = HashMap::new();
|
||||||
@@ -857,6 +882,11 @@ impl StreamingProcessor {
|
|||||||
|
|
||||||
match gen_response.into_response() {
|
match gen_response.into_response() {
|
||||||
ProtoResponseVariant::Chunk(chunk) => {
|
ProtoResponseVariant::Chunk(chunk) => {
|
||||||
|
// Track TTFT immediately on first chunk received from backend
|
||||||
|
if first_token_time.is_none() {
|
||||||
|
first_token_time = Some(Instant::now());
|
||||||
|
}
|
||||||
|
|
||||||
let index = chunk.index();
|
let index = chunk.index();
|
||||||
|
|
||||||
// Both backends send delta token_ids, so accumulate for both
|
// Both backends send delta token_ids, so accumulate for both
|
||||||
@@ -880,7 +910,7 @@ impl StreamingProcessor {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Generate unique ID per index
|
// Generate unique ID per index
|
||||||
let index_id = format!("{}-{}", request_id, index);
|
let index_id = format!("{}-{}", ctx.request_id, index);
|
||||||
|
|
||||||
// Build streaming response chunk with cumulative logprobs
|
// Build streaming response chunk with cumulative logprobs
|
||||||
let current_output_logprobs = accumulated_output_logprobs
|
let current_output_logprobs = accumulated_output_logprobs
|
||||||
@@ -894,7 +924,7 @@ impl StreamingProcessor {
|
|||||||
"id": index_id,
|
"id": index_id,
|
||||||
"finish_reason": null,
|
"finish_reason": null,
|
||||||
"prompt_tokens": chunk.prompt_tokens(),
|
"prompt_tokens": chunk.prompt_tokens(),
|
||||||
"weight_version": &weight_version,
|
"weight_version": &ctx.weight_version,
|
||||||
"input_token_logprobs": input_token_logprobs.as_ref(),
|
"input_token_logprobs": input_token_logprobs.as_ref(),
|
||||||
"output_token_logprobs": current_output_logprobs,
|
"output_token_logprobs": current_output_logprobs,
|
||||||
"completion_tokens": current_completion_tokens,
|
"completion_tokens": current_completion_tokens,
|
||||||
@@ -921,7 +951,7 @@ impl StreamingProcessor {
|
|||||||
let final_output_logprobs = accumulated_output_logprobs
|
let final_output_logprobs = accumulated_output_logprobs
|
||||||
.get(&index)
|
.get(&index)
|
||||||
.and_then(|o| o.as_ref());
|
.and_then(|o| o.as_ref());
|
||||||
let index_id = format!("{}-{}", request_id, index);
|
let index_id = format!("{}-{}", ctx.request_id, index);
|
||||||
let e2e_latency = start_time.elapsed().as_secs_f64();
|
let e2e_latency = start_time.elapsed().as_secs_f64();
|
||||||
|
|
||||||
// Parse finish_reason
|
// Parse finish_reason
|
||||||
@@ -938,7 +968,7 @@ impl StreamingProcessor {
|
|||||||
"id": index_id,
|
"id": index_id,
|
||||||
"finish_reason": finish_reason,
|
"finish_reason": finish_reason,
|
||||||
"prompt_tokens": complete.prompt_tokens(),
|
"prompt_tokens": complete.prompt_tokens(),
|
||||||
"weight_version": &weight_version,
|
"weight_version": &ctx.weight_version,
|
||||||
"input_token_logprobs": input_token_logprobs.as_ref(),
|
"input_token_logprobs": input_token_logprobs.as_ref(),
|
||||||
"output_token_logprobs": final_output_logprobs,
|
"output_token_logprobs": final_output_logprobs,
|
||||||
"completion_tokens": completion_tokens,
|
"completion_tokens": completion_tokens,
|
||||||
@@ -967,6 +997,10 @@ impl StreamingProcessor {
|
|||||||
// Mark stream as completed successfully to prevent abort on drop
|
// Mark stream as completed successfully to prevent abort on drop
|
||||||
stream.mark_completed();
|
stream.mark_completed();
|
||||||
|
|
||||||
|
// Record streaming metrics
|
||||||
|
let total_completion: u32 = completion_tokens_map.values().sum();
|
||||||
|
Self::record_generate_metrics(start_time, first_token_time, total_completion, &ctx);
|
||||||
|
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -974,6 +1008,25 @@ impl StreamingProcessor {
|
|||||||
// Helper Methods
|
// Helper Methods
|
||||||
// ========================================================================
|
// ========================================================================
|
||||||
|
|
||||||
|
/// Record streaming metrics for generate endpoint
|
||||||
|
fn record_generate_metrics(
|
||||||
|
start_time: Instant,
|
||||||
|
first_token_time: Option<Instant>,
|
||||||
|
total_completion: u32,
|
||||||
|
ctx: &GenerateStreamContext,
|
||||||
|
) {
|
||||||
|
SmgMetrics::record_streaming_metrics(StreamingMetricsParams {
|
||||||
|
router_type: smg_labels::ROUTER_GRPC,
|
||||||
|
backend_type: ctx.backend_type,
|
||||||
|
model_id: &ctx.model,
|
||||||
|
endpoint: smg_labels::ENDPOINT_GENERATE,
|
||||||
|
ttft: first_token_time.map(|t| t.duration_since(start_time)),
|
||||||
|
generation_duration: start_time.elapsed(),
|
||||||
|
input_tokens: None, // generate endpoint doesn't expose prompt tokens in streaming
|
||||||
|
output_tokens: total_completion as u64,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
/// Process a chunk of tokens through the stop decoder
|
/// Process a chunk of tokens through the stop decoder
|
||||||
fn process_chunk_tokens(
|
fn process_chunk_tokens(
|
||||||
stop_decoder: &mut StopSequenceDecoder,
|
stop_decoder: &mut StopSequenceDecoder,
|
||||||
|
|||||||
Reference in New Issue
Block a user