support rust sglang server (#29799)
This commit is contained in:
@@ -2,6 +2,8 @@
|
||||
python3 -m unittest test_srt_endpoint.TestSRTEndpoint.test_simple_decode
|
||||
python3 -m unittest test_srt_endpoint.TestSRTEndpoint.test_logprob_with_chunked_prefill
|
||||
python3 -m unittest test_srt_endpoint.TestTokenizeDetokenize
|
||||
python3 -m unittest test_srt_endpoint.TestRustServerEndpoint
|
||||
python3 -m unittest test_srt_endpoint.TestRustServerLogprob
|
||||
"""
|
||||
|
||||
import json
|
||||
@@ -24,17 +26,22 @@ from sglang.test.test_utils import (
|
||||
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
DEFAULT_URL_FOR_TEST,
|
||||
CustomTestCase,
|
||||
is_rust_server_built,
|
||||
popen_launch_server,
|
||||
run_logprob_check,
|
||||
)
|
||||
|
||||
register_cuda_ci(est_time=160, stage="base-b", runner_config="1-gpu-small")
|
||||
register_amd_ci(est_time=160, suite="stage-b-test-1-gpu-small-amd")
|
||||
register_cuda_ci(est_time=260, stage="base-b", runner_config="1-gpu-small")
|
||||
register_amd_ci(est_time=260, suite="stage-b-test-1-gpu-small-amd")
|
||||
|
||||
SERVER_ENV = {"SGLANG_USE_PICKLE_IPC": "0"}
|
||||
|
||||
|
||||
class TestSRTEndpoint(CustomTestCase):
|
||||
# Extra server-launch env; subclasses override to run the same suite
|
||||
# against a different server flavor (e.g. SGLANG_RUST_SERVER=1).
|
||||
env = {}
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
|
||||
@@ -46,7 +53,7 @@ class TestSRTEndpoint(CustomTestCase):
|
||||
# The tiny logprob chunk size routes this file's logprob tests
|
||||
# through the multi-chunk stitching path (requests at or below 64
|
||||
# rows still cover the non-chunked path).
|
||||
env={**SERVER_ENV, "SGLANG_LOGPROB_CHUNK_SIZE": "64"},
|
||||
env={**cls.env, **SERVER_ENV, "SGLANG_LOGPROB_CHUNK_SIZE": "64"},
|
||||
other_args=(
|
||||
"--enable-custom-logit-processor",
|
||||
"--mem-fraction-static",
|
||||
@@ -844,5 +851,74 @@ class TestTokenizeDetokenize(CustomTestCase):
|
||||
self.assertEqual(r2.status_code, 500)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Embedded Rust server (SGLANG_RUST_SERVER=1): rerun the whole endpoint suite
|
||||
# against the rust api-server/tokenizer/detokenizer stack — the logprob tests
|
||||
# exercise the columnar egress wire (`push_generation` extras -> Rust
|
||||
# `BatchHeader`/`for_each_chunk` -> detok reshape) end to end. Suite
|
||||
# surface the rust server does not implement yet is skipped explicitly below.
|
||||
# ---------------------------------------------------------------------------
|
||||
@unittest.skipUnless(
|
||||
is_rust_server_built(),
|
||||
"embedded rust server extension not built (e.g. AMD suite)",
|
||||
)
|
||||
class TestRustServerEndpoint(TestSRTEndpoint):
|
||||
env = {"SGLANG_RUST_SERVER": "1"}
|
||||
|
||||
_RUST_TODO = "not implemented by the embedded Rust server yet"
|
||||
|
||||
@unittest.skip(f"custom_logit_processor request field {_RUST_TODO}")
|
||||
def test_custom_logit_processor(self):
|
||||
pass
|
||||
|
||||
@unittest.skip(f"custom_logit_processor request field {_RUST_TODO}")
|
||||
def test_custom_logit_processor_batch_mixed(self):
|
||||
pass
|
||||
|
||||
@unittest.skip(f"custom_logit_processor request field {_RUST_TODO}")
|
||||
def test_stateful_custom_logit_processor(self):
|
||||
pass
|
||||
|
||||
@unittest.skip(f"custom_logit_processor request field {_RUST_TODO}")
|
||||
def test_stateful_custom_logit_processor_batch_mixed(self):
|
||||
pass
|
||||
|
||||
@unittest.skip(f"/flush_cache endpoint + cached_tokens meta {_RUST_TODO}")
|
||||
def test_cache_tokens(self):
|
||||
pass
|
||||
|
||||
def test_greedy_token_equals_top1(self):
|
||||
"""Cross-column alignment guard for the columnar logprob wire: at
|
||||
temperature 0 the chosen token must BE the top-1 entry of its own
|
||||
position. A column shifted across requests or positions (the failure
|
||||
mode a truncation-tolerant reader would mask) breaks this instantly."""
|
||||
response = requests.post(
|
||||
self.base_url + "/generate",
|
||||
json={
|
||||
"text": ["The capital of France is", "I have a very good idea on"],
|
||||
"sampling_params": {"temperature": 0, "max_new_tokens": 8},
|
||||
"return_logprob": True,
|
||||
"top_logprobs_num": 5,
|
||||
"logprob_start_len": 0,
|
||||
},
|
||||
)
|
||||
self.assertEqual(response.status_code, 200, response.text)
|
||||
for res in response.json():
|
||||
meta = res["meta_info"]
|
||||
out_lp = meta["output_token_logprobs"]
|
||||
top = meta["output_top_logprobs"]
|
||||
self.assertEqual(len(out_lp), meta["completion_tokens"])
|
||||
self.assertEqual(len(top), len(out_lp))
|
||||
# First prompt token's logprob is the None sentinel; it must
|
||||
# survive the NaN wire encoding and come back as null.
|
||||
self.assertIsNone(meta["input_token_logprobs"][0][0])
|
||||
for (lp, tid, _), pos_top in zip(out_lp, top):
|
||||
self.assertEqual(len(pos_top), 5)
|
||||
self.assertEqual(pos_top[0][1], tid)
|
||||
self.assertAlmostEqual(pos_top[0][0], lp, places=4)
|
||||
vals = [t[0] for t in pos_top]
|
||||
self.assertEqual(vals, sorted(vals, reverse=True))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user