Skip to content
Merged
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
69 changes: 38 additions & 31 deletions tensorrt_llm/commands/serve.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import json
import os
import signal # Added import
import socket
import subprocess # nosec B404
import sys
from pathlib import Path
Expand Down Expand Up @@ -176,37 +177,43 @@ def launch_server(

backend = llm_args["backend"]
model = llm_args["model"]
if backend == 'pytorch':
llm_args.pop("build_config", None)
llm = PyTorchLLM(**llm_args)
elif backend == '_autodeploy':
from tensorrt_llm._torch.auto_deploy import LLM as AutoDeployLLM

# AutoDeploy does not support build_config
llm_args.pop("build_config", None)
llm = AutoDeployLLM(**llm_args)
elif backend == 'tensorrt' or backend == 'trt':
llm_args.pop("backend")
llm = LLM(**llm_args)
else:
raise click.BadParameter(
f"{backend} is not a known backend, check help for available options.",
param_hint="backend")

server = OpenAIServer(llm=llm,
model=model,
tool_parser=tool_parser,
server_role=server_role,
metadata_server_cfg=metadata_server_cfg,
disagg_cluster_config=disagg_cluster_config,
multimodal_server_config=multimodal_server_config,
chat_template=chat_template)

# Optionally disable GC (default: not disabled)
if os.getenv("TRTLLM_SERVER_DISABLE_GC", "0") == "1":
gc.disable()

asyncio.run(server(host, port))
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
try:
s.bind((host, port))
except OSError as e:
raise RuntimeError(f"Failed to bind socket to {host}:{port}: {e}")

if backend == 'pytorch':
llm_args.pop("build_config", None)
llm = PyTorchLLM(**llm_args)
elif backend == '_autodeploy':
from tensorrt_llm._torch.auto_deploy import LLM as AutoDeployLLM

# AutoDeploy does not support build_config
llm_args.pop("build_config", None)
llm = AutoDeployLLM(**llm_args)
elif backend == 'tensorrt' or backend == 'trt':
llm_args.pop("backend")
llm = LLM(**llm_args)
else:
raise click.BadParameter(
f"{backend} is not a known backend, check help for available options.",
param_hint="backend")

server = OpenAIServer(llm=llm,
model=model,
tool_parser=tool_parser,
server_role=server_role,
metadata_server_cfg=metadata_server_cfg,
disagg_cluster_config=disagg_cluster_config,
multimodal_server_config=multimodal_server_config,
chat_template=chat_template)

# Optionally disable GC (default: not disabled)
if os.getenv("TRTLLM_SERVER_DISABLE_GC", "0") == "1":
gc.disable()

asyncio.run(server(host, port, sockets=[s]))


def launch_mm_encoder_server(
Expand Down
5 changes: 3 additions & 2 deletions tensorrt_llm/serve/openai_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import os
import re
import signal
import socket
import traceback
from collections import deque
from contextlib import asynccontextmanager
Expand Down Expand Up @@ -990,7 +991,7 @@ async def create_stream_response(generator, request: ResponsesRequest, sampling_
return JSONResponse(content={"detail": "None"})


async def __call__(self, host, port):
async def __call__(self, host, port, sockets: list[socket.socket] | None = None):
# Store the binding address for server registration
self.binding_addr = f"http://{host}:{port}"
self.host = host
Expand All @@ -1000,4 +1001,4 @@ async def __call__(self, host, port):
port=port,
log_level="info",
timeout_keep_alive=TIMEOUT_KEEP_ALIVE)
await uvicorn.Server(config).serve()
await uvicorn.Server(config).serve(sockets=sockets)
19 changes: 1 addition & 18 deletions tests/integration/defs/accuracy/test_disaggregated_serving.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,8 +13,8 @@
import pytest
import requests
import yaml
from defs.common import revise_disaggregated_server_config_urls_with_free_ports

from tensorrt_llm._utils import get_free_port
from tensorrt_llm.executor.result import GenerationResultBase
from tensorrt_llm.llmapi import CompletionOutput, RequestOutput, SamplingParams
from tensorrt_llm.llmapi.llm_args import LlmArgs
Expand Down Expand Up @@ -68,23 +68,6 @@ def __exit__(self, exc_type, exc_val, exc_tb):
return False


def revise_disaggregated_server_config_urls_with_free_ports(
disaggregated_server_config: Dict[str, Any]) -> Dict[str, Any]:
num_ctx_ports = len(disaggregated_server_config["context_servers"]["urls"])
num_gen_ports = len(
disaggregated_server_config["generation_servers"]["urls"])

disaggregated_server_config['port'] = get_free_port()
disaggregated_server_config["context_servers"]["urls"] = [
f"localhost:{get_free_port()}" for _ in range(num_ctx_ports)
]
disaggregated_server_config["generation_servers"]["urls"] = [
f"localhost:{get_free_port()}" for _ in range(num_gen_ports)
]

return disaggregated_server_config


@contextlib.contextmanager
def launch_disaggregated_llm(
disaggregated_server_config: Dict[str, Any],
Expand Down
42 changes: 42 additions & 0 deletions tests/integration/defs/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,13 +17,17 @@
import platform
import re
import socket
import tempfile
import time
from difflib import SequenceMatcher
from pathlib import Path
from typing import Any

import yaml
from packaging import version

from tensorrt_llm import LLM as LLM_torch
from tensorrt_llm._utils import get_free_port
from tensorrt_llm.executor.request import LoRARequest
from tensorrt_llm.lora_manager import LoraConfig
from tensorrt_llm.sampling_params import SamplingParams
Expand Down Expand Up @@ -1147,3 +1151,41 @@ def wait_for_server(host, port, timeout_seconds=180):
except (socket.error, ConnectionRefusedError, OSError):
time.sleep(2)
return False


def revise_disaggregated_server_config_urls_with_free_ports(
disaggregated_server_config: dict[str, Any]) -> dict[str, Any]:
# Revise serve port
disaggregated_server_config['port'] = get_free_port()

# Revise context and generation server urls
ctx_urls = disaggregated_server_config["context_servers"]["urls"]
gen_urls = disaggregated_server_config["generation_servers"]["urls"]
url_map = dict()
for url in set(ctx_urls + gen_urls):
url_map[url] = (url.split(':')[0], get_free_port())

for i, url in enumerate(ctx_urls):
disaggregated_server_config["context_servers"]["urls"][
i] = f"{url_map[url][0]}:{url_map[url][1]}"

for i, url in enumerate(gen_urls):
disaggregated_server_config["generation_servers"]["urls"][
i] = f"{url_map[url][0]}:{url_map[url][1]}"

return disaggregated_server_config


def revise_disagg_config_file_with_free_ports(disagg_config_file: str) -> str:
# Revise the config file to use free ports
new_config = None
with open(disagg_config_file, 'r') as f:
config = yaml.safe_load(f)
new_config = revise_disaggregated_server_config_urls_with_free_ports(
config)

temp_fd, new_config_file = tempfile.mkstemp(suffix='.yaml')
with os.fdopen(temp_fd, 'w') as f:
yaml.dump(new_config, f)

return new_config_file
38 changes: 22 additions & 16 deletions tests/integration/defs/disaggregated/test_disaggregated.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,12 +22,13 @@

import pytest
import yaml
from defs.common import wait_for_server
from defs.common import (revise_disagg_config_file_with_free_ports,
wait_for_server)
from defs.conftest import (get_sm_version, llm_models_root, skip_arm,
skip_no_hopper)
from defs.trt_test_alternative import check_call, check_output, popen

from tensorrt_llm._utils import mpi_disabled
from tensorrt_llm._utils import get_free_port, mpi_disabled
from tensorrt_llm.logger import logger


Expand Down Expand Up @@ -143,12 +144,12 @@ def validate_timing_metrics(perf_metrics_item, request_context=""):
return True


def get_disagg_server_url_from_cfg(config_file: str) -> str:
def get_disagg_server_url_from_cfg(config_file: str) -> tuple[str, int]:
with open(config_file, 'r') as file:
config = yaml.safe_load(file)
server_host = config.get('hostname', 'localhost')
server_port = config.get('port', 8000)
return f"http://{server_host}:{server_port}"
return server_host, server_port


def get_test_config(test_desc, example_dir, test_root):
Expand Down Expand Up @@ -277,7 +278,8 @@ def get_test_config(test_desc, example_dir, test_root):
raise ValueError(f"Invalid test description: {test_desc}, "
f"valid descriptions are: {config_map.keys()}")

return config_map[test_desc]
return (config_map[test_desc][0],
revise_disagg_config_file_with_free_ports(config_map[test_desc][1]))


def get_extra_llm_config(config, suffix, cwd):
Expand Down Expand Up @@ -481,7 +483,8 @@ def run_disaggregated_test(example_dir,
'trtllm-serve', 'disaggregated', '--server_start_timeout',
str(server_start_timeout), '-c', config_file
]
server_url = get_disagg_server_url_from_cfg(config_file)
server_host, server_port = get_disagg_server_url_from_cfg(config_file)
server_url = f"http://{server_host}:{server_port}"

try:
if not use_ray:
Expand Down Expand Up @@ -538,8 +541,8 @@ def run_disaggregated_test(example_dir,
env=run_env,
cwd=cwd))

if not wait_for_server("localhost",
8000,
if not wait_for_server(server_host,
server_port,
timeout_seconds=server_start_timeout):
raise RuntimeError(
f"Disaggregated server failed to start within {server_start_timeout} seconds"
Expand Down Expand Up @@ -1569,6 +1572,7 @@ def run_disaggregated_benchmark(example_dir,
'trtllm-serve', 'disaggregated', '--server_start_timeout',
str(server_start_timeout), '-c', config_file
]
server_host, server_port = get_disagg_server_url_from_cfg(config_file)
try:
with ( # Start workers
open('output_workers.log', 'w') as output_workers,
Expand Down Expand Up @@ -1622,9 +1626,9 @@ def run_disaggregated_benchmark(example_dir,
'--max-concurrency',
str(max_concurrency),
'--host',
'localhost',
server_host,
'--port',
'8000',
str(server_port),
'--ignore-eos',
'--no-test-input',
'--percentile-metrics',
Expand Down Expand Up @@ -1666,7 +1670,7 @@ def get_config_for_benchmark(model_root, backend):
serve_config = {
"model": model_root,
"hostname": "localhost",
"port": 8000,
"port": get_free_port(),
"backend": "pytorch",
"context_servers": {
"num_instances": 1,
Expand All @@ -1680,7 +1684,7 @@ def get_config_for_benchmark(model_root, backend):
"backend": backend,
"max_tokens_in_buffer": 512,
},
"urls": ["localhost:8001"]
"urls": [f"localhost:{get_free_port()}"]
},
"generation_servers": {
"num_instances": 1,
Expand All @@ -1693,7 +1697,7 @@ def get_config_for_benchmark(model_root, backend):
"backend": backend,
"max_tokens_in_buffer": 512,
},
"urls": ["localhost:8002"]
"urls": [f"localhost:{get_free_port()}"]
}
}
return serve_config
Expand Down Expand Up @@ -1724,6 +1728,7 @@ def run_disaggregated_genai_perf(config_file,
]

artifact_dir = os.path.join(cwd or ".", "benchmark-results")
server_host, server_port = get_disagg_server_url_from_cfg(config_file)

try:
with (open('output_workers.log', 'w') as output_workers,
Expand All @@ -1740,8 +1745,9 @@ def run_disaggregated_genai_perf(config_file,
cwd=cwd) as server_proc):

# Wait for server to be ready
if not wait_for_server(
"localhost", 8000, timeout_seconds=server_start_timeout):
if not wait_for_server(server_host,
server_port,
timeout_seconds=server_start_timeout):
raise RuntimeError(
f"Disaggregated server did not become ready within {server_start_timeout} seconds"
)
Expand All @@ -1751,7 +1757,7 @@ def run_disaggregated_genai_perf(config_file,
'genai-perf', 'profile', '--model', model_path, '--tokenizer',
model_path, '--endpoint-type', 'chat', '--endpoint',
'/v1/chat/completions', '--streaming', '--url',
'localhost:8000', '--synthetic-input-tokens-mean',
f'{server_host}:{server_port}', '--synthetic-input-tokens-mean',
str(input_tokens), '--synthetic-input-tokens-stddev', '0',
'--output-tokens-mean',
str(output_tokens), '--output-tokens-stddev', '0',
Expand Down
2 changes: 2 additions & 0 deletions tests/integration/defs/disaggregated/test_workers.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
import aiohttp
import pytest
import yaml
from defs.common import revise_disagg_config_file_with_free_ports
from defs.conftest import skip_no_hopper
from defs.trt_test_alternative import popen
from transformers import AutoTokenizer
Expand Down Expand Up @@ -42,6 +43,7 @@ def run_disaggregated_workers(
num_ranks: Optional[int] = None
) -> Tuple[Generator[subprocess.Popen, None, None], List[str], List[str]]:

config_file = revise_disagg_config_file_with_free_ports(config_file)
ctx_servers, gen_servers = get_ctx_gen_server_urls_from_cfg(config_file)

# TODO: auto detect num_ranks
Expand Down