[Test] test_session_latency: assert streaming tail/head stability (#26230)
This commit is contained in:
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user