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
35 changes: 6 additions & 29 deletions tensorrt_llm/commands/serve.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
import signal # Added import
import subprocess # nosec B404
import sys
from typing import Any, List, Optional
from typing import Any, Optional

import click
import torch
Expand All @@ -20,8 +20,7 @@
from tensorrt_llm.llmapi import (BuildConfig, CapacitySchedulerPolicy,
DynamicBatchConfig, KvCacheConfig,
SchedulerConfig)
from tensorrt_llm.llmapi.disagg_utils import (CtxGenServerConfig,
MetadataServerConfig, ServerRole,
from tensorrt_llm.llmapi.disagg_utils import (MetadataServerConfig, ServerRole,
parse_disagg_config_file,
parse_metadata_server_config_file)
from tensorrt_llm.llmapi.llm_utils import update_llm_args_with_extra_dict
Expand Down Expand Up @@ -442,19 +441,6 @@ def serve_encoder(model: str, host: str, port: int, log_level: str,
launch_mm_encoder_server(host, port, encoder_args, metadata_server_cfg)


def get_ctx_gen_server_urls(
server_configs: List[CtxGenServerConfig]) -> List[str]:
ctx_server_urls = []
gen_server_urls = []
for cfg in server_configs:
if cfg.type == "ctx":
ctx_server_urls.append(f"http://{cfg.hostname}:{cfg.port}")
else:
gen_server_urls.append(f"http://{cfg.hostname}:{cfg.port}")

return ctx_server_urls, gen_server_urls


@click.command("disaggregated")
@click.option("-c",
"--config_file",
Expand Down Expand Up @@ -491,22 +477,13 @@ def disaggregated(config_file: Optional[str],

disagg_cfg = parse_disagg_config_file(config_file)

ctx_server_urls, gen_server_urls = get_ctx_gen_server_urls(
disagg_cfg.server_configs)

metadata_server_cfg = parse_metadata_server_config_file(
metadata_server_config_file)

server = OpenAIDisaggServer(
ctx_servers=ctx_server_urls,
gen_servers=gen_server_urls,
req_timeout_secs=request_timeout,
server_start_timeout_secs=server_start_timeout,
max_retries=disagg_cfg.max_retries,
ctx_router_config=disagg_cfg.ctx_router_config,
gen_router_config=disagg_cfg.gen_router_config,
conditional_disagg_config=disagg_cfg.conditional_disagg_config,
metadata_server_cfg=metadata_server_cfg)
server = OpenAIDisaggServer(config=disagg_cfg,
req_timeout_secs=request_timeout,
server_start_timeout_secs=server_start_timeout,
metadata_server_cfg=metadata_server_cfg)

asyncio.run(server(disagg_cfg.hostname, disagg_cfg.port))

Expand Down
19 changes: 18 additions & 1 deletion tensorrt_llm/llmapi/disagg_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,7 @@ class DisaggServerConfig():
gen_router_config: Optional[RouterConfig] = None
conditional_disagg_config: Optional[ConditionalDisaggConfig] = None
max_retries: int = 3
perf_metrics_max_requests: int = 0


@dataclass
Expand All @@ -63,6 +64,20 @@ class MetadataServerConfig():
refresh_interval: float = 10.0


def get_ctx_gen_server_urls(
server_configs: list[CtxGenServerConfig]
) -> tuple[list[str], list[str]]:
ctx_server_urls = []
Comment thread
zhengd-nv marked this conversation as resolved.
gen_server_urls = []
for cfg in server_configs:
if cfg.type == "ctx":
ctx_server_urls.append(f"http://{cfg.hostname}:{cfg.port}")
else:
gen_server_urls.append(f"http://{cfg.hostname}:{cfg.port}")

return ctx_server_urls, gen_server_urls


def parse_disagg_config_file(yaml_config_file: str):

with open(yaml_config_file, 'r') as file:
Expand All @@ -77,6 +92,7 @@ def parse_disagg_config_file(yaml_config_file: str):
def extract_disagg_cfg(hostname: str = 'localhost',
port: int = 8000,
max_retries: int = 3,
perf_metrics_max_requests: int = 0,
context_servers: Optional[dict] = None,
generation_servers: Optional[dict] = None,
conditional_disagg_config: Optional[dict] = None,
Expand Down Expand Up @@ -115,7 +131,8 @@ def extract_disagg_cfg(hostname: str = 'localhost',

config = DisaggServerConfig(server_configs, hostname, port,
ctx_router_config, gen_router_config,
conditional_disagg_config, max_retries)
conditional_disagg_config, max_retries,
perf_metrics_max_requests)

return config

Expand Down
6 changes: 6 additions & 0 deletions tensorrt_llm/llmapi/llm_args.py
Original file line number Diff line number Diff line change
Expand Up @@ -2115,6 +2115,12 @@ class TorchLlmArgs(BaseLlmArgs):
description="Print iteration logs.",
status="beta")

perf_metrics_max_requests: int = Field(
default=0,
description=
"The maximum number of requests for perf metrics. Must also set request_perf_metrics to true to get perf metrics.",
status="prototype")

batch_wait_timeout_ms: float = Field(
default=0,
description=
Expand Down
110 changes: 91 additions & 19 deletions tensorrt_llm/serve/openai_disagg_server.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,14 @@
#!/usr/bin/env python
import asyncio
import copy
import itertools
Comment thread
zhengd-nv marked this conversation as resolved.
import os
import signal
import traceback
from collections import deque
from contextlib import asynccontextmanager
from http import HTTPStatus
from typing import List, Optional, Type, Union
from typing import Optional, Type, Union

import aiohttp
import uvicorn
Expand All @@ -17,9 +19,9 @@

# yapf: disable
from tensorrt_llm.executor import CppExecutorError
from tensorrt_llm.llmapi.disagg_utils import (ConditionalDisaggConfig,
from tensorrt_llm.llmapi.disagg_utils import (DisaggServerConfig,
MetadataServerConfig,
RouterConfig)
get_ctx_gen_server_urls)
from tensorrt_llm.logger import logger
from tensorrt_llm.serve.metadata_server import create_metadata_server
from tensorrt_llm.serve.openai_protocol import (ChatCompletionRequest,
Expand All @@ -37,35 +39,45 @@
class OpenAIDisaggServer:

def __init__(self,
ctx_servers: List[str],
gen_servers: List[str],
config: DisaggServerConfig,
req_timeout_secs: int = 180,
server_start_timeout_secs: int = 180,
max_retries: int = 3,
ctx_router_config: Optional[RouterConfig] = None,
gen_router_config: Optional[RouterConfig] = None,
conditional_disagg_config: Optional[ConditionalDisaggConfig] = None,
metadata_server_cfg: Optional[MetadataServerConfig] = None):

self.ctx_servers = ctx_servers
self.gen_servers = gen_servers
self.ctx_servers, self.gen_servers = get_ctx_gen_server_urls(config.server_configs)
self.metadata_server = create_metadata_server(metadata_server_cfg)
self.ctx_router = create_router(ctx_router_config, ctx_servers, metadata_server_cfg, self.metadata_server)
self.gen_router = create_router(gen_router_config, gen_servers, metadata_server_cfg, self.metadata_server)
self.conditional_disagg_config = conditional_disagg_config
self.ctx_router = create_router(
config.ctx_router_config, self.ctx_servers, metadata_server_cfg, self.metadata_server)
self.gen_router = create_router(
config.gen_router_config, self.gen_servers, metadata_server_cfg, self.metadata_server)
self.conditional_disagg_config = config.conditional_disagg_config

self.perf_metrics_max_requests = config.perf_metrics_max_requests
if self.perf_metrics_max_requests > 0:
# record corresponding keys of context and generation servers for perf metrics
# (ctx_server, gen_server, ctx_request_id)
self.perf_metrics_keys = deque(maxlen=self.perf_metrics_max_requests)
self.perf_metrics_keys_lock = asyncio.Lock()
# server_key -> {ctx_request_id: perf_metrics}
self.server_perf_metrics: dict[str, dict[int, dict]] = {}
else:
self.perf_metrics_keys = None
self.perf_metrics_keys_lock = None
self.server_perf_metrics = None
Comment thread
zhengd-nv marked this conversation as resolved.

if max_retries < 0:
raise ValueError(f"Max retries {max_retries} must be greater than or equal to 0")
self.max_retries = max_retries
if config.max_retries < 0:
raise ValueError(f"Max retries {config.max_retries} must be greater than or equal to 0")
self.max_retries = config.max_retries
logger.info(f"Server max retries: {self.max_retries}")

if (len(self.gen_servers) == 0):
raise ValueError("At least one generation server must be provided")

if os.getenv("TRTLLM_DISAGG_BENCHMARK_GEN_ONLY") != "1" and len(ctx_servers) == 0:
if os.getenv("TRTLLM_DISAGG_BENCHMARK_GEN_ONLY") != "1" and len(self.ctx_servers) == 0:
raise ValueError("At least one context server must be provided")

if self.conditional_disagg_config is not None and not isinstance(self.gen_router, KvCacheAwareRouter):
if self.conditional_disagg_config is not None and \
not isinstance(self.gen_router, KvCacheAwareRouter):
raise ValueError("Generation router must be a KvCacheAwareRouter to enable conditional disaggregation")

# Session will be initialized in lifespan
Expand Down Expand Up @@ -112,6 +124,7 @@ def create_error_response(
def register_routes(self):
self.app.add_api_route("/health", self.health, methods=["GET"])
self.app.add_api_route("/version", self.version, methods=["GET"])
self.app.add_api_route("/perf_metrics", self.perf_metrics, methods=["GET"])
self.app.add_api_route("/v1/completions",
self.openai_completion,
methods=["POST"])
Expand All @@ -126,6 +139,61 @@ async def version(self) -> JSONResponse:
ver = {"version": VERSION}
return JSONResponse(content=ver)

async def _add_perf_metrics_keys(self, ctx_server: str, gen_server: str, ctx_request_id: int):
async with self.perf_metrics_keys_lock:
self.perf_metrics_keys.append((ctx_server, gen_server, ctx_request_id))

async def perf_metrics(self) -> JSONResponse:
if self.perf_metrics_keys is None:
return JSONResponse(content=[])

perf_metrics = {}
exc = None
try:
for server in self.ctx_servers + self.gen_servers:
async with self.session.get(f"{server}/perf_metrics") as response:
server_perf_metrics = await response.json()
perf_metrics[server] = server_perf_metrics
except Exception as e:
# Keep the exception to raise it after saving perf metrics
exc = e

return_metrics = []
async with self.perf_metrics_keys_lock:
for server in perf_metrics:
server_metrics = self.server_perf_metrics.setdefault(server, {})
for request_perf_metrics in perf_metrics[server]:
ctx_request_id = request_perf_metrics.get("ctx_request_id", None)
if ctx_request_id is None:
continue
server_metrics[ctx_request_id] = request_perf_metrics

if len(server_metrics) > self.perf_metrics_max_requests:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

If possible, you could use a ring buffer with size = perf_metrics_max_requests instead. Would be faster and easier to understand.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

ctx_request_id is used as a look-up key in later metrics aggregation, so server_metrics here has to be a map with max entries.

# Remove oldest requests and keep at most perf_metrics_max_requests
num_remove = len(server_metrics) - self.perf_metrics_max_requests
removed_keys = list(itertools.islice(server_metrics.keys(), num_remove))
for ctx_request_id in removed_keys:
server_metrics.pop(ctx_request_id)
if exc is not None:
raise exc

Comment thread
zhengd-nv marked this conversation as resolved.
remain_keys = []
for ctx_server, gen_server, ctx_request_id in self.perf_metrics_keys:
gen_perf_metrics = self.server_perf_metrics[gen_server].pop(ctx_request_id, None)
if gen_perf_metrics is None:
# generation not finished
remain_keys.append((ctx_server, gen_server, ctx_request_id))
continue
ctx_perf_metrics = self.server_perf_metrics[ctx_server].pop(ctx_request_id, None)
return_metrics.append({
"ctx_server": ctx_server,
"gen_server": gen_server,
"ctx_perf_metrics": ctx_perf_metrics,
"gen_perf_metrics": gen_perf_metrics})
self.perf_metrics_keys = deque(remain_keys, maxlen=self.perf_metrics_max_requests)

return JSONResponse(content=return_metrics)

async def merge_streaming_responses(self, ctx_response,
gen_server: str,
gen_req: Union[CompletionRequest, ChatCompletionRequest]):
Expand Down Expand Up @@ -263,6 +331,10 @@ async def _send_disagg_request(self, req: Union[CompletionRequest, ChatCompletio
gen_server, _ = await self.gen_router.get_next_server(req)
logger.debug("Sending request to gen server: %s", gen_server)

if need_ctx and self.perf_metrics_keys is not None:
asyncio.create_task(self._add_perf_metrics_keys(
ctx_server, gen_server, req.disaggregated_params.ctx_request_id))

if not req.stream:
try:
#If request finished after first token for reason other than length, return right away and skip gen
Expand Down
Loading