Support multi-round conversations in bench_serving (#6135)
This commit is contained in:
@@ -0,0 +1,100 @@
|
||||
import json
|
||||
import tempfile
|
||||
import time
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from sglang.bench_serving import run_benchmark
|
||||
from sglang.srt.utils import kill_process_tree
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
from sglang.test.test_utils import (
|
||||
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
DEFAULT_URL_FOR_TEST,
|
||||
CustomTestCase,
|
||||
get_benchmark_args,
|
||||
popen_launch_server,
|
||||
)
|
||||
|
||||
register_cuda_ci(est_time=300, suite="nightly-1-gpu", nightly=True)
|
||||
|
||||
MODEL = "Qwen/Qwen3-0.6B"
|
||||
NUM_CONVERSATIONS, NUM_TURNS = 4, 3
|
||||
|
||||
|
||||
class TestBenchServingFunctionality(CustomTestCase):
|
||||
def test_gsp_multi_turn(self):
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
process = popen_launch_server(
|
||||
MODEL,
|
||||
DEFAULT_URL_FOR_TEST,
|
||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
other_args=[
|
||||
"--mem-fraction-static",
|
||||
"0.7",
|
||||
"--log-requests",
|
||||
"--log-requests-level",
|
||||
"3",
|
||||
"--log-requests-format",
|
||||
"json",
|
||||
"--log-requests-target",
|
||||
"stdout",
|
||||
temp_dir,
|
||||
],
|
||||
)
|
||||
try:
|
||||
args = get_benchmark_args(
|
||||
base_url=DEFAULT_URL_FOR_TEST,
|
||||
backend="sglang-oai-chat",
|
||||
tokenizer=MODEL,
|
||||
dataset_name="generated-shared-prefix",
|
||||
num_prompts=NUM_CONVERSATIONS,
|
||||
request_rate=float("inf"),
|
||||
gsp_num_groups=2,
|
||||
gsp_prompts_per_group=2,
|
||||
gsp_system_prompt_len=64,
|
||||
gsp_question_len=16,
|
||||
gsp_output_len=16,
|
||||
gsp_num_turns=NUM_TURNS,
|
||||
)
|
||||
args.warmup_requests = 0
|
||||
res = run_benchmark(args)
|
||||
self.assertEqual(res["completed"], NUM_CONVERSATIONS * NUM_TURNS)
|
||||
|
||||
time.sleep(1)
|
||||
logs = "".join(f.read_text() for f in Path(temp_dir).glob("*.log"))
|
||||
self._verify_multi_turn_logs(logs)
|
||||
finally:
|
||||
kill_process_tree(process.pid)
|
||||
|
||||
def _verify_multi_turn_logs(self, content: str):
|
||||
reqs = []
|
||||
for line in content.splitlines():
|
||||
if not line.startswith("{"):
|
||||
continue
|
||||
obj = json.loads(line)
|
||||
if obj.get("event") != "request.finished":
|
||||
continue
|
||||
text = obj.get("obj", {}).get("text")
|
||||
rid = obj.get("rid", "")
|
||||
if text and not rid.startswith("HEALTH_CHECK"):
|
||||
reqs.append(text)
|
||||
|
||||
self.assertGreaterEqual(len(reqs), NUM_CONVERSATIONS * NUM_TURNS)
|
||||
|
||||
# Verify prefix relationships
|
||||
reqs_sorted = sorted(reqs, key=len)
|
||||
prefix_count = 0
|
||||
for i, text in enumerate(reqs_sorted):
|
||||
for j in range(i + 1, len(reqs_sorted)):
|
||||
if reqs_sorted[j].startswith(text):
|
||||
prefix_count += 1
|
||||
break
|
||||
|
||||
expected = NUM_CONVERSATIONS * (NUM_TURNS - 1)
|
||||
self.assertGreaterEqual(
|
||||
prefix_count, expected, f"Expected at least {expected} prefix pairs"
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user