streaming session: spec v2 bonus accounting + comprehensive test matrix (#22651)
This commit is contained in:
@@ -358,7 +358,8 @@ class SchedulerOutputProcessorMixin:
|
|||||||
stride = self.draft_worker.speculative_num_draft_tokens
|
stride = self.draft_worker.speculative_num_draft_tokens
|
||||||
|
|
||||||
for i, req in enumerate(batch.reqs):
|
for i, req in enumerate(batch.reqs):
|
||||||
req.kv_committed_len += accept_lens[i]
|
# -1 because prepare_for_decode pre-claimed the bonus slot.
|
||||||
|
req.kv_committed_len += accept_lens[i] - 1
|
||||||
predict_tokens.append(
|
predict_tokens.append(
|
||||||
next_token_ids[i * stride : i * stride + accept_lens[i]]
|
next_token_ids[i * stride : i * stride + accept_lens[i]]
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1075,15 +1075,6 @@ class TokenizerCommunicatorMixin:
|
|||||||
"Streaming sessions are disabled. "
|
"Streaming sessions are disabled. "
|
||||||
"Please relaunch with --enable-streaming-session."
|
"Please relaunch with --enable-streaming-session."
|
||||||
)
|
)
|
||||||
if (
|
|
||||||
self.server_args.speculative_algorithm is not None
|
|
||||||
and not self.server_args.disable_overlap_schedule
|
|
||||||
):
|
|
||||||
raise ValueError(
|
|
||||||
"Streaming sessions are incompatible with speculative decoding v2 "
|
|
||||||
"(overlap + speculative). Use --disable-overlap-schedule or "
|
|
||||||
"disable speculative decoding."
|
|
||||||
)
|
|
||||||
|
|
||||||
if obj.session_id is None:
|
if obj.session_id is None:
|
||||||
obj.session_id = uuid.uuid4().hex
|
obj.session_id = uuid.uuid4().hex
|
||||||
|
|||||||
@@ -127,6 +127,8 @@ class EagleDraftInputV2Mixin:
|
|||||||
num_needed_tokens += x
|
num_needed_tokens += x
|
||||||
r.kv_allocated_len += x
|
r.kv_allocated_len += x
|
||||||
r.decode_batch_idx += 1
|
r.decode_batch_idx += 1
|
||||||
|
# Pre-claim bonus slot here (like normal decode); resolve subtracts 1.
|
||||||
|
r.kv_committed_len += 1
|
||||||
|
|
||||||
cur_kv_lens_cpu = torch.tensor(cur_kv_lens_cpu, dtype=torch.int32, device="cpu")
|
cur_kv_lens_cpu = torch.tensor(cur_kv_lens_cpu, dtype=torch.int32, device="cpu")
|
||||||
nxt_kv_lens_cpu = torch.tensor(nxt_kv_lens_cpu, dtype=torch.int32, device="cpu")
|
nxt_kv_lens_cpu = torch.tensor(nxt_kv_lens_cpu, dtype=torch.int32, device="cpu")
|
||||||
|
|||||||
@@ -1,14 +1,3 @@
|
|||||||
"""
|
|
||||||
Streaming session tests: KV cache mechanics, logprob leak, chunked prefill leak.
|
|
||||||
|
|
||||||
All tests share a single server (DEFAULT_SMALL_MODEL) with streaming sessions
|
|
||||||
and chunked prefill enabled.
|
|
||||||
|
|
||||||
Usage:
|
|
||||||
python -m pytest test_streaming_session.py -xvs
|
|
||||||
python -m unittest test_streaming_session.TestStreamingSession
|
|
||||||
"""
|
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
import time
|
import time
|
||||||
@@ -23,21 +12,16 @@ from sglang.srt.utils import kill_process_tree
|
|||||||
from sglang.srt.utils.hf_transformers_utils import get_tokenizer
|
from sglang.srt.utils.hf_transformers_utils import get_tokenizer
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
from sglang.test.test_utils import (
|
from sglang.test.test_utils import (
|
||||||
|
DEFAULT_DRAFT_MODEL_EAGLE3,
|
||||||
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
|
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
|
||||||
|
DEFAULT_TARGET_MODEL_EAGLE3,
|
||||||
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
DEFAULT_URL_FOR_TEST,
|
DEFAULT_URL_FOR_TEST,
|
||||||
CustomTestCase,
|
CustomTestCase,
|
||||||
popen_launch_server,
|
popen_launch_server,
|
||||||
)
|
)
|
||||||
|
|
||||||
register_cuda_ci(est_time=67, suite="stage-b-test-1-gpu-large")
|
register_cuda_ci(est_time=560, suite="stage-b-test-1-gpu-large")
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# Logprob leak constants
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
LOGPROB_NUM_TURNS = 5
|
|
||||||
LOGPROB_NUM_ROUNDS = 30
|
|
||||||
|
|
||||||
LOGPROB_PROMPTS = [
|
LOGPROB_PROMPTS = [
|
||||||
"The quick brown fox jumps over the lazy dog.",
|
"The quick brown fox jumps over the lazy dog.",
|
||||||
@@ -47,15 +31,7 @@ LOGPROB_PROMPTS = [
|
|||||||
"The five boxing wizards jump quickly.",
|
"The five boxing wizards jump quickly.",
|
||||||
]
|
]
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# Long enough to trigger chunked prefill at 200+ tokens per slice.
|
||||||
# Chunked prefill leak constants
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
LEAK_NUM_SESSIONS = 4
|
|
||||||
LEAK_NUM_TURNS = 5
|
|
||||||
LEAK_GEN_LEN = 16
|
|
||||||
|
|
||||||
# Filler text to trigger chunked prefill (200+ tokens per turn)
|
|
||||||
LEAK_FILLER = (
|
LEAK_FILLER = (
|
||||||
"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. "
|
||||||
@@ -67,12 +43,8 @@ LEAK_FILLER = (
|
|||||||
"We promptly judged antique ivory buckles for the next prize. "
|
"We promptly judged antique ivory buckles for the next prize. "
|
||||||
) * 20
|
) * 20
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# Abort-heavy chunked prefill leak repro constants
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
ABORT_REPRO_CONTEXT_LEN = 512
|
ABORT_REPRO_CONTEXT_LEN = 512
|
||||||
ABORT_REPRO_PAGE_SIZE = 16
|
ABORT_REPRO_PAGE_SIZE = 256
|
||||||
ABORT_REPRO_GEN_LEN = 4
|
ABORT_REPRO_GEN_LEN = 4
|
||||||
ABORT_REPRO_SESSIONS = 4
|
ABORT_REPRO_SESSIONS = 4
|
||||||
ABORT_REPRO_WARMUP_TURNS = 1
|
ABORT_REPRO_WARMUP_TURNS = 1
|
||||||
@@ -80,146 +52,16 @@ ABORT_REPRO_ROUNDS = 8
|
|||||||
ABORT_REPRO_STREAM_TOKENS = 16
|
ABORT_REPRO_STREAM_TOKENS = 16
|
||||||
ABORT_REPRO_ABORT_TOKENS = 600
|
ABORT_REPRO_ABORT_TOKENS = 600
|
||||||
ABORT_REPRO_NON_STREAMING_TOKENS = 16
|
ABORT_REPRO_NON_STREAMING_TOKENS = 16
|
||||||
ABORT_REPRO_CHUNKED_PREFILL_SIZE = 128
|
ABORT_REPRO_CHUNKED_PREFILL_SIZE = 4096
|
||||||
|
|
||||||
|
CONCURRENT_LOGPROB_SESSIONS = 6
|
||||||
|
CONCURRENT_LOGPROB_TURNS = 5
|
||||||
|
CONCURRENT_LOGPROB_ROUNDS = 10
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
STRESS_NUM_SESSIONS = 8
|
||||||
# Logprob leak helpers
|
STRESS_NUM_NON_STREAMING = 4
|
||||||
# ---------------------------------------------------------------------------
|
STRESS_NUM_TURNS = 6
|
||||||
|
STRESS_GEN_LEN = 16
|
||||||
|
|
||||||
def _logprob_generate(base_url, input_ids, **kwargs) -> dict:
|
|
||||||
payload: dict[str, Any] = {
|
|
||||||
"input_ids": input_ids,
|
|
||||||
"sampling_params": {
|
|
||||||
"temperature": 0,
|
|
||||||
"max_new_tokens": kwargs.get("max_new_tokens", 8),
|
|
||||||
"no_stop_trim": True,
|
|
||||||
"skip_special_tokens": False,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
if kwargs.get("return_logprob"):
|
|
||||||
payload["return_logprob"] = True
|
|
||||||
if kwargs.get("logprob_start_len") is not None:
|
|
||||||
payload["logprob_start_len"] = kwargs["logprob_start_len"]
|
|
||||||
if kwargs.get("session_params"):
|
|
||||||
payload["session_params"] = kwargs["session_params"]
|
|
||||||
resp = requests.post(base_url + "/generate", json=payload, timeout=120)
|
|
||||||
assert resp.status_code == 200, f"Generate failed: {resp.text}"
|
|
||||||
return resp.json()
|
|
||||||
|
|
||||||
|
|
||||||
def _logprob_run_one_session(base_url, tokenizer, **gen_kwargs):
|
|
||||||
"""Open session -> N turns -> close."""
|
|
||||||
resp = requests.post(
|
|
||||||
base_url + "/open_session",
|
|
||||||
json={"capacity_of_str_len": 50000, "streaming": True},
|
|
||||||
)
|
|
||||||
assert resp.status_code == 200
|
|
||||||
session_id = resp.json()
|
|
||||||
|
|
||||||
rid = None
|
|
||||||
for turn in range(LOGPROB_NUM_TURNS):
|
|
||||||
turn_ids = tokenizer.encode(
|
|
||||||
f"Turn {turn}: {LOGPROB_PROMPTS[turn % len(LOGPROB_PROMPTS)]}"
|
|
||||||
)
|
|
||||||
result = _logprob_generate(
|
|
||||||
base_url,
|
|
||||||
turn_ids,
|
|
||||||
session_params={"id": session_id, "rid": rid},
|
|
||||||
**gen_kwargs,
|
|
||||||
)
|
|
||||||
rid = result["meta_info"]["id"]
|
|
||||||
|
|
||||||
requests.post(base_url + "/close_session", json={"session_id": session_id})
|
|
||||||
|
|
||||||
|
|
||||||
def _logprob_assert_no_leak(base_url, tokenizer, **gen_kwargs):
|
|
||||||
"""Run many session rounds and verify server stays healthy."""
|
|
||||||
requests.post(base_url + "/flush_cache")
|
|
||||||
for _ in range(LOGPROB_NUM_ROUNDS):
|
|
||||||
_logprob_run_one_session(base_url, tokenizer, **gen_kwargs)
|
|
||||||
time.sleep(3)
|
|
||||||
assert (
|
|
||||||
requests.get(base_url + "/health").status_code == 200
|
|
||||||
), "Server unhealthy — likely a token memory leak."
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
# Chunked prefill leak helpers
|
|
||||||
# ---------------------------------------------------------------------------
|
|
||||||
|
|
||||||
|
|
||||||
async def _leak_async_generate(
|
|
||||||
base_url: str,
|
|
||||||
session: aiohttp.ClientSession,
|
|
||||||
input_ids: list[int],
|
|
||||||
session_params: Optional[dict[str, Any]] = None,
|
|
||||||
) -> Any:
|
|
||||||
payload: dict[str, Any] = {
|
|
||||||
"input_ids": input_ids,
|
|
||||||
"sampling_params": {
|
|
||||||
"temperature": 0,
|
|
||||||
"max_new_tokens": LEAK_GEN_LEN,
|
|
||||||
"no_stop_trim": True,
|
|
||||||
"skip_special_tokens": False,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
if session_params:
|
|
||||||
payload["session_params"] = session_params
|
|
||||||
timeout = aiohttp.ClientTimeout(total=300)
|
|
||||||
async with session.post(
|
|
||||||
base_url + "/generate", json=payload, timeout=timeout
|
|
||||||
) as resp:
|
|
||||||
assert resp.status == 200, f"Generate failed: {await resp.text()}"
|
|
||||||
return await resp.json()
|
|
||||||
|
|
||||||
|
|
||||||
async def _leak_run_all(base_url: str, tokenizer: Any) -> None:
|
|
||||||
"""Fire all requests per turn simultaneously to create mixed batches."""
|
|
||||||
timeout = aiohttp.ClientTimeout(total=300)
|
|
||||||
async with aiohttp.ClientSession(timeout=timeout) as http:
|
|
||||||
# Open all sessions
|
|
||||||
sids = []
|
|
||||||
for s in range(LEAK_NUM_SESSIONS):
|
|
||||||
async with http.post(
|
|
||||||
base_url + "/open_session",
|
|
||||||
json={"capacity_of_str_len": 50000, "streaming": True},
|
|
||||||
) as resp:
|
|
||||||
sids.append(await resp.json())
|
|
||||||
|
|
||||||
# For each turn, fire ALL streaming + non-streaming requests at once
|
|
||||||
for turn in range(LEAK_NUM_TURNS):
|
|
||||||
tasks = []
|
|
||||||
# Streaming requests for all sessions
|
|
||||||
for s in range(LEAK_NUM_SESSIONS):
|
|
||||||
offset = (s * LEAK_NUM_TURNS + turn) * 200
|
|
||||||
text = f"Session {s} turn {turn}: {LEAK_FILLER[offset : offset + 1500]}"
|
|
||||||
ids = tokenizer.encode(text)
|
|
||||||
tasks.append(
|
|
||||||
_leak_async_generate(
|
|
||||||
base_url,
|
|
||||||
http,
|
|
||||||
ids,
|
|
||||||
session_params={"id": sids[s], "rid": None},
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
# Non-streaming requests interleaved
|
|
||||||
for ns in range(LEAK_NUM_SESSIONS // 2):
|
|
||||||
text = f"Non-streaming {ns} turn {turn}: {LEAK_FILLER[ns * 100 : ns * 100 + 500]}"
|
|
||||||
ids = tokenizer.encode(text)
|
|
||||||
tasks.append(_leak_async_generate(base_url, http, ids))
|
|
||||||
|
|
||||||
# Fire all at once — creates mixed batch of streaming + non-streaming
|
|
||||||
await asyncio.gather(*tasks)
|
|
||||||
|
|
||||||
# Close all sessions
|
|
||||||
for sid in sids:
|
|
||||||
async with http.post(
|
|
||||||
base_url + "/close_session", json={"session_id": sid}
|
|
||||||
) as resp:
|
|
||||||
assert resp.status == 200
|
|
||||||
|
|
||||||
|
|
||||||
def _make_token_sized_ids(
|
def _make_token_sized_ids(
|
||||||
@@ -402,9 +244,144 @@ async def _abort_repro_run_all(base_url: str, tokenizer: Any) -> None:
|
|||||||
assert resp.status == 200, await resp.text()
|
assert resp.status == 200, await resp.text()
|
||||||
|
|
||||||
|
|
||||||
# ===================================================================
|
async def _async_generate(
|
||||||
# Test class
|
base_url: str,
|
||||||
# ===================================================================
|
session: aiohttp.ClientSession,
|
||||||
|
input_ids: list[int],
|
||||||
|
max_new_tokens: int = 8,
|
||||||
|
session_params: Optional[dict[str, Any]] = None,
|
||||||
|
return_logprob: bool = False,
|
||||||
|
logprob_start_len: Optional[int] = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
payload: dict[str, Any] = {
|
||||||
|
"input_ids": input_ids,
|
||||||
|
"sampling_params": {
|
||||||
|
"temperature": 0,
|
||||||
|
"max_new_tokens": max_new_tokens,
|
||||||
|
"no_stop_trim": True,
|
||||||
|
"skip_special_tokens": False,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
if session_params:
|
||||||
|
payload["session_params"] = session_params
|
||||||
|
if return_logprob:
|
||||||
|
payload["return_logprob"] = True
|
||||||
|
if logprob_start_len is not None:
|
||||||
|
payload["logprob_start_len"] = logprob_start_len
|
||||||
|
timeout = aiohttp.ClientTimeout(total=300)
|
||||||
|
async with session.post(
|
||||||
|
base_url + "/generate", json=payload, timeout=timeout
|
||||||
|
) as resp:
|
||||||
|
assert resp.status == 200, f"Generate failed: {await resp.text()}"
|
||||||
|
return await resp.json()
|
||||||
|
|
||||||
|
|
||||||
|
async def _concurrent_logprob_run(base_url: str, tokenizer: Any, **gen_kwargs) -> None:
|
||||||
|
"""N sessions per round, all requests fired simultaneously per turn so
|
||||||
|
the running batch has real concurrency (retract can actually kick one).
|
||||||
|
"""
|
||||||
|
timeout = aiohttp.ClientTimeout(total=300)
|
||||||
|
async with aiohttp.ClientSession(timeout=timeout) as http:
|
||||||
|
for _ in range(CONCURRENT_LOGPROB_ROUNDS):
|
||||||
|
sids: list[str] = []
|
||||||
|
for _ in range(CONCURRENT_LOGPROB_SESSIONS):
|
||||||
|
async with http.post(
|
||||||
|
base_url + "/open_session",
|
||||||
|
json={"capacity_of_str_len": 50000, "streaming": True},
|
||||||
|
) as resp:
|
||||||
|
assert resp.status == 200
|
||||||
|
sids.append(await resp.json())
|
||||||
|
|
||||||
|
rids: list[Optional[str]] = [None] * CONCURRENT_LOGPROB_SESSIONS
|
||||||
|
for turn in range(CONCURRENT_LOGPROB_TURNS):
|
||||||
|
tasks = []
|
||||||
|
for s in range(CONCURRENT_LOGPROB_SESSIONS):
|
||||||
|
text = (
|
||||||
|
f"S{s} T{turn}: "
|
||||||
|
f"{LOGPROB_PROMPTS[turn % len(LOGPROB_PROMPTS)]}"
|
||||||
|
)
|
||||||
|
ids = tokenizer.encode(text)
|
||||||
|
tasks.append(
|
||||||
|
_async_generate(
|
||||||
|
base_url,
|
||||||
|
http,
|
||||||
|
ids,
|
||||||
|
session_params={"id": sids[s], "rid": rids[s]},
|
||||||
|
**gen_kwargs,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
results = await asyncio.gather(*tasks)
|
||||||
|
for s in range(CONCURRENT_LOGPROB_SESSIONS):
|
||||||
|
rids[s] = results[s]["meta_info"]["id"]
|
||||||
|
|
||||||
|
for sid in sids:
|
||||||
|
async with http.post(
|
||||||
|
base_url + "/close_session", json={"session_id": sid}
|
||||||
|
) as resp:
|
||||||
|
assert resp.status == 200
|
||||||
|
|
||||||
|
|
||||||
|
async def _stress_run_all(base_url: str, tokenizer: Any) -> None:
|
||||||
|
"""Streaming + non-streaming mixed batches under retract pressure.
|
||||||
|
Long prompts (~200+ tokens) trigger chunked prefill so retract can
|
||||||
|
interrupt mid-extend.
|
||||||
|
"""
|
||||||
|
timeout = aiohttp.ClientTimeout(total=300)
|
||||||
|
async with aiohttp.ClientSession(timeout=timeout) as http:
|
||||||
|
sids: list[str] = []
|
||||||
|
for _ in range(STRESS_NUM_SESSIONS):
|
||||||
|
async with http.post(
|
||||||
|
base_url + "/open_session",
|
||||||
|
json={"capacity_of_str_len": 50000, "streaming": True},
|
||||||
|
) as resp:
|
||||||
|
assert resp.status == 200
|
||||||
|
sids.append(await resp.json())
|
||||||
|
|
||||||
|
rids: list[Optional[str]] = [None] * STRESS_NUM_SESSIONS
|
||||||
|
for turn in range(STRESS_NUM_TURNS):
|
||||||
|
tasks = []
|
||||||
|
# Streaming requests — long prompts to trigger chunked prefill.
|
||||||
|
for s in range(STRESS_NUM_SESSIONS):
|
||||||
|
offset = (s * STRESS_NUM_TURNS + turn) * 200
|
||||||
|
text = (
|
||||||
|
f"Session {s} turn {turn}: " f"{LEAK_FILLER[offset : offset + 800]}"
|
||||||
|
)
|
||||||
|
ids = tokenizer.encode(text)
|
||||||
|
tasks.append(
|
||||||
|
_async_generate(
|
||||||
|
base_url,
|
||||||
|
http,
|
||||||
|
ids,
|
||||||
|
max_new_tokens=STRESS_GEN_LEN,
|
||||||
|
session_params={"id": sids[s], "rid": rids[s]},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
# Non-streaming requests interleaved.
|
||||||
|
for ns in range(STRESS_NUM_NON_STREAMING):
|
||||||
|
text = (
|
||||||
|
f"Non-streaming {ns} turn {turn}: "
|
||||||
|
f"{LEAK_FILLER[ns * 100 : ns * 100 + 400]}"
|
||||||
|
)
|
||||||
|
ids = tokenizer.encode(text)
|
||||||
|
tasks.append(
|
||||||
|
_async_generate(
|
||||||
|
base_url,
|
||||||
|
http,
|
||||||
|
ids,
|
||||||
|
max_new_tokens=STRESS_GEN_LEN,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
results = await asyncio.gather(*tasks)
|
||||||
|
for s in range(STRESS_NUM_SESSIONS):
|
||||||
|
rids[s] = results[s]["meta_info"]["id"]
|
||||||
|
|
||||||
|
for sid in sids:
|
||||||
|
async with http.post(
|
||||||
|
base_url + "/close_session", json={"session_id": sid}
|
||||||
|
) as resp:
|
||||||
|
assert resp.status == 200
|
||||||
|
|
||||||
|
|
||||||
class TestStreamingSession(CustomTestCase):
|
class TestStreamingSession(CustomTestCase):
|
||||||
@@ -429,8 +406,13 @@ class TestStreamingSession(CustomTestCase):
|
|||||||
def tearDownClass(cls):
|
def tearDownClass(cls):
|
||||||
kill_process_tree(cls.process.pid)
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
# -1 for non-overlap subclasses: the last sampled token isn't committed
|
||||||
|
# before max_new stops, so slot.kv_committed_len = input + output - 1.
|
||||||
|
kv_inherit_offset = 0
|
||||||
|
|
||||||
def test_kv_cache_inheritance(self, gen_len=12):
|
def test_kv_cache_inheritance(self, gen_len=12):
|
||||||
"""Verify KV inheritance, radix cache insertion, and flush reclamation."""
|
"""Each turn's cached_tokens must equal previous turn's prompt+completion
|
||||||
|
(modulo kv_inherit_offset)."""
|
||||||
chunks = [
|
chunks = [
|
||||||
"Let me tell you something about France.",
|
"Let me tell you something about France.",
|
||||||
"The capital of France is",
|
"The capital of France is",
|
||||||
@@ -470,102 +452,58 @@ class TestStreamingSession(CustomTestCase):
|
|||||||
completion_tokens = response["meta_info"]["completion_tokens"]
|
completion_tokens = response["meta_info"]["completion_tokens"]
|
||||||
|
|
||||||
if turn_idx == 0:
|
if turn_idx == 0:
|
||||||
# Turn 1 should have no cache hit (cache was flushed).
|
# Turn 1: cache flushed, no hit.
|
||||||
self.assertEqual(
|
self.assertEqual(cached, 0, "Turn 1: clean start, no cache hit")
|
||||||
cached, 0, "Turn 1 should have 0 cached tokens (clean start)"
|
|
||||||
)
|
|
||||||
else:
|
else:
|
||||||
# Turns 2+ inherit KV from the previous turn (via inherit_kv_states,
|
# Turns 2+: cached_tokens reflects KV inherited from previous turn
|
||||||
# not radix tree matching). cached_tokens reflects the inherited prefix.
|
# (via inherit_kv_states, not radix tree matching).
|
||||||
|
expected = prev_kv_len + self.kv_inherit_offset
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
cached,
|
cached,
|
||||||
prev_kv_len,
|
expected,
|
||||||
f"Turn {turn_idx + 1}: should inherit {prev_kv_len} KV tokens from previous turn",
|
f"Turn {turn_idx + 1}: inherited {cached} != expected {expected}",
|
||||||
)
|
)
|
||||||
prev_kv_len = prompt_tokens + completion_tokens
|
prev_kv_len = prompt_tokens + completion_tokens
|
||||||
|
|
||||||
# Close the session before checking cache/memory state.
|
# Close the session.
|
||||||
ret = requests.post(
|
ret = requests.post(
|
||||||
self.base_url + "/close_session",
|
self.base_url + "/close_session",
|
||||||
json={"session_id": session_id},
|
json={"session_id": session_id},
|
||||||
)
|
)
|
||||||
self.assertEqual(ret.status_code, 200)
|
self.assertEqual(ret.status_code, 200)
|
||||||
|
|
||||||
# === Cache verification (after close, before flush) ===
|
def test_leak_logprob_concurrent(self) -> None:
|
||||||
|
"""Concurrent multi-session × 3 logprob modes (output / input / none),
|
||||||
# Turn 1's prompt was inserted to the cache.
|
watch for KV leak."""
|
||||||
verify_resp = requests.post(
|
|
||||||
self.base_url + "/generate",
|
|
||||||
json={
|
|
||||||
"input_ids": chunks_ids[0],
|
|
||||||
"sampling_params": {"temperature": 0, "max_new_tokens": 1},
|
|
||||||
},
|
|
||||||
).json()
|
|
||||||
self.assertGreater(
|
|
||||||
verify_resp["meta_info"]["cached_tokens"],
|
|
||||||
0,
|
|
||||||
"Turn 1's prompt should be cached in the radix tree",
|
|
||||||
)
|
|
||||||
|
|
||||||
# Turn 2's prompt tokens should NOT be in cache.
|
|
||||||
# The tree should only contain turn 1's extent (prompt + output from
|
|
||||||
# cache_unfinished_req during decode). Turn 2's prompt starts fresh tokens
|
|
||||||
# that were never inserted.
|
|
||||||
verify_resp2 = requests.post(
|
|
||||||
self.base_url + "/generate",
|
|
||||||
json={
|
|
||||||
"input_ids": chunks_ids[1],
|
|
||||||
"sampling_params": {"temperature": 0, "max_new_tokens": 1},
|
|
||||||
},
|
|
||||||
).json()
|
|
||||||
self.assertEqual(
|
|
||||||
verify_resp2["meta_info"]["cached_tokens"],
|
|
||||||
0,
|
|
||||||
"Turn 2's prompt should not be in cache (no insertion for turns 2+)",
|
|
||||||
)
|
|
||||||
|
|
||||||
# === Flush reclamation ===
|
|
||||||
|
|
||||||
requests.post(self.base_url + "/flush_cache")
|
requests.post(self.base_url + "/flush_cache")
|
||||||
verify_resp3 = requests.post(
|
# Output logprob
|
||||||
self.base_url + "/generate",
|
asyncio.run(
|
||||||
json={
|
_concurrent_logprob_run(self.base_url, self.tokenizer, return_logprob=True)
|
||||||
"input_ids": chunks_ids[0],
|
|
||||||
"sampling_params": {"temperature": 0, "max_new_tokens": 1},
|
|
||||||
},
|
|
||||||
).json()
|
|
||||||
self.assertEqual(
|
|
||||||
verify_resp3["meta_info"]["cached_tokens"],
|
|
||||||
0,
|
|
||||||
"After session close + flush, cache should be fully reclaimed",
|
|
||||||
)
|
)
|
||||||
|
# Input logprob (logprob_start_len=0)
|
||||||
def test_leak_logprob_none(self) -> None:
|
asyncio.run(
|
||||||
"""Streaming sessions without logprobs must not leak tokens."""
|
_concurrent_logprob_run(
|
||||||
_logprob_assert_no_leak(self.base_url, self.tokenizer)
|
self.base_url,
|
||||||
|
self.tokenizer,
|
||||||
def test_leak_logprob_output(self) -> None:
|
return_logprob=True,
|
||||||
"""Streaming sessions with output logprobs must not leak tokens."""
|
logprob_start_len=0,
|
||||||
_logprob_assert_no_leak(self.base_url, self.tokenizer, return_logprob=True)
|
)
|
||||||
|
|
||||||
def test_leak_logprob_input(self) -> None:
|
|
||||||
"""Streaming sessions with logprob_start_len=0 must not leak tokens."""
|
|
||||||
_logprob_assert_no_leak(
|
|
||||||
self.base_url,
|
|
||||||
self.tokenizer,
|
|
||||||
return_logprob=True,
|
|
||||||
logprob_start_len=0,
|
|
||||||
)
|
)
|
||||||
|
# No logprob
|
||||||
|
asyncio.run(_concurrent_logprob_run(self.base_url, self.tokenizer))
|
||||||
|
time.sleep(3)
|
||||||
|
assert (
|
||||||
|
requests.get(self.base_url + "/health").status_code == 200
|
||||||
|
), "Server unhealthy after concurrent logprob sessions."
|
||||||
|
|
||||||
def test_leak_chunked_prefill(self) -> None:
|
def test_stress_concurrent_sessions(self) -> None:
|
||||||
"""Concurrent multi-turn streaming sessions then idle health check."""
|
"""High concurrency streaming + non-streaming with retract pressure;
|
||||||
|
scheduler must roll back streaming KV without leaking."""
|
||||||
requests.post(self.base_url + "/flush_cache")
|
requests.post(self.base_url + "/flush_cache")
|
||||||
|
asyncio.run(_stress_run_all(self.base_url, self.tokenizer))
|
||||||
|
|
||||||
asyncio.run(_leak_run_all(self.base_url, self.tokenizer))
|
|
||||||
|
|
||||||
# Run a few non-streaming requests to flush state
|
|
||||||
for i in range(3):
|
for i in range(3):
|
||||||
ids = self.tokenizer.encode(f"Flush request {i}: final cleanup.")
|
ids = self.tokenizer.encode(f"Post-stress cleanup {i}.")
|
||||||
requests.post(
|
requests.post(
|
||||||
self.base_url + "/generate",
|
self.base_url + "/generate",
|
||||||
json={
|
json={
|
||||||
@@ -574,19 +512,18 @@ class TestStreamingSession(CustomTestCase):
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
# Wait for server to go idle and run memory check
|
|
||||||
time.sleep(5)
|
time.sleep(5)
|
||||||
health = requests.get(self.base_url + "/health")
|
health = requests.get(self.base_url + "/health")
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
health.status_code,
|
health.status_code,
|
||||||
200,
|
200,
|
||||||
"Server unhealthy after streaming session close — "
|
"Server unhealthy after concurrent stress test — "
|
||||||
"likely a token memory leak from streaming session lifecycle.",
|
"likely a token leak from retract/mixed-chunk + streaming session.",
|
||||||
)
|
)
|
||||||
|
|
||||||
def test_nth_mid_abort_recovery(self) -> None:
|
def test_nth_mid_abort_recovery(self) -> None:
|
||||||
"""Abort a running streaming session request (nth turn) via the
|
"""Abort an Nth-turn request mid-decode; session rolls back to last
|
||||||
abort API. Session rolls back to last successful turn."""
|
successful turn."""
|
||||||
requests.post(self.base_url + "/flush_cache")
|
requests.post(self.base_url + "/flush_cache")
|
||||||
|
|
||||||
resp = requests.post(
|
resp = requests.post(
|
||||||
@@ -691,9 +628,8 @@ class TestStreamingSession(CustomTestCase):
|
|||||||
self.assertEqual(health.status_code, 200)
|
self.assertEqual(health.status_code, 200)
|
||||||
|
|
||||||
def test_first_mid_abort_recovery(self) -> None:
|
def test_first_mid_abort_recovery(self) -> None:
|
||||||
"""Abort the very first request on a streaming session mid-decode.
|
"""Abort the very first request mid-decode (no slot yet; ephemeral
|
||||||
No slot exists yet (ephemeral slot created and nuked).
|
slot is created and nuked). Session must still be usable."""
|
||||||
Verify the session is still usable afterward."""
|
|
||||||
requests.post(self.base_url + "/flush_cache")
|
requests.post(self.base_url + "/flush_cache")
|
||||||
|
|
||||||
resp = requests.post(
|
resp = requests.post(
|
||||||
@@ -776,8 +712,8 @@ class TestStreamingSession(CustomTestCase):
|
|||||||
self.assertEqual(health.status_code, 200)
|
self.assertEqual(health.status_code, 200)
|
||||||
|
|
||||||
def test_preabort_recovery(self) -> None:
|
def test_preabort_recovery(self) -> None:
|
||||||
"""Pre-aborted request (unsupported offset) does not corrupt session.
|
"""Pre-abort (rejected by create_req) preserves the slot; next turn
|
||||||
The slot is preserved, and the next turn inherits correctly."""
|
inherits correctly."""
|
||||||
requests.post(self.base_url + "/flush_cache")
|
requests.post(self.base_url + "/flush_cache")
|
||||||
|
|
||||||
resp = requests.post(
|
resp = requests.post(
|
||||||
@@ -853,14 +789,121 @@ class TestStreamingSession(CustomTestCase):
|
|||||||
self.assertEqual(health.status_code, 200)
|
self.assertEqual(health.status_code, 200)
|
||||||
|
|
||||||
|
|
||||||
class TestStreamingSessionMixedChunk(TestStreamingSession):
|
class TestStreamingSessionRetractMixedChunk(TestStreamingSession):
|
||||||
"""Streaming session with --enable-mixed-chunk."""
|
"""Retract + --enable-mixed-chunk."""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def setUpClass(cls):
|
def setUpClass(cls):
|
||||||
cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
|
cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
|
||||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
with envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2):
|
with envs.SGLANG_TEST_RETRACT.override(
|
||||||
|
True
|
||||||
|
), envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2):
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
cls.model,
|
||||||
|
cls.base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=[
|
||||||
|
"--enable-streaming-session",
|
||||||
|
"--chunked-prefill-size",
|
||||||
|
"128",
|
||||||
|
"--enable-mixed-chunk",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
cls.tokenizer = get_tokenizer(cls.model)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
|
||||||
|
class TestStreamingSessionRetractLargePage(TestStreamingSession):
|
||||||
|
"""Retract + page=256: exercises page-aligned `_free_tail`. Partial-page
|
||||||
|
free would corrupt pages still holding committed tokens."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
|
||||||
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
with envs.SGLANG_TEST_RETRACT.override(
|
||||||
|
True
|
||||||
|
), envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2):
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
cls.model,
|
||||||
|
cls.base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=[
|
||||||
|
"--enable-streaming-session",
|
||||||
|
"--chunked-prefill-size",
|
||||||
|
"4096",
|
||||||
|
"--page-size",
|
||||||
|
"256",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
cls.tokenizer = get_tokenizer(cls.model)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
|
||||||
|
class TestStreamingSessionEagle(TestStreamingSession):
|
||||||
|
"""EAGLE3 spec v1 (overlap disabled); offset=-1 — see base class note."""
|
||||||
|
|
||||||
|
kv_inherit_offset = -1
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.model = DEFAULT_TARGET_MODEL_EAGLE3
|
||||||
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
with envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(
|
||||||
|
2
|
||||||
|
), envs.SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN.override(True):
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
cls.model,
|
||||||
|
cls.base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=[
|
||||||
|
"--enable-streaming-session",
|
||||||
|
"--disable-overlap-schedule",
|
||||||
|
"--chunked-prefill-size",
|
||||||
|
"512",
|
||||||
|
"--dtype=float16",
|
||||||
|
"--speculative-algorithm",
|
||||||
|
"EAGLE3",
|
||||||
|
"--speculative-draft-model",
|
||||||
|
DEFAULT_DRAFT_MODEL_EAGLE3,
|
||||||
|
"--speculative-num-steps",
|
||||||
|
"3",
|
||||||
|
"--speculative-eagle-topk",
|
||||||
|
"1",
|
||||||
|
"--speculative-num-draft-tokens",
|
||||||
|
"4",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.7",
|
||||||
|
],
|
||||||
|
)
|
||||||
|
cls.tokenizer = get_tokenizer(cls.model)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
|
||||||
|
class TestStreamingSessionEagleV2(TestStreamingSession):
|
||||||
|
"""EAGLE3 spec v2 (overlap on)."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.model = DEFAULT_TARGET_MODEL_EAGLE3
|
||||||
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
with envs.SGLANG_ENABLE_SPEC_V2.override(
|
||||||
|
True
|
||||||
|
), envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(
|
||||||
|
2
|
||||||
|
), envs.SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN.override(
|
||||||
|
True
|
||||||
|
):
|
||||||
cls.process = popen_launch_server(
|
cls.process = popen_launch_server(
|
||||||
cls.model,
|
cls.model,
|
||||||
cls.base_url,
|
cls.base_url,
|
||||||
@@ -869,7 +912,19 @@ class TestStreamingSessionMixedChunk(TestStreamingSession):
|
|||||||
"--enable-streaming-session",
|
"--enable-streaming-session",
|
||||||
"--chunked-prefill-size",
|
"--chunked-prefill-size",
|
||||||
"512",
|
"512",
|
||||||
"--enable-mixed-chunk",
|
"--dtype=float16",
|
||||||
|
"--speculative-algorithm",
|
||||||
|
"EAGLE3",
|
||||||
|
"--speculative-draft-model",
|
||||||
|
DEFAULT_DRAFT_MODEL_EAGLE3,
|
||||||
|
"--speculative-num-steps",
|
||||||
|
"3",
|
||||||
|
"--speculative-eagle-topk",
|
||||||
|
"1",
|
||||||
|
"--speculative-num-draft-tokens",
|
||||||
|
"4",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.7",
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
cls.tokenizer = get_tokenizer(cls.model)
|
cls.tokenizer = get_tokenizer(cls.model)
|
||||||
@@ -879,25 +934,47 @@ class TestStreamingSessionMixedChunk(TestStreamingSession):
|
|||||||
kill_process_tree(cls.process.pid)
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
|
||||||
@unittest.skip("streaming session + retract has a token leak — tracked separately")
|
class TestStreamingSessionEagleRetractLargePage(TestStreamingSession):
|
||||||
class TestStreamingSessionRetract(TestStreamingSession):
|
"""EAGLE3 spec v1 + retract + page=256: max-pressure on `_free_tail`
|
||||||
"""Streaming session under retract decode pressure."""
|
(spec tail + retract alloc-commit gap + page alignment)."""
|
||||||
|
|
||||||
|
kv_inherit_offset = -1
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def setUpClass(cls):
|
def setUpClass(cls):
|
||||||
cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
|
cls.model = DEFAULT_TARGET_MODEL_EAGLE3
|
||||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
with envs.SGLANG_TEST_RETRACT.override(
|
with envs.SGLANG_TEST_RETRACT.override(
|
||||||
True
|
True
|
||||||
), envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2):
|
), envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(
|
||||||
|
2
|
||||||
|
), envs.SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN.override(
|
||||||
|
True
|
||||||
|
):
|
||||||
cls.process = popen_launch_server(
|
cls.process = popen_launch_server(
|
||||||
cls.model,
|
cls.model,
|
||||||
cls.base_url,
|
cls.base_url,
|
||||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
other_args=[
|
other_args=[
|
||||||
"--enable-streaming-session",
|
"--enable-streaming-session",
|
||||||
|
"--disable-overlap-schedule",
|
||||||
"--chunked-prefill-size",
|
"--chunked-prefill-size",
|
||||||
"128",
|
"4096",
|
||||||
|
"--dtype=float16",
|
||||||
|
"--speculative-algorithm",
|
||||||
|
"EAGLE3",
|
||||||
|
"--speculative-draft-model",
|
||||||
|
DEFAULT_DRAFT_MODEL_EAGLE3,
|
||||||
|
"--speculative-num-steps",
|
||||||
|
"3",
|
||||||
|
"--speculative-eagle-topk",
|
||||||
|
"1",
|
||||||
|
"--speculative-num-draft-tokens",
|
||||||
|
"4",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.7",
|
||||||
|
"--page-size",
|
||||||
|
"256",
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
cls.tokenizer = get_tokenizer(cls.model)
|
cls.tokenizer = get_tokenizer(cls.model)
|
||||||
@@ -907,17 +984,22 @@ class TestStreamingSessionRetract(TestStreamingSession):
|
|||||||
kill_process_tree(cls.process.pid)
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
|
||||||
@unittest.skip("streaming session + retract has a token leak — tracked separately")
|
class TestStreamingSessionEagleV2RetractLargePage(TestStreamingSession):
|
||||||
class TestStreamingSessionRetractMixedChunk(TestStreamingSession):
|
"""EAGLE3 spec v2 + retract + page=256."""
|
||||||
"""Streaming session under retract decode with --enable-mixed-chunk."""
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def setUpClass(cls):
|
def setUpClass(cls):
|
||||||
cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
|
cls.model = DEFAULT_TARGET_MODEL_EAGLE3
|
||||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
with envs.SGLANG_TEST_RETRACT.override(
|
with envs.SGLANG_ENABLE_SPEC_V2.override(
|
||||||
True
|
True
|
||||||
), envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2):
|
), envs.SGLANG_TEST_RETRACT.override(
|
||||||
|
True
|
||||||
|
), envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(
|
||||||
|
2
|
||||||
|
), envs.SGLANG_ALLOW_OVERWRITE_LONGER_CONTEXT_LEN.override(
|
||||||
|
True
|
||||||
|
):
|
||||||
cls.process = popen_launch_server(
|
cls.process = popen_launch_server(
|
||||||
cls.model,
|
cls.model,
|
||||||
cls.base_url,
|
cls.base_url,
|
||||||
@@ -925,8 +1007,22 @@ class TestStreamingSessionRetractMixedChunk(TestStreamingSession):
|
|||||||
other_args=[
|
other_args=[
|
||||||
"--enable-streaming-session",
|
"--enable-streaming-session",
|
||||||
"--chunked-prefill-size",
|
"--chunked-prefill-size",
|
||||||
"128",
|
"4096",
|
||||||
"--enable-mixed-chunk",
|
"--dtype=float16",
|
||||||
|
"--speculative-algorithm",
|
||||||
|
"EAGLE3",
|
||||||
|
"--speculative-draft-model",
|
||||||
|
DEFAULT_DRAFT_MODEL_EAGLE3,
|
||||||
|
"--speculative-num-steps",
|
||||||
|
"3",
|
||||||
|
"--speculative-eagle-topk",
|
||||||
|
"1",
|
||||||
|
"--speculative-num-draft-tokens",
|
||||||
|
"4",
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.7",
|
||||||
|
"--page-size",
|
||||||
|
"256",
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
cls.tokenizer = get_tokenizer(cls.model)
|
cls.tokenizer = get_tokenizer(cls.model)
|
||||||
@@ -941,24 +1037,25 @@ class TestStreamingSessionAbortLeakRepro(CustomTestCase):
|
|||||||
def setUpClass(cls):
|
def setUpClass(cls):
|
||||||
cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
|
cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
|
||||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
cls.process = popen_launch_server(
|
with envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2):
|
||||||
cls.model,
|
cls.process = popen_launch_server(
|
||||||
cls.base_url,
|
cls.model,
|
||||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
cls.base_url,
|
||||||
other_args=[
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
"--enable-streaming-session",
|
other_args=[
|
||||||
"--chunked-prefill-size",
|
"--enable-streaming-session",
|
||||||
str(ABORT_REPRO_CHUNKED_PREFILL_SIZE),
|
"--chunked-prefill-size",
|
||||||
"--context-length",
|
str(ABORT_REPRO_CHUNKED_PREFILL_SIZE),
|
||||||
str(ABORT_REPRO_CONTEXT_LEN),
|
"--context-length",
|
||||||
"--page-size",
|
str(ABORT_REPRO_CONTEXT_LEN),
|
||||||
str(ABORT_REPRO_PAGE_SIZE),
|
"--page-size",
|
||||||
"--max-running-requests",
|
str(ABORT_REPRO_PAGE_SIZE),
|
||||||
"32",
|
"--max-running-requests",
|
||||||
"--log-level",
|
"32",
|
||||||
"info",
|
"--log-level",
|
||||||
],
|
"info",
|
||||||
)
|
],
|
||||||
|
)
|
||||||
cls.tokenizer = get_tokenizer(cls.model)
|
cls.tokenizer = get_tokenizer(cls.model)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
|
|||||||
@@ -0,0 +1,160 @@
|
|||||||
|
import os
|
||||||
|
import sys
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
from sglang.srt.utils import kill_process_tree
|
||||||
|
from sglang.srt.utils.hf_transformers_utils import get_tokenizer
|
||||||
|
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,
|
||||||
|
popen_launch_server,
|
||||||
|
)
|
||||||
|
|
||||||
|
# test/ has no __init__.py; add sibling dir so sibling module is importable
|
||||||
|
# when this file is run as a script via `python3 <path>`.
|
||||||
|
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||||
|
|
||||||
|
from test_streaming_session import ( # noqa: E402
|
||||||
|
ABORT_REPRO_CHUNKED_PREFILL_SIZE,
|
||||||
|
ABORT_REPRO_CONTEXT_LEN,
|
||||||
|
ABORT_REPRO_PAGE_SIZE,
|
||||||
|
TestStreamingSession,
|
||||||
|
TestStreamingSessionAbortLeakRepro,
|
||||||
|
)
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=450, suite="stage-b-test-1-gpu-large")
|
||||||
|
|
||||||
|
|
||||||
|
SWA_MODEL = "openai/gpt-oss-20b"
|
||||||
|
|
||||||
|
# Common gpt-oss-20b launch args. Matches TestSessionLatency/TestSWARadixCacheKL.
|
||||||
|
SWA_COMMON_ARGS = [
|
||||||
|
"--mem-fraction-static",
|
||||||
|
"0.70",
|
||||||
|
"--disable-piecewise-cuda-graph",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class TestStreamingSessionSWA(TestStreamingSession):
|
||||||
|
"""Baseline streaming session on a hybrid-SWA model."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.model = SWA_MODEL
|
||||||
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
with envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2):
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
cls.model,
|
||||||
|
cls.base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=[
|
||||||
|
"--enable-streaming-session",
|
||||||
|
"--chunked-prefill-size",
|
||||||
|
"512",
|
||||||
|
*SWA_COMMON_ARGS,
|
||||||
|
],
|
||||||
|
)
|
||||||
|
cls.tokenizer = get_tokenizer(cls.model)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
|
||||||
|
class TestStreamingSessionSWARetractLargePage(TestStreamingSession):
|
||||||
|
"""SWA under retract decode with page=256."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.model = SWA_MODEL
|
||||||
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
with envs.SGLANG_TEST_RETRACT.override(
|
||||||
|
True
|
||||||
|
), envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2):
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
cls.model,
|
||||||
|
cls.base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=[
|
||||||
|
"--enable-streaming-session",
|
||||||
|
"--chunked-prefill-size",
|
||||||
|
"4096",
|
||||||
|
"--page-size",
|
||||||
|
"256",
|
||||||
|
*SWA_COMMON_ARGS,
|
||||||
|
],
|
||||||
|
)
|
||||||
|
cls.tokenizer = get_tokenizer(cls.model)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
|
||||||
|
class TestStreamingSessionSWARetractMixedChunk(TestStreamingSession):
|
||||||
|
"""SWA under retract decode with --enable-mixed-chunk."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.model = SWA_MODEL
|
||||||
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
with envs.SGLANG_TEST_RETRACT.override(
|
||||||
|
True
|
||||||
|
), envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2):
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
cls.model,
|
||||||
|
cls.base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=[
|
||||||
|
"--enable-streaming-session",
|
||||||
|
"--chunked-prefill-size",
|
||||||
|
"128",
|
||||||
|
"--enable-mixed-chunk",
|
||||||
|
*SWA_COMMON_ARGS,
|
||||||
|
],
|
||||||
|
)
|
||||||
|
cls.tokenizer = get_tokenizer(cls.model)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
|
||||||
|
class TestStreamingSessionSWAAbortLeakRepro(TestStreamingSessionAbortLeakRepro):
|
||||||
|
"""SWA abort-heavy chunked prefill leak repro."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
cls.model = SWA_MODEL
|
||||||
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
with envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY.override(2):
|
||||||
|
cls.process = popen_launch_server(
|
||||||
|
cls.model,
|
||||||
|
cls.base_url,
|
||||||
|
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||||
|
other_args=[
|
||||||
|
"--enable-streaming-session",
|
||||||
|
"--chunked-prefill-size",
|
||||||
|
str(ABORT_REPRO_CHUNKED_PREFILL_SIZE),
|
||||||
|
"--context-length",
|
||||||
|
str(ABORT_REPRO_CONTEXT_LEN),
|
||||||
|
"--page-size",
|
||||||
|
str(ABORT_REPRO_PAGE_SIZE),
|
||||||
|
"--max-running-requests",
|
||||||
|
"32",
|
||||||
|
"--log-level",
|
||||||
|
"info",
|
||||||
|
*SWA_COMMON_ARGS,
|
||||||
|
],
|
||||||
|
)
|
||||||
|
cls.tokenizer = get_tokenizer(cls.model)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls):
|
||||||
|
kill_process_tree(cls.process.pid)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user