[misc] Share bench HTTP-client base-URL resolution with IPv6-compatible formatting (#28598)
This commit is contained in:
@@ -48,7 +48,7 @@ from sglang.benchmark.utils import (
|
|||||||
set_ulimit,
|
set_ulimit,
|
||||||
)
|
)
|
||||||
from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST
|
from sglang.srt.disaggregation.utils import FAKE_BOOTSTRAP_HOST
|
||||||
from sglang.srt.utils.network import NetworkAddress
|
from sglang.srt.utils.network import NetworkAddress, resolve_base_url
|
||||||
|
|
||||||
_ROUTING_KEY_HEADER = "X-SMG-Routing-Key"
|
_ROUTING_KEY_HEADER = "X-SMG-Routing-Key"
|
||||||
|
|
||||||
@@ -916,6 +916,22 @@ ASYNC_REQUEST_FUNCS = {
|
|||||||
"truss": async_request_truss,
|
"truss": async_request_truss,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# API path appended to the base URL per backend. gserver is special (bare
|
||||||
|
# host:port, no path) and is handled separately, so it is not listed here.
|
||||||
|
_BACKEND_API_PATHS = {
|
||||||
|
"sglang": "/generate",
|
||||||
|
"sglang-native": "/generate",
|
||||||
|
"sglang-oai": "/v1/completions",
|
||||||
|
"sglang-oai-chat": "/v1/chat/completions",
|
||||||
|
"sglang-embedding": "/v1/embeddings",
|
||||||
|
"vllm": "/v1/completions",
|
||||||
|
"vllm-chat": "/v1/chat/completions",
|
||||||
|
"lmdeploy": "/v1/completions",
|
||||||
|
"lmdeploy-chat": "/v1/chat/completions",
|
||||||
|
"trt": "/v2/models/ensemble/generate_stream",
|
||||||
|
"truss": "/v1/models/model:predict",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class BenchmarkMetrics:
|
class BenchmarkMetrics:
|
||||||
@@ -1924,59 +1940,23 @@ def run_benchmark(args_: argparse.Namespace):
|
|||||||
"truss": 8080,
|
"truss": 8080,
|
||||||
}.get(args.backend, 30000)
|
}.get(args.backend, 30000)
|
||||||
|
|
||||||
# Build base URL with proper IPv6 bracket wrapping (only when base_url is not provided)
|
# Base URL the client sends to: --base-url if given, else http://host:port
|
||||||
if not args.base_url:
|
# (IPv6-correct). NetworkAddress is also kept for gserver's host:port form.
|
||||||
_na = NetworkAddress(args.host, args.port)
|
base_url = resolve_base_url(args.base_url, args.host, args.port)
|
||||||
_host_base = _na.to_url()
|
_na = NetworkAddress(args.host, args.port)
|
||||||
else:
|
|
||||||
_na = None
|
|
||||||
_host_base = None
|
|
||||||
|
|
||||||
model_url = (
|
model_url = f"{base_url}/v1/models"
|
||||||
f"{args.base_url}/v1/models" if args.base_url else f"{_host_base}/v1/models"
|
|
||||||
)
|
|
||||||
|
|
||||||
if args.backend == "sglang-embedding":
|
if args.backend == "gserver":
|
||||||
api_url = (
|
# gRPC server takes a bare host:port, not an http URL.
|
||||||
f"{args.base_url}/v1/embeddings"
|
|
||||||
if args.base_url
|
|
||||||
else f"http://{args.host}:{args.port}/v1/embeddings"
|
|
||||||
)
|
|
||||||
elif args.backend in ["sglang", "sglang-native"]:
|
|
||||||
api_url = (
|
|
||||||
f"{args.base_url}/generate" if args.base_url else f"{_host_base}/generate"
|
|
||||||
)
|
|
||||||
elif args.backend in ["sglang-oai", "vllm", "lmdeploy"]:
|
|
||||||
api_url = (
|
|
||||||
f"{args.base_url}/v1/completions"
|
|
||||||
if args.base_url
|
|
||||||
else f"{_host_base}/v1/completions"
|
|
||||||
)
|
|
||||||
elif args.backend in ["sglang-oai-chat", "vllm-chat", "lmdeploy-chat"]:
|
|
||||||
api_url = (
|
|
||||||
f"{args.base_url}/v1/chat/completions"
|
|
||||||
if args.base_url
|
|
||||||
else f"{_host_base}/v1/chat/completions"
|
|
||||||
)
|
|
||||||
elif args.backend == "trt":
|
|
||||||
api_url = (
|
|
||||||
f"{args.base_url}/v2/models/ensemble/generate_stream"
|
|
||||||
if args.base_url
|
|
||||||
else f"{_host_base}/v2/models/ensemble/generate_stream"
|
|
||||||
)
|
|
||||||
if args.model is None:
|
|
||||||
print("Please provide a model using `--model` when using `trt` backend.")
|
|
||||||
sys.exit(1)
|
|
||||||
elif args.backend == "gserver":
|
|
||||||
api_url = args.base_url if args.base_url else _na.to_host_port_str()
|
api_url = args.base_url if args.base_url else _na.to_host_port_str()
|
||||||
args.model = args.model or "default"
|
args.model = args.model or "default"
|
||||||
elif args.backend == "truss":
|
else:
|
||||||
api_url = (
|
api_url = f"{base_url}{_BACKEND_API_PATHS[args.backend]}"
|
||||||
f"{args.base_url}/v1/models/model:predict"
|
|
||||||
if args.base_url
|
if args.backend == "trt" and args.model is None:
|
||||||
else f"{_host_base}/v1/models/model:predict"
|
print("Please provide a model using `--model` when using `trt` backend.")
|
||||||
)
|
sys.exit(1)
|
||||||
base_url = _host_base if args.base_url is None else args.base_url
|
|
||||||
|
|
||||||
# Wait for server to be ready
|
# Wait for server to be ready
|
||||||
if args.ready_check_timeout_sec > 0:
|
if args.ready_check_timeout_sec > 0:
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ import requests
|
|||||||
from sglang.srt.entrypoints.http_server import launch_server
|
from sglang.srt.entrypoints.http_server import launch_server
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.utils import kill_process_tree
|
from sglang.srt.utils import kill_process_tree
|
||||||
|
from sglang.srt.utils.network import resolve_base_url
|
||||||
|
|
||||||
DEFAULT_TIMEOUT = 600
|
DEFAULT_TIMEOUT = 600
|
||||||
|
|
||||||
@@ -47,7 +48,7 @@ def _launch_server_target(launch_server_func: Callable, server_args: ServerArgs)
|
|||||||
|
|
||||||
|
|
||||||
def launch_or_reuse_server(launch_server_func: Callable, server_args: ServerArgs):
|
def launch_or_reuse_server(launch_server_func: Callable, server_args: ServerArgs):
|
||||||
base_url = f"http://{server_args.host}:{server_args.port}"
|
base_url = resolve_base_url("", server_args.host, server_args.port)
|
||||||
|
|
||||||
# Reuse an already-running server instead of forking a second one onto the
|
# Reuse an already-running server instead of forking a second one onto the
|
||||||
# occupied port, where it would orphan, compete for the GPU, and OOM.
|
# occupied port, where it would orphan, compete for the GPU, and OOM.
|
||||||
|
|||||||
@@ -543,3 +543,11 @@ class NetworkAddress:
|
|||||||
|
|
||||||
def __repr__(self) -> str:
|
def __repr__(self) -> str:
|
||||||
return f"NetworkAddress({self.host!r}, {self.port})"
|
return f"NetworkAddress({self.host!r}, {self.port})"
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_base_url(base_url: str, host: str, port: int) -> str:
|
||||||
|
"""Base URL a client sends to: ``base_url`` if set, else ``http://host:port``
|
||||||
|
(IPv6-correct via :class:`NetworkAddress`)."""
|
||||||
|
if base_url:
|
||||||
|
return base_url
|
||||||
|
return NetworkAddress(host, port).to_url()
|
||||||
|
|||||||
@@ -18,12 +18,14 @@ import requests
|
|||||||
import tabulate
|
import tabulate
|
||||||
|
|
||||||
from sglang.profiler import run_profile
|
from sglang.profiler import run_profile
|
||||||
|
from sglang.srt.utils.network import resolve_base_url
|
||||||
|
|
||||||
|
|
||||||
@dataclasses.dataclass
|
@dataclasses.dataclass
|
||||||
class BenchArgs:
|
class BenchArgs:
|
||||||
host: str = "localhost"
|
host: str = "localhost"
|
||||||
port: int = 30000
|
port: int = 30000
|
||||||
|
base_url: str = ""
|
||||||
batch_size: int = 1
|
batch_size: int = 1
|
||||||
different_prompts: bool = False
|
different_prompts: bool = False
|
||||||
random_input_len: Optional[int] = None
|
random_input_len: Optional[int] = None
|
||||||
@@ -51,6 +53,12 @@ class BenchArgs:
|
|||||||
def add_cli_args(parser: argparse.ArgumentParser):
|
def add_cli_args(parser: argparse.ArgumentParser):
|
||||||
parser.add_argument("--host", type=str, default=BenchArgs.host)
|
parser.add_argument("--host", type=str, default=BenchArgs.host)
|
||||||
parser.add_argument("--port", type=int, default=BenchArgs.port)
|
parser.add_argument("--port", type=int, default=BenchArgs.port)
|
||||||
|
parser.add_argument(
|
||||||
|
"--base-url",
|
||||||
|
type=str,
|
||||||
|
default=BenchArgs.base_url,
|
||||||
|
help="Server base url. Overrides --host/--port when set.",
|
||||||
|
)
|
||||||
parser.add_argument("--batch-size", type=int, default=BenchArgs.batch_size)
|
parser.add_argument("--batch-size", type=int, default=BenchArgs.batch_size)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--different-prompts",
|
"--different-prompts",
|
||||||
@@ -110,7 +118,7 @@ def send_one_prompt(
|
|||||||
label: Optional[str] = None,
|
label: Optional[str] = None,
|
||||||
print_output: bool = True,
|
print_output: bool = True,
|
||||||
):
|
):
|
||||||
base_url = f"http://{args.host}:{args.port}"
|
base_url = resolve_base_url(args.base_url, args.host, args.port)
|
||||||
|
|
||||||
# Construct the input
|
# Construct the input
|
||||||
if args.random_input_len is not None:
|
if args.random_input_len is not None:
|
||||||
|
|||||||
Reference in New Issue
Block a user