From 02690d63471cb24bc61f2519c909e65c82215371 Mon Sep 17 00:00:00 2001 From: ai-jz Date: Sun, 12 Apr 2026 21:51:35 -0700 Subject: [PATCH] feat(grpc): support --sampling-defaults in sglang gRPC servicer The gRPC path does not support --sampling-defaults model or --preferred-sampling-params. The HTTP path reads generation_config.json and applies defaults for unset request fields; the gRPC path uses OpenAI defaults regardless. Fix across three layers: - Proto: make 7 sampling fields optional (HasField detection) - Rust client: pass Option through in all 5 sampling param builders - Python servicer: load defaults at init, apply 3-tier cascade (user-set > preferred > generation_config > OpenAI defaults) Signed-off-by: ai-jz --- .../grpc_client/proto/sglang_scheduler.proto | 23 +-- crates/grpc_client/src/sglang_scheduler.rs | 151 +++++++++++------- .../smg_grpc_servicer/sglang/server.py | 1 + .../smg_grpc_servicer/sglang/servicer.py | 71 +++++++- 4 files changed, 170 insertions(+), 76 deletions(-) diff --git a/crates/grpc_client/proto/sglang_scheduler.proto b/crates/grpc_client/proto/sglang_scheduler.proto index 55d58a1bba..bfed3097bd 100644 --- a/crates/grpc_client/proto/sglang_scheduler.proto +++ b/crates/grpc_client/proto/sglang_scheduler.proto @@ -44,18 +44,19 @@ service SglangScheduler { // Sampling parameters matching SGLang's SamplingParams // -// IMPORTANT: Do not use SamplingParams::default() directly! -// The proto3 defaults (0 for numeric fields) do NOT match the semantic defaults -// (temperature=1.0, top_p=1.0, top_k=-1, etc.). Always construct with explicit values -// or use the conversion functions in sglang_scheduler.rs / grpc_server.py. +// The 7 core sampling fields (temperature, top_p, top_k, min_p, +// frequency_penalty, presence_penalty, repetition_penalty) are proto3 +// `optional` so the servicer can distinguish "user set this to 0" from +// "user did not set this". Unset fields fall back to: +// preferred_sampling_params > model generation_config > OpenAI defaults. message SamplingParams { - float temperature = 1; - float top_p = 2; - int32 top_k = 3; - float min_p = 4; - float frequency_penalty = 5; - float presence_penalty = 6; - float repetition_penalty = 7; + optional float temperature = 1; + optional float top_p = 2; + optional int32 top_k = 3; + optional float min_p = 4; + optional float frequency_penalty = 5; + optional float presence_penalty = 6; + optional float repetition_penalty = 7; optional uint32 max_new_tokens = 8; repeated string stop = 9; diff --git a/crates/grpc_client/src/sglang_scheduler.rs b/crates/grpc_client/src/sglang_scheduler.rs index d29299aaca..f75c74dd3c 100644 --- a/crates/grpc_client/src/sglang_scheduler.rs +++ b/crates/grpc_client/src/sglang_scheduler.rs @@ -446,14 +446,18 @@ impl SglangSchedulerClient { // Detokenization happens on the SMG Rust side (StopDecoder/Sequence). let skip_special_tokens = true; + // The 7 optional sampling fields pass Option through to proto `optional`, + // letting the servicer's HasField-based cascade resolve defaults from + // generation_config. Deploy note: update the servicer before the gateway + // so HasField logic is in place when None (wire-absent) values arrive. Ok(proto::SamplingParams { - temperature: request.temperature.unwrap_or(1.0), - top_p: request.top_p.unwrap_or(1.0), - top_k: request.top_k.unwrap_or(-1), - min_p: request.min_p.unwrap_or(0.0), - frequency_penalty: request.frequency_penalty.unwrap_or(0.0), - presence_penalty: request.presence_penalty.unwrap_or(0.0), - repetition_penalty: request.repetition_penalty.unwrap_or(1.0), + temperature: request.temperature, + top_p: request.top_p, + top_k: request.top_k, + min_p: request.min_p, + frequency_penalty: request.frequency_penalty, + presence_penalty: request.presence_penalty, + repetition_penalty: request.repetition_penalty, max_new_tokens, stop: stop_sequences, stop_token_ids: request.stop_token_ids.clone().unwrap_or_default(), @@ -551,14 +555,17 @@ impl SglangSchedulerClient { let max_new_tokens = request.max_output_tokens; + // top_k, min_p, repetition_penalty are non-Option on ResponsesRequest + // (can't distinguish user-set from default). Pass None to let the + // engine apply its own defaults, same as unset params via HTTP. Ok(proto::SamplingParams { - temperature: request.temperature.unwrap_or(1.0), - top_p: request.top_p.unwrap_or(1.0), - top_k: -1, // ResponsesRequest doesn't expose top_k - min_p: 0.0, // ResponsesRequest doesn't expose min_p - frequency_penalty: 0.0, // ResponsesRequest doesn't expose frequency_penalty - presence_penalty: 0.0, // ResponsesRequest doesn't expose presence_penalty - repetition_penalty: 1.0, // ResponsesRequest doesn't expose repetition_penalty + temperature: request.temperature, + top_p: request.top_p, + top_k: None, + min_p: None, + frequency_penalty: request.frequency_penalty, + presence_penalty: request.presence_penalty, + repetition_penalty: None, max_new_tokens, stop: vec![], // No stop sequences in Responses API stop_token_ids: vec![], // Handled by Harmony stop tokens @@ -645,13 +652,13 @@ impl SglangSchedulerClient { let skip_special_tokens = true; Ok(proto::SamplingParams { - temperature: request.temperature.unwrap_or(1.0) as f32, - top_p: request.top_p.unwrap_or(1.0) as f32, - top_k: request.top_k.map(|v| v as i32).unwrap_or(-1), - min_p: 0.0, - frequency_penalty: 0.0, - presence_penalty: 0.0, - repetition_penalty: 1.0, + temperature: request.temperature.map(|v| v as f32), + top_p: request.top_p.map(|v| v as f32), + top_k: request.top_k.map(|v| v as i32), + min_p: None, // Messages API doesn't expose min_p + frequency_penalty: None, // Messages API doesn't expose frequency_penalty + presence_penalty: None, // Messages API doesn't expose presence_penalty + repetition_penalty: None, // Messages API doesn't expose repetition_penalty max_new_tokens: Some(request.max_tokens), stop: stop_sequences, stop_token_ids: vec![], @@ -710,13 +717,13 @@ impl SglangSchedulerClient { let constraint = Self::build_single_constraint_from_completion(request)?; Ok(proto::SamplingParams { - temperature: request.temperature.unwrap_or(1.0), - top_p: request.top_p.unwrap_or(1.0), - top_k: request.top_k.unwrap_or(-1), - min_p: request.min_p.unwrap_or(0.0), - frequency_penalty: request.frequency_penalty.unwrap_or(0.0), - presence_penalty: request.presence_penalty.unwrap_or(0.0), - repetition_penalty: request.repetition_penalty.unwrap_or(1.0), + temperature: request.temperature, + top_p: request.top_p, + top_k: request.top_k, + min_p: request.min_p, + frequency_penalty: request.frequency_penalty, + presence_penalty: request.presence_penalty, + repetition_penalty: request.repetition_penalty, max_new_tokens: request.max_tokens, min_new_tokens: request.min_tokens.unwrap_or(0), stop: stop_sequences, @@ -785,10 +792,6 @@ impl SglangSchedulerClient { params: Option<&GenerateSamplingParams>, ) -> Result { let mut sampling = proto::SamplingParams { - temperature: 1.0, - top_p: 1.0, - top_k: -1, - repetition_penalty: 1.0, n: 1, skip_special_tokens: true, spaces_between_special_tokens: true, @@ -799,8 +802,17 @@ impl SglangSchedulerClient { return Ok(sampling); }; - // Simple field mappings using a macro - macro_rules! map_field { + // Pass optional sampling fields through directly (Option → Option) + sampling.temperature = p.temperature; + sampling.top_p = p.top_p; + sampling.top_k = p.top_k; + sampling.frequency_penalty = p.frequency_penalty; + sampling.presence_penalty = p.presence_penalty; + sampling.repetition_penalty = p.repetition_penalty; + sampling.min_p = p.min_p; + + // Bool fields: unwrap Option into proto bool + macro_rules! map_bool_field { ($field:ident) => { if let Some(val) = p.$field { sampling.$field = val; @@ -808,16 +820,9 @@ impl SglangSchedulerClient { }; } - map_field!(temperature); - map_field!(top_p); - map_field!(top_k); - map_field!(frequency_penalty); - map_field!(presence_penalty); - map_field!(repetition_penalty); - map_field!(min_p); - map_field!(ignore_eos); - map_field!(skip_special_tokens); - map_field!(no_stop_trim); + map_bool_field!(ignore_eos); + map_bool_field!(skip_special_tokens); + map_bool_field!(no_stop_trim); // Handle stop sequences if let Some(stop) = &p.stop { @@ -897,10 +902,10 @@ mod tests { #[test] fn test_generate_request_construction() { let sampling_params = proto::SamplingParams { - temperature: 0.7, + temperature: Some(0.7), max_new_tokens: Some(128), - top_p: 0.9, - top_k: 50, + top_p: Some(0.9), + top_k: Some(50), stop: vec!["".to_string()], ..Default::default() }; @@ -926,7 +931,7 @@ mod tests { assert_eq!(gen_req.top_logprobs_num, 5); let params = gen_req.sampling_params.unwrap(); - assert_eq!(params.temperature, 0.7); + assert_eq!(params.temperature, Some(0.7)); assert_eq!(params.max_new_tokens, Some(128)); assert_eq!(params.stop, vec![""]); } @@ -950,11 +955,15 @@ mod tests { #[test] fn test_sampling_params_defaults() { let params = proto::SamplingParams::default(); - // Numeric fields have proto defaults (0) - assert_eq!(params.temperature, 0.0); - assert_eq!(params.top_p, 0.0); - assert_eq!(params.top_k, 0); - assert_eq!(params.repetition_penalty, 0.0); + // Optional sampling fields default to None (unset) + assert_eq!(params.temperature, None); + assert_eq!(params.top_p, None); + assert_eq!(params.top_k, None); + assert_eq!(params.min_p, None); + assert_eq!(params.frequency_penalty, None); + assert_eq!(params.presence_penalty, None); + assert_eq!(params.repetition_penalty, None); + // Non-optional numeric fields have proto defaults (0) assert_eq!(params.n, 0); // Bool fields have proto defaults (false) assert!(!params.skip_special_tokens); @@ -964,13 +973,41 @@ mod tests { // Optional int fields should be None assert_eq!(params.max_new_tokens, None); assert_eq!(params.stream_interval, None); - // Other non-optional fields - assert_eq!(params.min_p, 0.0); - assert_eq!(params.frequency_penalty, 0.0); - assert_eq!(params.presence_penalty, 0.0); assert!(params.stop.is_empty()); } + #[test] + fn test_optional_sampling_params_set_vs_unset() { + // Verify that optional fields distinguish "set to 0" from "not set" + let unset = proto::SamplingParams::default(); + assert_eq!(unset.temperature, None); + + let set_to_zero = proto::SamplingParams { + temperature: Some(0.0), + ..Default::default() + }; + assert_eq!(set_to_zero.temperature, Some(0.0)); + + // Both serialize differently — the servicer uses HasField() to distinguish + let set_to_value = proto::SamplingParams { + temperature: Some(0.7), + top_p: Some(0.9), + top_k: Some(50), + min_p: Some(0.1), + frequency_penalty: Some(0.5), + presence_penalty: Some(0.3), + repetition_penalty: Some(1.2), + ..Default::default() + }; + assert_eq!(set_to_value.temperature, Some(0.7)); + assert_eq!(set_to_value.top_p, Some(0.9)); + assert_eq!(set_to_value.top_k, Some(50)); + assert_eq!(set_to_value.min_p, Some(0.1)); + assert_eq!(set_to_value.frequency_penalty, Some(0.5)); + assert_eq!(set_to_value.presence_penalty, Some(0.3)); + assert_eq!(set_to_value.repetition_penalty, Some(1.2)); + } + #[test] fn test_multimodal_inputs() { let mm_inputs = proto::MultimodalInputs { diff --git a/grpc_servicer/smg_grpc_servicer/sglang/server.py b/grpc_servicer/smg_grpc_servicer/sglang/server.py index 99230711bb..837e83ce56 100644 --- a/grpc_servicer/smg_grpc_servicer/sglang/server.py +++ b/grpc_servicer/smg_grpc_servicer/sglang/server.py @@ -133,6 +133,7 @@ async def serve_grpc( model_info=model_info, scheduler_info=scheduler_info, health_servicer=health_servicer, + model_config=model_config, ) sglang_scheduler_pb2_grpc.add_SglangSchedulerServicer_to_server(servicer, server) diff --git a/grpc_servicer/smg_grpc_servicer/sglang/servicer.py b/grpc_servicer/smg_grpc_servicer/sglang/servicer.py index aa2ee916e4..6252f794d8 100644 --- a/grpc_servicer/smg_grpc_servicer/sglang/servicer.py +++ b/grpc_servicer/smg_grpc_servicer/sglang/servicer.py @@ -162,6 +162,7 @@ def __init__( model_info: dict, scheduler_info: dict, health_servicer: SGLangHealthServicer | None = None, + model_config=None, ): """Initialize the standalone gRPC service.""" self.request_manager = request_manager @@ -170,6 +171,32 @@ def __init__( self.scheduler_info = scheduler_info self.start_time = time.time() self.health_servicer = health_servicer + + # Build merged sampling defaults for generation-config support. + # Priority: user-set (HasField) > preferred > generation_config > OpenAI defaults. + # Computed once here since these are immutable after init. + openai_defaults = { + "temperature": 1.0, + "top_p": 1.0, + "top_k": -1, + "min_p": 0.0, + "frequency_penalty": 0.0, + "presence_penalty": 0.0, + "repetition_penalty": 1.0, + } + default_sampling_params = ( + dict(model_config.get_default_sampling_params()) if model_config else {} + ) + preferred_sampling_params = ( + dict(server_args.preferred_sampling_params) + if server_args.preferred_sampling_params + else {} + ) + self._sampling_base = { + **openai_defaults, + **default_sampling_params, + **preferred_sampling_params, + } self.mm_receiver = None if ( self.server_args.language_only @@ -887,7 +914,35 @@ def _convert_embed_request( def _convert_sampling_params( self, grpc_params: sglang_scheduler_pb2.SamplingParams ) -> SGLSamplingParams: - """Convert gRPC SamplingParams to internal format.""" + """Convert gRPC SamplingParams to internal format. + + For the 7 optional sampling fields, applies a 3-tier default cascade: + user-set (HasField) > preferred_sampling_params > generation_config > OpenAI defaults + """ + base = self._sampling_base + + # Resolve each optional field: user-set value wins, else fall back to base + temperature = ( + grpc_params.temperature if grpc_params.HasField("temperature") else base["temperature"] + ) + top_p = grpc_params.top_p if grpc_params.HasField("top_p") else base["top_p"] + top_k = grpc_params.top_k if grpc_params.HasField("top_k") else base["top_k"] + min_p = grpc_params.min_p if grpc_params.HasField("min_p") else base["min_p"] + frequency_penalty = ( + grpc_params.frequency_penalty + if grpc_params.HasField("frequency_penalty") + else base["frequency_penalty"] + ) + presence_penalty = ( + grpc_params.presence_penalty + if grpc_params.HasField("presence_penalty") + else base["presence_penalty"] + ) + repetition_penalty = ( + grpc_params.repetition_penalty + if grpc_params.HasField("repetition_penalty") + else base["repetition_penalty"] + ) # Handle constraint types regex = None @@ -921,13 +976,13 @@ def _convert_sampling_params( stop_token_ids = list(grpc_params.stop_token_ids) if grpc_params.stop_token_ids else None return SGLSamplingParams( - temperature=grpc_params.temperature, - top_p=grpc_params.top_p, - top_k=grpc_params.top_k, - min_p=grpc_params.min_p, - frequency_penalty=grpc_params.frequency_penalty, - presence_penalty=grpc_params.presence_penalty, - repetition_penalty=grpc_params.repetition_penalty, + temperature=temperature, + top_p=top_p, + top_k=top_k, + min_p=min_p, + frequency_penalty=frequency_penalty, + presence_penalty=presence_penalty, + repetition_penalty=repetition_penalty, max_new_tokens=max_new_tokens, min_new_tokens=grpc_params.min_new_tokens, stop=stop,