[Rust] Keep sampling and scheduler wire schemas in sync (#37222)

This commit is contained in:
Lianmin Zheng
2026-08-31 12:31:07 -07:00
committed by GitHub
parent 48098b5f23
commit 2138494272
3 changed files with 144 additions and 19 deletions
+93 -13
View File
@@ -28,6 +28,32 @@ const MAX_STOP_REGEX_LEN: usize = 256;
/// (`re._MAXCACHE`), so past that every pattern recompiles on every decode step. /// (`re._MAXCACHE`), so past that every pattern recompiles on every decode step.
const MAX_STOP_REGEX_COUNT: usize = 32; const MAX_STOP_REGEX_COUNT: usize = 32;
/// JSON values accepted by Python's `CustomParamValue`: a scalar, a list of
/// scalars, or a string-keyed object whose values are scalars.
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum CustomParamValue {
Null(()),
Bool(bool),
Signed(i64),
Unsigned(u64),
Float(f64),
String(String),
List(Vec<JsonScalar>),
Object(BTreeMap<String, JsonScalar>),
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum JsonScalar {
Null(()),
Bool(bool),
Signed(i64),
Unsigned(u64),
Float(f64),
String(String),
}
/// One module per field default, each exposing the two hooks serde needs under /// One module per field default, each exposing the two hooks serde needs under
/// one name: `default` (key absent) and `deserialize` (key present — including /// one name: `default` (key absent) and `deserialize` (key present — including
/// an explicit `null`, which Python's `__post_init__` maps back to the default: /// an explicit `null`, which Python's `__post_init__` maps back to the default:
@@ -180,10 +206,10 @@ pub struct SamplingParams {
pub logit_bias: Option<BTreeMap<String, f64>>, pub logit_bias: Option<BTreeMap<String, f64>>,
#[serde(default)] #[serde(default)]
pub sampling_seed: Option<i64>, pub sampling_seed: Option<i64>,
/// Opaque JSON object forwarded to a custom logit processor. Python types it /// JSON object forwarded to a custom logit processor. Its values match
/// as `Dict[str, JsonScalar | list | dict]`; it is never inspected here. /// Python's `CustomParamValue` exactly.
#[serde(default)] #[serde(default)]
pub custom_params: Option<serde_json::Value>, pub custom_params: Option<BTreeMap<String, CustomParamValue>>,
// Normalized internal fields. // Normalized internal fields.
// //
@@ -329,6 +355,16 @@ impl SamplingParams {
if self.top_k == -1 { if self.top_k == -1 {
self.top_k = TOP_K_ALL; // -1 disables top_k → whole vocabulary self.top_k = TOP_K_ALL; // -1 disables top_k → whole vocabulary
} }
for constraint in [
&mut self.json_schema,
&mut self.regex,
&mut self.ebnf,
&mut self.structural_tag,
] {
if constraint.as_deref() == Some("") {
*constraint = None;
}
}
} }
/// Python `normalize(tokenizer)`: size the stop match windows, reject /// Python `normalize(tokenizer)`: size the stop match windows, reject
@@ -402,6 +438,13 @@ impl SamplingParams {
/// Python `verify(vocab_size)` — the same ranges, messages and mutual /// Python `verify(vocab_size)` — the same ranges, messages and mutual
/// exclusions, plus the rust-server `n == 1` restriction. /// exclusions, plus the rust-server `n == 1` restriction.
fn verify(&self, vocab_size: u64) -> Result<(), Error> { fn verify(&self, vocab_size: u64) -> Result<(), Error> {
if let Some(beam_width) = self.beam_width
&& beam_width < 1
{
return Err(bad(format!(
"beam_width must be at least 1, got {beam_width}."
)));
}
if !self.temperature.is_finite() || self.temperature < 0.0 { if !self.temperature.is_finite() || self.temperature < 0.0 {
return Err(bad(format!( return Err(bad(format!(
"temperature must be a non-negative finite number, got {}", "temperature must be a non-negative finite number, got {}",
@@ -477,13 +520,18 @@ impl SamplingParams {
} }
} }
// Grammars are mutually exclusive. // Grammars are mutually exclusive.
let grammars = [&self.json_schema, &self.regex, &self.ebnf] let grammars = [
.iter() &self.json_schema,
.filter(|g| g.is_some()) &self.regex,
.count(); &self.ebnf,
&self.structural_tag,
]
.iter()
.filter(|g| g.is_some())
.count();
if grammars > 1 { if grammars > 1 {
return Err(bad( return Err(bad(
"Only one of regex, json_schema, or ebnf can be set".into() "Only one of json_schema, regex, ebnf, or structural_tag can be set".into(),
)); ));
} }
// Not a Python restriction: the rust from_scheduler maps one rid to one response, // Not a Python restriction: the rust from_scheduler maps one rid to one response,
@@ -497,11 +545,6 @@ impl SamplingParams {
))); )));
} }
if let Some(beam_width) = self.beam_width { if let Some(beam_width) = self.beam_width {
if beam_width < 1 {
return Err(bad(format!(
"beam_width must be at least 1, got {beam_width}."
)));
}
// Also not a Python restriction: beam search returns its candidates // Also not a Python restriction: beam search returns its candidates
// in `meta_info.beam_results`, which from_scheduler does not carry. // in `meta_info.beam_results`, which from_scheduler does not carry.
if beam_width > 1 { if beam_width > 1 {
@@ -912,6 +955,43 @@ mod tests {
assert_eq!(once.top_k, twice.top_k); assert_eq!(once.top_k, twice.top_k);
} }
#[test]
fn empty_grammar_constraints_are_unset() {
let sp = norm(r#"{"json_schema":"","regex":"","ebnf":"","structural_tag":""}"#);
assert!(sp.json_schema.is_none());
assert!(sp.regex.is_none());
assert!(sp.ebnf.is_none());
assert!(sp.structural_tag.is_none());
}
#[test]
fn structural_tag_is_mutually_exclusive_with_other_grammars() {
for field in ["json_schema", "regex", "ebnf"] {
let json = format!(r#"{{"structural_tag":"tag","{field}":"other"}}"#);
assert!(norm_err(&json).to_string().contains("Only one of"));
}
}
#[test]
fn custom_params_matches_python_shape() {
assert!(
norm(r#"{"custom_params":{"null":null,"bool":true,"int":1,"float":1.5,"str":"x","list":[1,"x",null],"object":{"x":1}}}"#)
.custom_params
.is_some()
);
for json in [
r#"{"custom_params":[]}"#,
r#"{"custom_params":{"nested_list":[[1]]}}"#,
r#"{"custom_params":{"nested_object":{"x":{"y":1}}}}"#,
] {
assert!(
serde_json::from_str::<SamplingParams>(json).is_err(),
"{json} must not parse"
);
}
}
/// `skip_tokenizer_init` has no tokenizer, so the text-matching stop features /// `skip_tokenizer_init` has no tokenizer, so the text-matching stop features
/// and `min_new_tokens` (needs eos_token_id) are 400s, not silent no-ops. /// and `min_new_tokens` (needs eos_token_id) are 400s, not silent no-ops.
/// Mirrors Python `raise_if_tokenizer_required`. /// Mirrors Python `raise_if_tokenizer_required`.
@@ -1,7 +1,9 @@
import copy import copy
import re
import unittest import unittest
import weakref import weakref
from array import array from array import array
from pathlib import Path
import msgspec import msgspec
import numpy as np import numpy as np
@@ -41,6 +43,33 @@ register_cpu_ci(est_time=6, suite="base-c-test-cpu")
class TestTokenizedReqInputMsgpack(unittest.TestCase): class TestTokenizedReqInputMsgpack(unittest.TestCase):
def test_rust_tokenized_generate_schema_stays_in_lockstep(self):
"""Compare the Rust wire declaration with the imported Python schema."""
rust_path = (
Path(__file__).resolve().parents[4]
/ "rust/sglang-server/src/message/io_struct.rs"
)
source = rust_path.read_text()
start = source.index("pub(super) TokenizedGenerateReqInput<'a> {")
end = source.index("\n }\n}", start)
rust_fields = (
"rid",
"http_worker_ipc",
*re.findall(r"^\s*([a-z][a-z0-9_]*):", source[start:end], re.MULTILINE),
)
python_fields = TokenizedGenerateReqInput.__struct_fields__
self.assertEqual(python_fields[: len(rust_fields)], rust_fields)
self.assertTrue(
all(
default is not msgspec.NODEFAULT
for default in TokenizedGenerateReqInput.__struct_defaults__[
len(rust_fields) :
]
),
"Rust may omit only a defaulted suffix of the Python wire schema",
)
def _make_mm_inputs(self, device="cpu"): def _make_mm_inputs(self, device="cpu"):
return MultimodalProcessorOutput( return MultimodalProcessorOutput(
mm_items=[ mm_items=[
@@ -7,7 +7,9 @@ register_cpu_ci(est_time=8, suite="base-c-test-cpu")
register_xpu_ci(est_time=10, suite="stage-a-test-1-gpu-xpu") register_xpu_ci(est_time=10, suite="stage-a-test-1-gpu-xpu")
import copy import copy
import re
import unittest import unittest
from pathlib import Path
from unittest.mock import MagicMock from unittest.mock import MagicMock
import msgspec import msgspec
@@ -25,7 +27,6 @@ from sglang.test.test_utils import CustomTestCase
class TestSamplingParamsInit(CustomTestCase): class TestSamplingParamsInit(CustomTestCase):
def test_zero_temperature_becomes_greedy(self): def test_zero_temperature_becomes_greedy(self):
"""Test greedy conversion when temperature is 0.""" """Test greedy conversion when temperature is 0."""
sp = SamplingParams(temperature=0.0) sp = SamplingParams(temperature=0.0)
@@ -88,7 +89,6 @@ class TestSamplingParamsInit(CustomTestCase):
class TestSamplingParamsVerify(CustomTestCase): class TestSamplingParamsVerify(CustomTestCase):
VOCAB_SIZE = 32000 VOCAB_SIZE = 32000
GRAMMAR_VALUES = { GRAMMAR_VALUES = {
"json_schema": '{"type":"object"}', "json_schema": '{"type":"object"}',
@@ -314,13 +314,12 @@ class TestSamplingParamsVerify(CustomTestCase):
class TestSamplingParamsNormalize(CustomTestCase): class TestSamplingParamsNormalize(CustomTestCase):
def _mock_tokenizer(self, encode_map=None): def _mock_tokenizer(self, encode_map=None):
"""Create a mock tokenizer that returns predetermined token lists.""" """Create a mock tokenizer that returns predetermined token lists."""
tokenizer = MagicMock() tokenizer = MagicMock()
if encode_map: if encode_map:
tokenizer.encode.side_effect = ( tokenizer.encode.side_effect = lambda s, add_special_tokens=False: (
lambda s, add_special_tokens=False: encode_map.get(s, [1]) encode_map.get(s, [1])
) )
else: else:
tokenizer.encode.return_value = [1] # Default: 1 token tokenizer.encode.return_value = [1] # Default: 1 token
@@ -428,6 +427,24 @@ class TestSamplingParamsNormalize(CustomTestCase):
class TestSamplingParamsMsgspecStruct(CustomTestCase): class TestSamplingParamsMsgspecStruct(CustomTestCase):
def test_rust_sampling_schema_stays_in_lockstep(self):
"""Compare Rust fields with the imported Python wire schema."""
rust_path = (
Path(__file__).resolve().parents[4]
/ "rust/sglang-server/src/message/sampling.rs"
)
source = rust_path.read_text()
start = source.index("pub struct SamplingParams {")
end = source.index("\n}\n\n/// The `/generate`", start)
rust_fields = tuple(
re.findall(
r"^\s*pub ([a-z][a-z0-9_]*):",
source[start:end],
re.MULTILINE,
)
)
self.assertEqual(SamplingParams.__struct_fields__, rust_fields)
def test_copy_remains_mutable_and_independent(self): def test_copy_remains_mutable_and_independent(self):
sp = SamplingParams(max_new_tokens=8, custom_params={"a": 1}) sp = SamplingParams(max_new_tokens=8, custom_params={"a": 1})
@@ -501,7 +518,6 @@ class TestSamplingParamsMsgspecStruct(CustomTestCase):
class TestRegexMaxLength(CustomTestCase): class TestRegexMaxLength(CustomTestCase):
def test_literal_string(self): def test_literal_string(self):
"""Test that plain string 'abc' gives max length 3.""" """Test that plain string 'abc' gives max length 3."""
self.assertEqual(get_max_seq_length("abc"), 3) self.assertEqual(get_max_seq_length("abc"), 3)