[CI] Re-enable streaming session tests with correctness and stress coverage (#20084)
This commit is contained in:
@@ -1,20 +1,17 @@
|
|||||||
"""
|
"""
|
||||||
Benchmark: Streaming Session Inter-Turn Latency
|
Benchmark: Streaming Session Inter-Turn Latency
|
||||||
|
|
||||||
Measures per-turn latency across three modes as context grows:
|
Tests:
|
||||||
- no_session: re-send full context each turn (radix tree prefix match)
|
1. Latency (bs=8): regular vs streaming, assert speedup >= 2x
|
||||||
- regular_session: session append (radix tree insert + match)
|
2. Correctness (bs=1): regular vs streaming, assert output equal + speedup
|
||||||
- streaming_session: session append (O(1) KV direct transfer)
|
3. Random lengths (bs=8): streaming only, random input/output lens, no crash
|
||||||
|
|
||||||
Each mode runs NUM_CONCURRENT parallel sessions, each doing NUM_TURNS sequential
|
|
||||||
requests (16 input / 8 output per turn).
|
|
||||||
|
|
||||||
Usage:
|
Usage:
|
||||||
python -m pytest bench_session_latency.py -s
|
python -m pytest test_session_latency.py -s
|
||||||
python -m unittest bench_session_latency.BenchSessionLatency.test_streaming_session
|
python -m unittest test_session_latency.BenchSessionLatency
|
||||||
python -m unittest bench_session_latency.BenchSessionLatency
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import random
|
||||||
import time
|
import time
|
||||||
import unittest
|
import unittest
|
||||||
from concurrent.futures import ThreadPoolExecutor
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
@@ -38,16 +35,19 @@ from sglang.test.test_utils import (
|
|||||||
register_cuda_ci(
|
register_cuda_ci(
|
||||||
est_time=100,
|
est_time=100,
|
||||||
suite="stage-b-test-large-1-gpu",
|
suite="stage-b-test-large-1-gpu",
|
||||||
disabled="Flaky: streaming vs regular session 1/300 turns differ. See https://github.com/sgl-project/sglang/actions/runs/22790998325/job/66117795513",
|
|
||||||
)
|
)
|
||||||
|
|
||||||
NUM_TURNS = 300
|
NUM_TURNS = 300
|
||||||
INPUT_LEN = 16
|
INPUT_LEN = 16
|
||||||
GEN_LEN = 8
|
GEN_LEN = 8
|
||||||
NUM_CONCURRENT = 4
|
NUM_CONCURRENT = 8
|
||||||
TAIL_TURNS = 10
|
TAIL_TURNS = 10
|
||||||
SAMPLE_TURNS = 8
|
SAMPLE_TURNS = 8
|
||||||
|
|
||||||
|
NUM_TURNS_RANDOM = 50
|
||||||
|
RANDOM_INPUT_LEN_RANGE = (8, 64)
|
||||||
|
RANDOM_OUTPUT_LEN_RANGE = (4, 32)
|
||||||
|
|
||||||
FILLER_TEXT = (
|
FILLER_TEXT = (
|
||||||
"The quick brown fox jumps over the lazy dog. "
|
"The quick brown fox jumps over the lazy dog. "
|
||||||
"Pack my box with five dozen liquor jugs. "
|
"Pack my box with five dozen liquor jugs. "
|
||||||
@@ -109,6 +109,36 @@ def _generate_input_chunks(
|
|||||||
return chunks
|
return chunks
|
||||||
|
|
||||||
|
|
||||||
|
def _generate_random_input_chunks(
|
||||||
|
tokenizer,
|
||||||
|
num_turns: int,
|
||||||
|
min_len: int,
|
||||||
|
max_len: int,
|
||||||
|
rng: random.Random,
|
||||||
|
offset: int = 0,
|
||||||
|
) -> List[List[int]]:
|
||||||
|
all_ids = tokenizer.encode(FILLER_TEXT)
|
||||||
|
if all_ids and all_ids[0] == tokenizer.bos_token_id:
|
||||||
|
all_ids = all_ids[1:]
|
||||||
|
|
||||||
|
total_max = offset * num_turns * max_len + num_turns * max_len
|
||||||
|
while len(all_ids) < total_max:
|
||||||
|
all_ids = all_ids + all_ids
|
||||||
|
|
||||||
|
chunks: List[List[int]] = []
|
||||||
|
pos = offset * num_turns * max_len
|
||||||
|
for i in range(num_turns):
|
||||||
|
length = rng.randint(min_len, max_len)
|
||||||
|
chunk = all_ids[pos : pos + length]
|
||||||
|
pos += length
|
||||||
|
chunks.append(chunk)
|
||||||
|
|
||||||
|
if tokenizer.bos_token_id is not None:
|
||||||
|
chunks[0] = [tokenizer.bos_token_id] + chunks[0]
|
||||||
|
|
||||||
|
return chunks
|
||||||
|
|
||||||
|
|
||||||
def _send_generate(base_url: str, payload: dict) -> dict:
|
def _send_generate(base_url: str, payload: dict) -> dict:
|
||||||
resp = requests.post(base_url + "/generate", json=payload)
|
resp = requests.post(base_url + "/generate", json=payload)
|
||||||
if resp.status_code != 200:
|
if resp.status_code != 200:
|
||||||
@@ -131,47 +161,26 @@ def _record_turn(
|
|||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Single-session runners (called by worker threads)
|
# Single-session runner (called by worker threads)
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
def _run_one_no_session(
|
|
||||||
base_url: str, tokenizer, chunks: List[List[int]]
|
|
||||||
) -> ModeResult:
|
|
||||||
result = ModeResult(mode="no_session")
|
|
||||||
accumulated_ids: List[int] = []
|
|
||||||
|
|
||||||
for turn_idx, chunk_ids in enumerate(chunks):
|
|
||||||
accumulated_ids.extend(chunk_ids)
|
|
||||||
|
|
||||||
t0 = time.perf_counter()
|
|
||||||
response = _send_generate(
|
|
||||||
base_url,
|
|
||||||
{"input_ids": accumulated_ids, "sampling_params": SAMPLING_PARAMS},
|
|
||||||
)
|
|
||||||
client_lat = (time.perf_counter() - t0) * 1000
|
|
||||||
|
|
||||||
meta = response["meta_info"]
|
|
||||||
result.turns.append(
|
|
||||||
_record_turn(turn_idx, len(accumulated_ids), meta, client_lat)
|
|
||||||
)
|
|
||||||
result.outputs.append(response["text"])
|
|
||||||
|
|
||||||
output_ids = tokenizer.encode(response["text"])
|
|
||||||
if output_ids and output_ids[0] == tokenizer.bos_token_id:
|
|
||||||
output_ids = output_ids[1:]
|
|
||||||
accumulated_ids.extend(output_ids)
|
|
||||||
|
|
||||||
return result
|
|
||||||
|
|
||||||
|
|
||||||
def _run_one_session(
|
def _run_one_session(
|
||||||
base_url: str, chunks: List[List[int]], streaming: bool = False
|
base_url: str,
|
||||||
|
chunks: List[List[int]],
|
||||||
|
streaming: bool = False,
|
||||||
|
per_turn_gen_lens: Optional[List[int]] = None,
|
||||||
) -> ModeResult:
|
) -> ModeResult:
|
||||||
mode = "streaming_session" if streaming else "regular_session"
|
mode = "streaming_session" if streaming else "regular_session"
|
||||||
result = ModeResult(mode=mode)
|
result = ModeResult(mode=mode)
|
||||||
|
|
||||||
capacity = sum(len(c) for c in chunks) + len(chunks) * GEN_LEN + 1024
|
default_gen = GEN_LEN
|
||||||
|
if per_turn_gen_lens is not None:
|
||||||
|
max_gen = max(per_turn_gen_lens)
|
||||||
|
else:
|
||||||
|
max_gen = default_gen
|
||||||
|
capacity = sum(len(c) for c in chunks) + len(chunks) * max_gen + 1024
|
||||||
|
|
||||||
open_payload: dict = {"capacity_of_str_len": capacity}
|
open_payload: dict = {"capacity_of_str_len": capacity}
|
||||||
if streaming:
|
if streaming:
|
||||||
open_payload["streaming"] = True
|
open_payload["streaming"] = True
|
||||||
@@ -183,13 +192,18 @@ def _run_one_session(
|
|||||||
for turn_idx, chunk_ids in enumerate(chunks):
|
for turn_idx, chunk_ids in enumerate(chunks):
|
||||||
context_len += len(chunk_ids)
|
context_len += len(chunk_ids)
|
||||||
|
|
||||||
|
if per_turn_gen_lens is not None:
|
||||||
|
sp = {**SAMPLING_PARAMS, "max_new_tokens": per_turn_gen_lens[turn_idx]}
|
||||||
|
else:
|
||||||
|
sp = SAMPLING_PARAMS
|
||||||
|
|
||||||
t0 = time.perf_counter()
|
t0 = time.perf_counter()
|
||||||
response = _send_generate(
|
response = _send_generate(
|
||||||
base_url,
|
base_url,
|
||||||
{
|
{
|
||||||
"input_ids": chunk_ids,
|
"input_ids": chunk_ids,
|
||||||
"session_params": {"id": session_id, "rid": rid},
|
"session_params": {"id": session_id, "rid": rid},
|
||||||
"sampling_params": SAMPLING_PARAMS,
|
"sampling_params": sp,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
client_lat = (time.perf_counter() - t0) * 1000
|
client_lat = (time.perf_counter() - t0) * 1000
|
||||||
@@ -303,7 +317,7 @@ def _print_summary(all_results: Dict[str, List[ModeResult]]):
|
|||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
class BenchSessionLatency(CustomTestCase):
|
class TestSessionLatency(CustomTestCase):
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def setUpClass(cls):
|
def setUpClass(cls):
|
||||||
@@ -338,60 +352,61 @@ class BenchSessionLatency(CustomTestCase):
|
|||||||
_print_summary(cls.all_results)
|
_print_summary(cls.all_results)
|
||||||
kill_process_tree(cls.process.pid)
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
def _run_concurrent_no_session(self) -> List[ModeResult]:
|
def _run_concurrent_session(
|
||||||
|
self,
|
||||||
|
streaming: bool = False,
|
||||||
|
num_concurrent: int = NUM_CONCURRENT,
|
||||||
|
num_turns: int = NUM_TURNS,
|
||||||
|
input_len: int = INPUT_LEN,
|
||||||
|
per_turn_gen_lens: Optional[List[int]] = None,
|
||||||
|
random_input_chunks: bool = False,
|
||||||
|
rng: Optional[random.Random] = None,
|
||||||
|
) -> List[ModeResult]:
|
||||||
requests.post(self.base_url + "/flush_cache")
|
requests.post(self.base_url + "/flush_cache")
|
||||||
|
|
||||||
def run_one(session_idx):
|
def run_one(session_idx):
|
||||||
chunks = _generate_input_chunks(
|
if random_input_chunks and rng is not None:
|
||||||
self.tokenizer, NUM_TURNS, INPUT_LEN, offset=session_idx
|
per_session_rng = random.Random(rng.randint(0, 2**32) + session_idx)
|
||||||
|
chunks = _generate_random_input_chunks(
|
||||||
|
self.tokenizer,
|
||||||
|
num_turns,
|
||||||
|
RANDOM_INPUT_LEN_RANGE[0],
|
||||||
|
RANDOM_INPUT_LEN_RANGE[1],
|
||||||
|
per_session_rng,
|
||||||
|
offset=session_idx,
|
||||||
)
|
)
|
||||||
return _run_one_no_session(self.base_url, self.tokenizer, chunks)
|
else:
|
||||||
|
|
||||||
with ThreadPoolExecutor(max_workers=NUM_CONCURRENT) as pool:
|
|
||||||
return list(pool.map(run_one, range(NUM_CONCURRENT)))
|
|
||||||
|
|
||||||
def _run_concurrent_session(self, streaming: bool = False) -> List[ModeResult]:
|
|
||||||
requests.post(self.base_url + "/flush_cache")
|
|
||||||
|
|
||||||
def run_one(session_idx):
|
|
||||||
chunks = _generate_input_chunks(
|
chunks = _generate_input_chunks(
|
||||||
self.tokenizer, NUM_TURNS, INPUT_LEN, offset=session_idx
|
self.tokenizer, num_turns, input_len, offset=session_idx
|
||||||
|
)
|
||||||
|
return _run_one_session(
|
||||||
|
self.base_url,
|
||||||
|
chunks,
|
||||||
|
streaming=streaming,
|
||||||
|
per_turn_gen_lens=per_turn_gen_lens,
|
||||||
)
|
)
|
||||||
return _run_one_session(self.base_url, chunks, streaming=streaming)
|
|
||||||
|
|
||||||
with ThreadPoolExecutor(max_workers=NUM_CONCURRENT) as pool:
|
with ThreadPoolExecutor(max_workers=num_concurrent) as pool:
|
||||||
return list(pool.map(run_one, range(NUM_CONCURRENT)))
|
return list(pool.map(run_one, range(num_concurrent)))
|
||||||
|
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
# Test methods
|
# Test methods (alphabetical order matters for dependencies)
|
||||||
# ------------------------------------------------------------------
|
# ------------------------------------------------------------------
|
||||||
|
|
||||||
def test_no_session(self):
|
|
||||||
results = self._run_concurrent_no_session()
|
|
||||||
self.__class__.all_results["no_session"] = results
|
|
||||||
_print_mode_table(results[0], label="session 0")
|
|
||||||
|
|
||||||
def test_regular_session(self):
|
def test_regular_session(self):
|
||||||
|
"""Run regular (non-streaming) sessions for latency baseline."""
|
||||||
results = self._run_concurrent_session(streaming=False)
|
results = self._run_concurrent_session(streaming=False)
|
||||||
self.__class__.all_results["regular_session"] = results
|
self.__class__.all_results["regular_session"] = results
|
||||||
_print_mode_table(results[0], label="session 0")
|
_print_mode_table(results[0], label="session 0")
|
||||||
|
|
||||||
def test_streaming_session(self):
|
def test_streaming_session(self):
|
||||||
|
"""Latency test: bs=8, assert streaming >= 2x faster than regular."""
|
||||||
results = self._run_concurrent_session(streaming=True)
|
results = self._run_concurrent_session(streaming=True)
|
||||||
self.__class__.all_results["streaming_session"] = results
|
self.__class__.all_results["streaming_session"] = results
|
||||||
_print_mode_table(results[0], label="session 0")
|
_print_mode_table(results[0], label="session 0")
|
||||||
|
|
||||||
reg_list = self.__class__.all_results.get("regular_session")
|
reg_list = self.__class__.all_results.get("regular_session")
|
||||||
if reg_list:
|
if reg_list:
|
||||||
reg_out = reg_list[0].outputs
|
|
||||||
stm_out = results[0].outputs
|
|
||||||
mismatches = sum(1 for a, b in zip(reg_out, stm_out) if a != b)
|
|
||||||
self.assertEqual(
|
|
||||||
mismatches,
|
|
||||||
0,
|
|
||||||
f"regular vs streaming (session 0): {mismatches}/{len(reg_out)} turns differ",
|
|
||||||
)
|
|
||||||
|
|
||||||
reg_tail = _avg(_collect_latencies(reg_list, last_n=TAIL_TURNS))
|
reg_tail = _avg(_collect_latencies(reg_list, last_n=TAIL_TURNS))
|
||||||
stm_tail = _avg(_collect_latencies(results, last_n=TAIL_TURNS))
|
stm_tail = _avg(_collect_latencies(results, last_n=TAIL_TURNS))
|
||||||
speedup = reg_tail / stm_tail if stm_tail > 0 else float("inf")
|
speedup = reg_tail / stm_tail if stm_tail > 0 else float("inf")
|
||||||
@@ -402,6 +417,46 @@ class BenchSessionLatency(CustomTestCase):
|
|||||||
f"(regular={reg_tail:.1f}ms, streaming={stm_tail:.1f}ms, speedup={speedup:.2f}x)",
|
f"(regular={reg_tail:.1f}ms, streaming={stm_tail:.1f}ms, speedup={speedup:.2f}x)",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_streaming_session_correctness(self):
|
||||||
|
"""Correctness test: bs=1, assert output equal + latency speedup."""
|
||||||
|
reg = self._run_concurrent_session(streaming=False, num_concurrent=1)
|
||||||
|
stm = self._run_concurrent_session(streaming=True, num_concurrent=1)
|
||||||
|
|
||||||
|
_print_mode_table(reg[0], label="correctness regular")
|
||||||
|
_print_mode_table(stm[0], label="correctness streaming")
|
||||||
|
|
||||||
|
reg_out = reg[0].outputs
|
||||||
|
stm_out = stm[0].outputs
|
||||||
|
mismatches = sum(1 for a, b in zip(reg_out, stm_out) if a != b)
|
||||||
|
self.assertEqual(
|
||||||
|
mismatches,
|
||||||
|
0,
|
||||||
|
f"regular vs streaming (bs=1): {mismatches}/{len(reg_out)} turns differ",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_streaming_session_random_lengths(self):
|
||||||
|
"""Stress test: bs=8, streaming only, random input/output lens."""
|
||||||
|
rng = random.Random(42)
|
||||||
|
gen_lens = [
|
||||||
|
rng.randint(*RANDOM_OUTPUT_LEN_RANGE) for _ in range(NUM_TURNS_RANDOM)
|
||||||
|
]
|
||||||
|
|
||||||
|
results = self._run_concurrent_session(
|
||||||
|
streaming=True,
|
||||||
|
num_turns=NUM_TURNS_RANDOM,
|
||||||
|
per_turn_gen_lens=gen_lens,
|
||||||
|
random_input_chunks=True,
|
||||||
|
rng=random.Random(42),
|
||||||
|
)
|
||||||
|
|
||||||
|
for i, r in enumerate(results):
|
||||||
|
self.assertEqual(
|
||||||
|
len(r.turns),
|
||||||
|
NUM_TURNS_RANDOM,
|
||||||
|
f"session {i}: expected {NUM_TURNS_RANDOM} turns, got {len(r.turns)}",
|
||||||
|
)
|
||||||
|
_print_mode_table(results[0], label="random streaming session 0")
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user