Abort during chunked prefill + PD peer-liveness abort (#28086)
This commit is contained in:
@@ -406,16 +406,55 @@ class SchedulerDisaggregationPrefillMixin:
|
||||
if room is not None and room in kv_mgr.transfer_infos:
|
||||
prefetch(room)
|
||||
|
||||
def resolve_waiting_queue_bootstrap(self: Scheduler) -> None:
|
||||
"""Resolve bootstrap status for waiting prefill requests before admission.
|
||||
|
||||
Covers the window between leaving the bootstrap queue and being admitted
|
||||
into a running batch: aborts requests whose decode peer died, and
|
||||
finalizes optimistic requests whose bootstrap completed so they skip
|
||||
the post-forward bootstrap check.
|
||||
"""
|
||||
candidates = [req for req in self.waiting_queue if not is_aborted(req)]
|
||||
if not candidates:
|
||||
return
|
||||
polls = poll_and_all_reduce_attn_cp_tp_group(
|
||||
[req.disagg_kv_sender for req in candidates],
|
||||
self.attn_cp_cpu_group,
|
||||
self.attn_tp_cpu_group,
|
||||
)
|
||||
failed = set()
|
||||
for req, poll in zip(candidates, polls):
|
||||
if poll == KVPoll.Failed:
|
||||
self.handle_bootstrap_failure(req)
|
||||
failed.add(req)
|
||||
elif (
|
||||
poll == KVPoll.WaitingForInput
|
||||
and req.pending_bootstrap
|
||||
and not should_force_retry(req)
|
||||
):
|
||||
# Optimistic requests reserved a metadata buffer when popped, so
|
||||
# finalize cannot fail here; if it ever does, the request stays
|
||||
# pending and the post-forward check resolves it.
|
||||
self.disagg_prefill_bootstrap_queue.finalize_bootstrap(req)
|
||||
if failed:
|
||||
self.waiting_queue = [
|
||||
req for req in self.waiting_queue if req not in failed
|
||||
]
|
||||
|
||||
@scheduler_nvtx_method("scheduler.get_next_batch_to_run")
|
||||
def get_next_disagg_prefill_batch_to_run(
|
||||
self: Scheduler,
|
||||
) -> Optional[ScheduleBatch]:
|
||||
self.process_pending_chunked_abort()
|
||||
|
||||
# HACK (byronhsu): reset the batch_is_full flag because we never enter update_running_batch which resets it
|
||||
# Otherwise, it hangs under high concurrency
|
||||
self.running_batch.batch_is_full = False
|
||||
|
||||
self.process_prefill_chunk()
|
||||
|
||||
self.resolve_waiting_queue_bootstrap()
|
||||
|
||||
batch = self.get_new_batch_prefill()
|
||||
batch = self.dp_attn_adapter.maybe_prepare_mlp_sync_batch(batch)
|
||||
|
||||
|
||||
@@ -985,6 +985,7 @@ class Scheduler(
|
||||
elif self.chunked_prefill_size is not None and self.chunked_prefill_size <= 0:
|
||||
self.chunked_prefill_size = None
|
||||
self.chunked_req = None
|
||||
self._pending_chunked_abort_req = None
|
||||
self.is_mixed_chunk = (
|
||||
self.chunked_prefill_size is not None
|
||||
and self.server_args.enable_mixed_chunk
|
||||
@@ -2424,6 +2425,52 @@ class Scheduler(
|
||||
def stash_chunked_request(self, req: Req):
|
||||
maybe_cache_unfinished_req(req, self.tree_cache, chunked=True)
|
||||
|
||||
def process_pending_chunked_abort(self) -> None:
|
||||
"""Abort an in-flight chunked-prefill request once it is safe to do so.
|
||||
|
||||
``abort_request`` only records the target in ``_pending_chunked_abort_req``
|
||||
(tearing it down mid-iteration is unsafe). Clearing ``chunked_req`` here at
|
||||
the top of the scheduling step stops the next chunk from launching; the
|
||||
chunk already launched is drained when its result is resolved. Under overlap
|
||||
the result lands a step later, so the batch-result processors keep
|
||||
``inflight_middle_chunks`` accounting intact and skip the aborted chunk:
|
||||
``process_batch_result_disagg_prefill`` via its ``is_aborted`` drop, and
|
||||
``process_batch_result_prefill`` via its chunked branch (the finished req
|
||||
is excluded from streaming and its logprob offset is still accounted).
|
||||
Mirrors ``handle_bootstrap_failure``.
|
||||
"""
|
||||
req = self._pending_chunked_abort_req
|
||||
if req is None:
|
||||
return
|
||||
if self.chunked_req is not req:
|
||||
# Already past chunked prefill; the running-batch abort path handles
|
||||
# it. Drop the marker once the request is actually gone.
|
||||
if req.finished() or req.req_pool_idx is None:
|
||||
self._pending_chunked_abort_req = None
|
||||
return
|
||||
|
||||
prepare_abort(req, "Aborted")
|
||||
req.time_stats.trace_ctx.abort(abort_info={"reason": "Aborted"})
|
||||
req.to_finish = None
|
||||
if self.disaggregation_mode == DisaggregationMode.PREFILL:
|
||||
req.disagg_kv_sender.abort()
|
||||
maybe_release_metadata_buffer(
|
||||
req, self.req_to_metadata_buffer_idx_allocator
|
||||
)
|
||||
req.pending_bootstrap = False
|
||||
if self.enable_hicache_storage:
|
||||
self.tree_cache.release_aborted_request(req.rid)
|
||||
if (
|
||||
req.req_pool_idx is not None or self.tree_cache.supports_mamba()
|
||||
) and not req.kv_committed_freed:
|
||||
release_kv_cache(req, self.tree_cache, is_insert=False)
|
||||
|
||||
self.chunked_req = None
|
||||
self._chunked_req_scheduled_last_iter = False
|
||||
self._pending_chunked_abort_req = None
|
||||
self.ipc_channels.send_to_tokenizer.send_output(AbortReq(rid=req.rid), req)
|
||||
logger.debug(f"Abort chunked prefill request. {req.rid=}")
|
||||
|
||||
def _build_hisparse_decode_batch(self, reqs):
|
||||
"""Build a ScheduleBatch for hisparse requests transitioning from staging to decode."""
|
||||
device = self.device
|
||||
@@ -2467,6 +2514,8 @@ class Scheduler(
|
||||
|
||||
@scheduler_nvtx_method("scheduler.get_next_batch_to_run")
|
||||
def get_next_batch_to_run(self) -> Optional[ScheduleBatch]:
|
||||
self.process_pending_chunked_abort()
|
||||
|
||||
if self.enable_fpm:
|
||||
self._fpm_batch_t0 = time.monotonic()
|
||||
self._abort_on_waiting_timeout()
|
||||
@@ -3684,6 +3733,10 @@ class Scheduler(
|
||||
return RpcReqOutput(success, "" if not exec else str(exec))
|
||||
|
||||
def abort_request(self, recv_req: AbortReq):
|
||||
if (chunked_req := self.chunked_req) is not None:
|
||||
if recv_req.abort_all or chunked_req.rid.startswith(recv_req.rid):
|
||||
self._pending_chunked_abort_req = chunked_req
|
||||
|
||||
# todo hisparse, release resources for abort requests in hisparse coordinator
|
||||
# Delete requests in the waiting queue
|
||||
to_del = []
|
||||
|
||||
@@ -216,8 +216,12 @@ class SchedulerBatchResultProcessor:
|
||||
logprob_pt = 0
|
||||
|
||||
for i, (req, next_token_id) in enumerate(zip(batch.reqs, next_token_ids)):
|
||||
if req.finished() or req.is_retracted:
|
||||
# decode req in mixed batch or retracted req
|
||||
if (
|
||||
req.finished() and req.inflight_middle_chunks <= 0
|
||||
) or req.is_retracted:
|
||||
# Decode req in a mixed batch, or a retracted req. Keep an
|
||||
# aborted middle chunk in the chunked branch long enough to
|
||||
# drain its accounting without streaming it.
|
||||
continue
|
||||
|
||||
if req.inflight_middle_chunks <= 0:
|
||||
|
||||
@@ -0,0 +1,134 @@
|
||||
import threading
|
||||
import time
|
||||
import unittest
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
import requests
|
||||
|
||||
from sglang.srt.utils import kill_process_tree
|
||||
from sglang.test.test_utils import (
|
||||
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
|
||||
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
DEFAULT_URL_FOR_TEST,
|
||||
CustomTestCase,
|
||||
popen_launch_server,
|
||||
)
|
||||
|
||||
CHUNKED_PREFILL_SIZE = 64
|
||||
LONG_PROMPT = (
|
||||
"The quick brown fox jumps over the lazy dog. "
|
||||
"Pack my box with five dozen liquor jugs. "
|
||||
"Sphinx of black quartz, judge my vow. "
|
||||
) * 900
|
||||
|
||||
|
||||
def _decode_response(response: requests.Response) -> Any:
|
||||
try:
|
||||
return response.json()
|
||||
except ValueError:
|
||||
return response.text
|
||||
|
||||
|
||||
def _is_abort_result(status_code: int, body: Any) -> bool:
|
||||
if status_code == 200:
|
||||
reason = (
|
||||
body.get("meta_info", {}).get("finish_reason", {})
|
||||
if isinstance(body, dict)
|
||||
else {}
|
||||
)
|
||||
return isinstance(reason, dict) and reason.get("type") == "abort"
|
||||
|
||||
if status_code not in (500, 503):
|
||||
return False
|
||||
|
||||
text = body if isinstance(body, str) else str(body)
|
||||
return "abort" in text.lower()
|
||||
|
||||
|
||||
class TestChunkedPrefillAbortE2E(CustomTestCase):
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||
cls.process = popen_launch_server(
|
||||
DEFAULT_SMALL_MODEL_NAME_FOR_TEST,
|
||||
cls.base_url,
|
||||
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
|
||||
other_args=[
|
||||
"--disable-cuda-graph",
|
||||
"--chunked-prefill-size",
|
||||
str(CHUNKED_PREFILL_SIZE),
|
||||
"--max-running-requests",
|
||||
"4",
|
||||
],
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
if cls.process:
|
||||
kill_process_tree(cls.process.pid)
|
||||
|
||||
def test_abort_mid_chunked_prefill_by_rid(self):
|
||||
rid = f"chunked-prefill-abort-{uuid.uuid4().hex}"
|
||||
result: dict[str, Any] = {}
|
||||
|
||||
def run_generate():
|
||||
try:
|
||||
response = requests.post(
|
||||
self.base_url + "/generate",
|
||||
json={
|
||||
"rid": rid,
|
||||
"text": f"{rid}\n{LONG_PROMPT}",
|
||||
"sampling_params": {
|
||||
"temperature": 0,
|
||||
"max_new_tokens": 4096,
|
||||
"ignore_eos": True,
|
||||
},
|
||||
},
|
||||
timeout=180,
|
||||
)
|
||||
result["status_code"] = response.status_code
|
||||
result["body"] = _decode_response(response)
|
||||
except requests.RequestException as exc:
|
||||
result["exception"] = repr(exc)
|
||||
|
||||
thread = threading.Thread(target=run_generate)
|
||||
thread.start()
|
||||
|
||||
time.sleep(0.5)
|
||||
abort_deadline = time.monotonic() + 8
|
||||
while thread.is_alive() and time.monotonic() < abort_deadline:
|
||||
requests.post(
|
||||
self.base_url + "/abort_request",
|
||||
json={"rid": rid, "abort_all": False},
|
||||
timeout=10,
|
||||
)
|
||||
time.sleep(0.2)
|
||||
|
||||
thread.join(timeout=60)
|
||||
self.assertFalse(thread.is_alive(), "Chunked-prefill abort request hung")
|
||||
self.assertNotIn("exception", result, result.get("exception"))
|
||||
self.assertTrue(
|
||||
_is_abort_result(result["status_code"], result["body"]),
|
||||
f"Expected chunked-prefill request to abort, got {result}",
|
||||
)
|
||||
|
||||
health = requests.get(self.base_url + "/health", timeout=10)
|
||||
self.assertEqual(health.status_code, 200, health.text)
|
||||
|
||||
follow_up = requests.post(
|
||||
self.base_url + "/generate",
|
||||
json={
|
||||
"rid": f"chunked-prefill-after-abort-{uuid.uuid4().hex}",
|
||||
"text": "The capital of France is",
|
||||
"sampling_params": {"temperature": 0, "max_new_tokens": 4},
|
||||
},
|
||||
timeout=60,
|
||||
)
|
||||
follow_up_body = _decode_response(follow_up)
|
||||
self.assertEqual(follow_up.status_code, 200, follow_up_body)
|
||||
self.assertFalse(_is_abort_result(follow_up.status_code, follow_up_body))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
+109
@@ -0,0 +1,109 @@
|
||||
import threading
|
||||
import time
|
||||
import unittest
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
import requests
|
||||
|
||||
from sglang.test.server_fixtures.disaggregation_fixture import (
|
||||
PDDisaggregationServerBase,
|
||||
)
|
||||
from sglang.test.test_utils import DEFAULT_SMALL_MODEL_NAME_FOR_TEST
|
||||
|
||||
PD_EXTRA_ARGS = ["--max-running-requests", "4", "--chunked-prefill-size", "64"]
|
||||
LONG_PROMPT = (
|
||||
"The quick brown fox jumps over the lazy dog. "
|
||||
"Pack my box with five dozen liquor jugs. "
|
||||
"Sphinx of black quartz, judge my vow. "
|
||||
) * 900
|
||||
|
||||
|
||||
def _decode_response(response: requests.Response) -> Any:
|
||||
try:
|
||||
return response.json()
|
||||
except ValueError:
|
||||
return response.text
|
||||
|
||||
|
||||
def _is_abort_result(status_code: int, body: Any) -> bool:
|
||||
if status_code == 200:
|
||||
reason = (
|
||||
body.get("meta_info", {}).get("finish_reason", {})
|
||||
if isinstance(body, dict)
|
||||
else {}
|
||||
)
|
||||
return isinstance(reason, dict) and reason.get("type") == "abort"
|
||||
|
||||
if status_code not in (500, 503):
|
||||
return False
|
||||
|
||||
text = body if isinstance(body, str) else str(body)
|
||||
return "abort" in text.lower()
|
||||
|
||||
|
||||
class TestDisaggChunkedPrefillAbort(PDDisaggregationServerBase):
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
super().setUpClass()
|
||||
cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST
|
||||
cls.extra_prefill_args = PD_EXTRA_ARGS
|
||||
cls.extra_decode_args = PD_EXTRA_ARGS
|
||||
cls.launch_all()
|
||||
|
||||
def _post_abort(self, rid: str):
|
||||
for url in (self.prefill_url, self.decode_url):
|
||||
requests.post(
|
||||
url + "/abort_request",
|
||||
json={"rid": rid, "abort_all": False},
|
||||
timeout=10,
|
||||
)
|
||||
|
||||
def test_abort_mid_chunked_prefill_by_rid(self):
|
||||
rid = f"pd-chunked-prefill-abort-{uuid.uuid4().hex}"
|
||||
result: dict[str, Any] = {}
|
||||
|
||||
def run_generate():
|
||||
try:
|
||||
response = requests.post(
|
||||
self.lb_url + "/generate",
|
||||
json={
|
||||
"rid": rid,
|
||||
"text": f"{rid}\n{LONG_PROMPT}",
|
||||
"sampling_params": {
|
||||
"temperature": 0,
|
||||
"max_new_tokens": 4096,
|
||||
"ignore_eos": True,
|
||||
},
|
||||
},
|
||||
timeout=180,
|
||||
)
|
||||
result["status_code"] = response.status_code
|
||||
result["body"] = _decode_response(response)
|
||||
except requests.RequestException as exc:
|
||||
result["exception"] = repr(exc)
|
||||
|
||||
thread = threading.Thread(target=run_generate)
|
||||
thread.start()
|
||||
|
||||
time.sleep(1.0)
|
||||
abort_deadline = time.monotonic() + 8
|
||||
while thread.is_alive() and time.monotonic() < abort_deadline:
|
||||
self._post_abort(rid)
|
||||
time.sleep(0.2)
|
||||
|
||||
thread.join(timeout=60)
|
||||
self.assertFalse(thread.is_alive(), "Chunked-prefill abort request hung")
|
||||
self.assertNotIn("exception", result, result.get("exception"))
|
||||
self.assertTrue(
|
||||
_is_abort_result(result["status_code"], result["body"]),
|
||||
f"Expected chunked-prefill request to abort, got {result}",
|
||||
)
|
||||
|
||||
for url in (self.lb_url, self.prefill_url, self.decode_url):
|
||||
health = requests.get(url + "/health", timeout=10)
|
||||
self.assertEqual(health.status_code, 200, health.text)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,316 @@
|
||||
import json
|
||||
import threading
|
||||
import time
|
||||
import unittest
|
||||
import uuid
|
||||
from concurrent.futures import ThreadPoolExecutor, TimeoutError, as_completed
|
||||
from typing import Any
|
||||
|
||||
import requests
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.test.server_fixtures.disaggregation_fixture import (
|
||||
PDDisaggregationServerBase,
|
||||
)
|
||||
from sglang.test.test_utils import DEFAULT_MODEL_NAME_FOR_TEST
|
||||
|
||||
FAILURE_PROB = 0.12
|
||||
PD_EXTRA_ARGS = [
|
||||
"--max-running-requests",
|
||||
"4",
|
||||
"--chunked-prefill-size",
|
||||
"128",
|
||||
]
|
||||
PD_LONG_PROMPT = (
|
||||
"The quick brown fox jumps over the lazy dog. "
|
||||
"Pack my box with five dozen liquor jugs. "
|
||||
"Sphinx of black quartz, judge my vow. "
|
||||
) * 240
|
||||
|
||||
|
||||
def _body_text(body: Any) -> str:
|
||||
if isinstance(body, str):
|
||||
return body
|
||||
return json.dumps(body, sort_keys=True)
|
||||
|
||||
|
||||
def _decode_response(response: requests.Response) -> Any:
|
||||
try:
|
||||
return response.json()
|
||||
except ValueError:
|
||||
return response.text
|
||||
|
||||
|
||||
def _finish_reason(body: Any) -> dict[str, Any]:
|
||||
if not isinstance(body, dict):
|
||||
return {}
|
||||
reason = body.get("meta_info", {}).get("finish_reason", {})
|
||||
return reason if isinstance(reason, dict) else {}
|
||||
|
||||
|
||||
def _is_abort_result(status_code: int, body: Any) -> bool:
|
||||
if status_code == 200:
|
||||
return _finish_reason(body).get("type") == "abort"
|
||||
|
||||
if status_code not in (500, 503):
|
||||
return False
|
||||
|
||||
text = _body_text(body).lower()
|
||||
return any(
|
||||
marker in text
|
||||
for marker in (
|
||||
"pd peer failed",
|
||||
"prefill bootstrap failed",
|
||||
"abort",
|
||||
"aborted",
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _is_success_result(status_code: int, body: Any) -> bool:
|
||||
return (
|
||||
status_code == 200
|
||||
and isinstance(body, dict)
|
||||
and "text" in body
|
||||
and _finish_reason(body).get("type") != "abort"
|
||||
)
|
||||
|
||||
|
||||
class TestDisaggregationPeerLivenessAbort(PDDisaggregationServerBase):
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
super().setUpClass()
|
||||
cls._disagg_failure_ctx = envs.SGLANG_TEST_DISAGG_FAILURE_PROB.override(
|
||||
FAILURE_PROB
|
||||
)
|
||||
cls._disagg_failure_ctx.__enter__()
|
||||
|
||||
cls.model = DEFAULT_MODEL_NAME_FOR_TEST
|
||||
cls.extra_prefill_args = PD_EXTRA_ARGS
|
||||
cls.extra_decode_args = PD_EXTRA_ARGS
|
||||
cls.launch_all()
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
try:
|
||||
super().tearDownClass()
|
||||
finally:
|
||||
if cls._disagg_failure_ctx:
|
||||
cls._disagg_failure_ctx.__exit__(None, None, None)
|
||||
|
||||
def _post_generate(self, rid: str, max_new_tokens: int = 32) -> dict[str, Any]:
|
||||
response = requests.post(
|
||||
self.lb_url + "/generate",
|
||||
json={
|
||||
"rid": rid,
|
||||
"text": f"{rid}\n{PD_LONG_PROMPT}",
|
||||
"sampling_params": {
|
||||
"temperature": 0,
|
||||
"max_new_tokens": max_new_tokens,
|
||||
"ignore_eos": True,
|
||||
},
|
||||
},
|
||||
timeout=75,
|
||||
)
|
||||
return {
|
||||
"rid": rid,
|
||||
"status_code": response.status_code,
|
||||
"body": _decode_response(response),
|
||||
}
|
||||
|
||||
def _assert_servers_healthy(self):
|
||||
for url in (self.lb_url, self.prefill_url, self.decode_url):
|
||||
response = requests.get(url + "/health", timeout=10)
|
||||
self.assertEqual(response.status_code, 200, response.text)
|
||||
|
||||
def _load_count(self, url: str) -> int:
|
||||
response = requests.get(
|
||||
url + "/v1/loads?include=core,disagg,queues",
|
||||
timeout=10,
|
||||
)
|
||||
response.raise_for_status()
|
||||
total = 0
|
||||
for load in response.json()["loads"]:
|
||||
total += int(load.get("num_running_reqs", 0))
|
||||
total += int(load.get("num_waiting_reqs", 0))
|
||||
disagg = load.get("disaggregation", {})
|
||||
total += int(disagg.get("prefill_bootstrap_queue_reqs", 0))
|
||||
total += int(disagg.get("prefill_inflight_queue_reqs", 0))
|
||||
total += int(disagg.get("decode_prealloc_queue_reqs", 0))
|
||||
total += int(disagg.get("decode_transfer_queue_reqs", 0))
|
||||
total += int(disagg.get("decode_retracted_queue_reqs", 0))
|
||||
queues = load.get("queues", {})
|
||||
total += int(queues.get("waiting", 0))
|
||||
total += int(queues.get("grammar", 0))
|
||||
total += int(queues.get("paused", 0))
|
||||
total += int(queues.get("retracted", 0))
|
||||
return total
|
||||
|
||||
def _assert_eventually_idle(self):
|
||||
deadline = time.monotonic() + 30
|
||||
last_counts = {}
|
||||
while time.monotonic() < deadline:
|
||||
last_counts = {
|
||||
self.prefill_url: self._load_count(self.prefill_url),
|
||||
self.decode_url: self._load_count(self.decode_url),
|
||||
}
|
||||
if all(count == 0 for count in last_counts.values()):
|
||||
return
|
||||
time.sleep(1)
|
||||
|
||||
self.fail(f"PD servers still have request load after cleanup: {last_counts}")
|
||||
|
||||
def _assert_one_successful_generate_with_retries(self):
|
||||
last_result = None
|
||||
for _ in range(20):
|
||||
rid = f"pd-peer-liveness-health-{uuid.uuid4().hex}"
|
||||
last_result = self._post_generate(rid, max_new_tokens=1)
|
||||
if _is_success_result(last_result["status_code"], last_result["body"]):
|
||||
return
|
||||
time.sleep(0.2)
|
||||
|
||||
self.fail(
|
||||
"Server stayed HTTP-healthy but did not complete a follow-up generate "
|
||||
f"under failure injection. Last result: {last_result}"
|
||||
)
|
||||
|
||||
def test_peer_failures_abort_without_hanging(self):
|
||||
num_requests = 36
|
||||
start = time.monotonic()
|
||||
results: list[dict[str, Any]] = []
|
||||
|
||||
with ThreadPoolExecutor(max_workers=18) as executor:
|
||||
futures = [
|
||||
executor.submit(
|
||||
self._post_generate,
|
||||
f"pd-peer-liveness-{i}-{uuid.uuid4().hex}",
|
||||
)
|
||||
for i in range(num_requests)
|
||||
]
|
||||
|
||||
try:
|
||||
for future in as_completed(futures, timeout=120):
|
||||
results.append(future.result())
|
||||
except TimeoutError:
|
||||
unfinished = sum(not future.done() for future in futures)
|
||||
self.fail(f"{unfinished} generate requests did not return in time")
|
||||
|
||||
elapsed = time.monotonic() - start
|
||||
self.assertLess(elapsed, 120)
|
||||
self.assertEqual(len(results), num_requests)
|
||||
|
||||
aborts = [
|
||||
result
|
||||
for result in results
|
||||
if _is_abort_result(result["status_code"], result["body"])
|
||||
]
|
||||
successes = [
|
||||
result
|
||||
for result in results
|
||||
if _is_success_result(result["status_code"], result["body"])
|
||||
]
|
||||
unexpected = [
|
||||
result
|
||||
for result in results
|
||||
if result not in aborts and result not in successes
|
||||
]
|
||||
|
||||
self.assertFalse(
|
||||
unexpected,
|
||||
"Expected only clean aborts or successful completions. "
|
||||
f"Unexpected results: {unexpected[:3]}",
|
||||
)
|
||||
self.assertGreater(
|
||||
len(aborts),
|
||||
0,
|
||||
"Injected KVPoll.Failed should abort at least one request",
|
||||
)
|
||||
|
||||
self._assert_servers_healthy()
|
||||
self._assert_eventually_idle()
|
||||
self._assert_one_successful_generate_with_retries()
|
||||
|
||||
|
||||
class TestDisaggregationPeerLivenessCleanRecovery(PDDisaggregationServerBase):
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
super().setUpClass()
|
||||
cls.model = DEFAULT_MODEL_NAME_FOR_TEST
|
||||
cls.extra_prefill_args = PD_EXTRA_ARGS
|
||||
cls.extra_decode_args = PD_EXTRA_ARGS
|
||||
cls.launch_all()
|
||||
|
||||
def _post_abort_to_engines(self, rid: str):
|
||||
for url in (self.prefill_url, self.decode_url):
|
||||
response = requests.post(
|
||||
url + "/abort_request",
|
||||
json={"rid": rid, "abort_all": False},
|
||||
timeout=10,
|
||||
)
|
||||
self.assertEqual(response.status_code, 200, response.text)
|
||||
|
||||
def _assert_servers_healthy(self):
|
||||
for url in (self.lb_url, self.prefill_url, self.decode_url):
|
||||
response = requests.get(url + "/health", timeout=10)
|
||||
self.assertEqual(response.status_code, 200, response.text)
|
||||
|
||||
def test_clean_generate_and_explicit_abort_by_rid(self):
|
||||
clean_response = requests.post(
|
||||
self.lb_url + "/generate",
|
||||
json={
|
||||
"rid": f"pd-clean-generate-{uuid.uuid4().hex}",
|
||||
"text": "The capital of France is",
|
||||
"sampling_params": {"temperature": 0, "max_new_tokens": 4},
|
||||
},
|
||||
timeout=60,
|
||||
)
|
||||
clean_body = _decode_response(clean_response)
|
||||
self.assertTrue(
|
||||
_is_success_result(clean_response.status_code, clean_body), clean_body
|
||||
)
|
||||
|
||||
rid = f"pd-explicit-abort-{uuid.uuid4().hex}"
|
||||
result: dict[str, Any] = {}
|
||||
|
||||
def run_generate():
|
||||
try:
|
||||
response = requests.post(
|
||||
self.lb_url + "/generate",
|
||||
json={
|
||||
"rid": rid,
|
||||
"text": f"{rid}\n{PD_LONG_PROMPT}",
|
||||
"sampling_params": {
|
||||
"temperature": 0,
|
||||
"max_new_tokens": 4096,
|
||||
"ignore_eos": True,
|
||||
},
|
||||
},
|
||||
timeout=180,
|
||||
)
|
||||
result["status_code"] = response.status_code
|
||||
result["body"] = _decode_response(response)
|
||||
except requests.RequestException as exc:
|
||||
result["exception"] = repr(exc)
|
||||
|
||||
thread = threading.Thread(target=run_generate)
|
||||
thread.start()
|
||||
|
||||
time.sleep(0.25)
|
||||
abort_deadline = time.monotonic() + 8
|
||||
while thread.is_alive() and time.monotonic() < abort_deadline:
|
||||
self._post_abort_to_engines(rid)
|
||||
time.sleep(0.25)
|
||||
|
||||
thread.join(timeout=60)
|
||||
self.assertFalse(thread.is_alive(), "Explicitly aborted PD request hung")
|
||||
self.assertNotIn("exception", result, result.get("exception"))
|
||||
self.assertTrue(
|
||||
_is_abort_result(result["status_code"], result["body"]),
|
||||
f"Expected explicit abort result, got {result}",
|
||||
)
|
||||
|
||||
self._assert_servers_healthy()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user