[Test] test_session_latency: assert streaming tail/head stability (#26230)

This commit is contained in:
Liangsheng Yin
2026-05-24 14:06:41 -07:00
committed by GitHub
parent fd94bd30b8
commit 030bd5d3ed
@@ -2,13 +2,9 @@
Benchmark: Streaming Session Inter-Turn Latency Benchmark: Streaming Session Inter-Turn Latency
Tests: Tests:
1. Latency (bs=8): regular vs streaming, assert speedup >= 2x 1. Stability (bs=8): streaming only, assert tail_avg / head_avg <= 1.15
2. Correctness (bs=1): regular vs streaming, assert output equal + speedup 2. Correctness (bs=1): regular vs streaming, assert output equal
3. Random lengths (bs=8): streaming only, random input/output lens, no crash 3. Random lengths (bs=8): streaming only, random input/output lens, no crash
Usage:
python -m pytest test_session_latency.py -s
python -m unittest test_session_latency.BenchSessionLatency
""" """
import random import random
@@ -16,7 +12,7 @@ import time
import unittest import unittest
from concurrent.futures import ThreadPoolExecutor from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Dict, List, Optional from typing import List, Optional
import requests import requests
from tabulate import tabulate from tabulate import tabulate
@@ -37,6 +33,7 @@ NUM_TURNS = 150
INPUT_LEN = 16 INPUT_LEN = 16
GEN_LEN = 8 GEN_LEN = 8
NUM_CONCURRENT = 8 NUM_CONCURRENT = 8
HEAD_TURNS = 10
TAIL_TURNS = 10 TAIL_TURNS = 10
SAMPLE_TURNS = 8 SAMPLE_TURNS = 8
@@ -65,8 +62,6 @@ class TurnResult:
turn: int turn: int
context_len: int context_len: int
cached_tokens: int cached_tokens: int
prompt_tokens: int
completion_tokens: int
client_latency_ms: float client_latency_ms: float
e2e_latency_ms: float e2e_latency_ms: float
@@ -78,11 +73,6 @@ class ModeResult:
outputs: List[str] = field(default_factory=list) outputs: List[str] = field(default_factory=list)
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _generate_input_chunks( def _generate_input_chunks(
tokenizer, num_turns: int, input_len: int, offset: int = 0 tokenizer, num_turns: int, input_len: int, offset: int = 0
) -> List[List[int]]: ) -> List[List[int]]:
@@ -149,18 +139,11 @@ def _record_turn(
turn=turn_idx + 1, turn=turn_idx + 1,
context_len=context_len, context_len=context_len,
cached_tokens=meta["cached_tokens"], cached_tokens=meta["cached_tokens"],
prompt_tokens=meta["prompt_tokens"],
completion_tokens=meta["completion_tokens"],
client_latency_ms=client_latency_ms, client_latency_ms=client_latency_ms,
e2e_latency_ms=meta.get("e2e_latency", 0) * 1000, e2e_latency_ms=meta.get("e2e_latency", 0) * 1000,
) )
# ---------------------------------------------------------------------------
# Single-session runner (called by worker threads)
# ---------------------------------------------------------------------------
def _run_one_session( def _run_one_session(
base_url: str, base_url: str,
chunks: List[List[int]], chunks: List[List[int]],
@@ -215,19 +198,20 @@ def _run_one_session(
return result return result
# ---------------------------------------------------------------------------
# Stats & reporting
# ---------------------------------------------------------------------------
def _collect_latencies( def _collect_latencies(
results: List[ModeResult], last_n: Optional[int] = None results: List[ModeResult],
last_n: Optional[int] = None,
first_n: Optional[int] = None,
) -> List[float]: ) -> List[float]:
lats = [] lats = []
for r in results: for r in results:
turns = r.turns[1:] # skip turn 1
if last_n is not None: if last_n is not None:
turns = r.turns[-last_n:] turns = r.turns[-last_n:]
elif first_n is not None:
# Skip turn 1 (includes prefill), then take next `first_n` turns.
turns = r.turns[1 : 1 + first_n]
else:
turns = r.turns[1:] # skip turn 1
lats.extend(t.client_latency_ms for t in turns) lats.extend(t.client_latency_ms for t in turns)
return lats return lats
@@ -270,49 +254,6 @@ def _print_mode_table(result: ModeResult, label: str = ""):
) )
def _print_summary(all_results: Dict[str, List[ModeResult]]):
stats = [
(
mode,
_avg(_collect_latencies(rs)),
_avg(_collect_latencies(rs, last_n=TAIL_TURNS)),
)
for mode, rs in all_results.items()
]
base_all, base_tail = (stats[0][1] or 1.0), (stats[0][2] or 1.0)
tail_label = f"last {TAIL_TURNS}"
print(f"\n SUMMARY ({NUM_CONCURRENT} sessions x {NUM_TURNS} turns)")
rows = [
[
mode,
f"{a:.1f}ms",
f"{t:.1f}ms",
f"{base_all / a:.2f}x" if a else "inf",
f"{base_tail / t:.2f}x" if t else "inf",
]
for mode, a, t in stats
]
print(
tabulate(
rows,
headers=[
"Mode",
"Avg (all)",
f"Avg ({tail_label})",
"Speedup (all)",
f"Speedup ({tail_label})",
],
colalign=("left", "right", "right", "right", "right"),
)
)
# ---------------------------------------------------------------------------
# Test class
# ---------------------------------------------------------------------------
class TestSessionLatency(CustomTestCase): class TestSessionLatency(CustomTestCase):
@classmethod @classmethod
def setUpClass(cls): def setUpClass(cls):
@@ -346,12 +287,8 @@ class TestSessionLatency(CustomTestCase):
}, },
) )
cls.all_results: Dict[str, List[ModeResult]] = {}
@classmethod @classmethod
def tearDownClass(cls): def tearDownClass(cls):
if len(cls.all_results) > 1:
_print_summary(cls.all_results)
kill_process_tree(cls.process.pid) kill_process_tree(cls.process.pid)
def _run_concurrent_session( def _run_concurrent_session(
@@ -391,36 +328,30 @@ class TestSessionLatency(CustomTestCase):
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 (alphabetical order matters for dependencies)
# ------------------------------------------------------------------
def test_regular_session(self):
"""Run regular (non-streaming) sessions for latency baseline."""
results = self._run_concurrent_session(streaming=False)
self.__class__.all_results["regular_session"] = results
_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.""" """Stability: streaming reuses KV across turns, so tail/head latency
should stay flat. Skip turn 1 (prefill) when computing head."""
results = self._run_concurrent_session(streaming=True) results = self._run_concurrent_session(streaming=True)
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") head_avg = _avg(_collect_latencies(results, first_n=HEAD_TURNS))
if reg_list: tail_avg = _avg(_collect_latencies(results, last_n=TAIL_TURNS))
reg_tail = _avg(_collect_latencies(reg_list, last_n=TAIL_TURNS)) ratio = tail_avg / head_avg if head_avg > 0 else float("inf")
stm_tail = _avg(_collect_latencies(results, last_n=TAIL_TURNS)) print(
speedup = reg_tail / stm_tail if stm_tail > 0 else float("inf") f"\n streaming_session "
self.assertGreaterEqual( f"head_avg(first {HEAD_TURNS})={head_avg:.1f}ms "
speedup, f"tail_avg(last {TAIL_TURNS})={tail_avg:.1f}ms "
1.4, f"ratio={ratio:.2f}"
f"streaming should be >=1.4x faster on last {TAIL_TURNS} turns " )
f"(regular={reg_tail:.1f}ms, streaming={stm_tail:.1f}ms, speedup={speedup:.2f}x)", self.assertLessEqual(
) ratio,
1.15,
f"streaming latency should stay flat across turns "
f"(head={head_avg:.1f}ms, tail={tail_avg:.1f}ms, ratio={ratio:.2f} > 1.15)",
)
def test_streaming_session_correctness(self): def test_streaming_session_correctness(self):
"""Correctness test: bs=1, assert output equal + latency speedup.""" """Correctness test: bs=1, assert regular and streaming outputs match."""
correctness_turns = 30 correctness_turns = 30
reg = self._run_concurrent_session( reg = self._run_concurrent_session(
streaming=False, num_concurrent=1, num_turns=correctness_turns streaming=False, num_concurrent=1, num_turns=correctness_turns