Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
51 changes: 50 additions & 1 deletion benchmarks/backend_request_func.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,13 +28,28 @@
import uuid
from dataclasses import dataclass, field
from typing import Optional
from urllib.parse import urlsplit, urlunsplit

import aiohttp
from tqdm.asyncio import tqdm

AIOHTTP_TIMEOUT = aiohttp.ClientTimeout(total=int(os.environ.get("AIOHTTP_TIMEOUT", 6 * 60 * 60)))


def _get_control_url(api_url: str, endpoint: str) -> str:
parsed = urlsplit(api_url)
return urlunsplit((parsed.scheme, parsed.netloc, endpoint, "", ""))


async def _close_session(session: aiohttp.ClientSession, api_url: str, session_id: str) -> None:
url = _get_control_url(api_url, "/close_session")
async with session.post(url=url, json={"session_id": session_id}) as response:
if response.status != 200:
raise RuntimeError(
f"Failed to close session {session_id}: {response.status} {await response.text()}"
)


@dataclass
class RequestFuncInput:
"""Input for requesting LLMs via API"""
Expand Down Expand Up @@ -64,6 +79,8 @@ class RequestFuncInput:
stream: bool = True
session_id: Optional[str] = None
turn_idx: Optional[int] = None
enable_session_control: bool = False
session_control_url: Optional[str] = None


@dataclass
Expand Down Expand Up @@ -411,6 +428,9 @@ async def async_request_eb_openai_chat_completions(
if request_func_input.ignore_eos:
payload["ignore_eos"] = request_func_input.ignore_eos

if request_func_input.session_id:
payload["session_id"] = request_func_input.session_id

headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {os.environ.get('OPENAI_API_KEY')}",
Expand Down Expand Up @@ -727,7 +747,7 @@ async def simple_tool_call(model_output, tool_url: str, timeout=60):
return str(e), False, tool_name, tool_id


async def async_request_eb_openai_chat_completions_multi_turn(
async def _async_request_eb_openai_chat_completions_multi_turn(
request_func_input: RequestFuncInput,
pbar: Optional[tqdm] = None,
):
Expand Down Expand Up @@ -1024,6 +1044,35 @@ async def async_request_eb_openai_chat_completions_multi_turn(
return outputs, metrics


async def async_request_eb_openai_chat_completions_multi_turn(
request_func_input: RequestFuncInput,
pbar: Optional[tqdm] = None,
):
if not request_func_input.enable_session_control:
return await _async_request_eb_openai_chat_completions_multi_turn(request_func_input, pbar)

request_func_input.session_id = request_func_input.session_id or uuid.uuid4().hex
result = None
try:
result = await _async_request_eb_openai_chat_completions_multi_turn(request_func_input, pbar)
finally:
try:
async with aiohttp.ClientSession(trust_env=True, timeout=AIOHTTP_TIMEOUT) as session:
await _close_session(
session,
request_func_input.session_control_url or request_func_input.api_url,
request_func_input.session_id,
)
except Exception:
close_error = traceback.format_exc()
logging.error("Failed to close session %s: %s", request_func_input.session_id, close_error)
if result and result[0]:
result[0][-1].success = False
result[0][-1].error += close_error

return result


async def async_request_eb_openai_completions(
request_func_input: RequestFuncInput,
pbar: Optional[tqdm] = None,
Expand Down
23 changes: 23 additions & 0 deletions benchmarks/benchmark_serving.py
Original file line number Diff line number Diff line change
Expand Up @@ -408,6 +408,8 @@ async def benchmark(
lora_modules: Optional[Iterable[str]],
extra_body: Optional[dict],
ip_list: Optional[list[str]] = None,
enable_session_control: bool = False,
session_control_url: Optional[str] = None,
):
"""Benchmarks an API endpoint using a given set of sample inputs and returns"""
if backend in ASYNC_REQUEST_FUNCS:
Expand Down Expand Up @@ -467,6 +469,8 @@ async def benchmark(
stream=args.stream,
session_id=f"warmup-{uuid.uuid4().hex}",
turn_idx=0,
enable_session_control=enable_session_control,
session_control_url=session_control_url,
)

if args.warmup:
Expand Down Expand Up @@ -581,6 +585,8 @@ async def limited_request_func(request_func_input, pbar):
tokenizer_model=args.tokenizer_model,
tokenizer_path=args.tokenizer_path,
stream=args.stream,
enable_session_control=enable_session_control,
session_control_url=session_control_url,
)
tasks.append(asyncio.create_task(limited_request_func(request_func_input=request_func_input, pbar=pbar)))

Expand Down Expand Up @@ -670,6 +676,8 @@ async def limited_request_func_per_ip(req_input, semaphore, pbar):
tokenizer_model=args.tokenizer_model,
tokenizer_path=args.tokenizer_path,
stream=args.stream,
enable_session_control=enable_session_control,
session_control_url=session_control_url,
)

tasks.append(asyncio.create_task(limited_request_func_per_ip(req_input, semaphore, pbar)))
Expand Down Expand Up @@ -1249,6 +1257,8 @@ def main(args: argparse.Namespace):
np.random.seed(args.seed)

backend = args.backend
if args.enable_session_control and not args.multi_turn:
raise ValueError("--enable-session-control requires --multi-turn.")
# 支持多轮对话方式请求,仅支持chat接口
if args.multi_turn:
backend = "openai-chat-multi-turn"
Expand Down Expand Up @@ -1362,6 +1372,8 @@ def main(args: argparse.Namespace):
lora_modules=args.lora_modules,
extra_body=sampling_params,
ip_list=ip_list,
enable_session_control=args.enable_session_control,
session_control_url=args.session_control_url,
)
)

Expand Down Expand Up @@ -1569,6 +1581,17 @@ def main(args: argparse.Namespace):
action="store_true",
help="按多轮对话方式请求",
)
parser.add_argument(
"--enable-session-control",
action="store_true",
help="为 SGLang 多轮请求携带 session id,并在 session 结束后调用 close_session 释放 radix cache",
)
parser.add_argument(
"--session-control-url",
type=str,
default=None,
help="SGLang session 控制接口的 base URL;通过不代理 close_session 的 router 发压时需要指定",
)
parser.add_argument(
"--no-warmup",
action="store_false",
Expand Down
Loading