From 1ae55aa10a1b38933cbed04318af6ca855bd43a4 Mon Sep 17 00:00:00 2001 From: Patrick Reiter Horn Date: Fri, 27 Jun 2025 16:14:26 -0700 Subject: [PATCH 1/2] Use decorator for request cancelation and handle CancelledError Signed-off-by: Patrick Reiter Horn --- tensorrt_llm/serve/openai_server.py | 147 ++++++++++++++++++--- tensorrt_llm/serve/postprocess_handlers.py | 31 ++--- 2 files changed, 138 insertions(+), 40 deletions(-) diff --git a/tensorrt_llm/serve/openai_server.py b/tensorrt_llm/serve/openai_server.py index 63f7e82c73d4..a02a4eb846dc 100644 --- a/tensorrt_llm/serve/openai_server.py +++ b/tensorrt_llm/serve/openai_server.py @@ -1,18 +1,23 @@ #!/usr/bin/env python import asyncio +import logging import signal import traceback from contextlib import asynccontextmanager from datetime import datetime +from functools import wraps from http import HTTPStatus from pathlib import Path -from typing import AsyncGenerator, AsyncIterator, List, Optional, Tuple +from typing import (AsyncGenerator, AsyncIterator, List, Optional, Tuple, + TypedDict, Callable, Awaitable, Any, Type) import uvicorn from fastapi import FastAPI, Request -from fastapi.exceptions import RequestValidationError +from fastapi.exceptions import RequestValidationError, HTTPException from fastapi.responses import JSONResponse, Response, StreamingResponse from transformers import AutoConfig, AutoProcessor +from openai.types.chat import ChatCompletionMessageParam +from pydantic import BaseModel from tensorrt_llm._tensorrt_engine import LLM # yapf: disable @@ -47,6 +52,73 @@ TIMEOUT_KEEP_ALIVE = 5 # seconds. +async def disconnect_poller(request: Request, result: Any): + """ + Poll for a disconnect. + If the request disconnects, stop polling and return. + """ + try: + while True: + message = await request.receive() + if message["type"] == "http.disconnect": + break + + print("Request disconnected") + + return result + except asyncio.CancelledError: + print("Stopping polling loop") + + +def cancel_on_disconnect(model_type: Type[BaseModel]): + """ + Decorator that will check if the client disconnects, + and cancel the task if required. + """ + + def cancel_on_disconnect_inner(handler: Callable): + + @wraps(handler) + async def cancel_on_disconnect_decorator(self, request: model_type, raw_request: Request): + sentinel = object() + + # Create two tasks, one to poll the request and check if the + # client disconnected, and another which is the request handler + poller_task = asyncio.ensure_future(disconnect_poller(raw_request, sentinel)) + handler_task = asyncio.ensure_future(handler(self, request=request, raw_request=raw_request)) + + done, pending = await asyncio.wait( + [poller_task, handler_task], return_when=asyncio.FIRST_COMPLETED + ) + + # Cancel any outstanding tasks + for t in pending: + t.cancel() + + try: + await t + except asyncio.CancelledError: + print(f"{t} was cancelled") + except Exception as exc: + print(f"{t} raised {exc} when being cancelled") + + # Return the result if the handler finished first + if handler_task in done: + return await handler_task + + # Otherwise, raise an exception + # This is not exactly needed, but it will prevent + # validation errors if your request handler is supposed + # to return something. + print("Raising an HTTP error because I was disconnected!!") + + raise HTTPException(503) + + return cancel_on_disconnect_decorator + + return cancel_on_disconnect_inner + + class OpenAIServer: def __init__(self, @@ -212,8 +284,11 @@ async def get_kv_cache_events(self) -> JSONResponse: pass return JSONResponse(content=events) + @cancel_on_disconnect(ChatCompletionRequest) async def openai_chat(self, request: ChatCompletionRequest, raw_request: Request) -> Response: + did_complete = False + def get_role() -> str: if request.add_generation_prompt: role = "assistant" @@ -223,14 +298,24 @@ def get_role() -> str: async def chat_stream_generator( promise: RequestOutput, postproc_params: PostprocParams) -> AsyncGenerator[str, None]: + nonlocal did_complete if not self.postproc_worker_enabled: post_processor, args = postproc_params.post_processor, postproc_params.postproc_args - async for res in promise: - pp_results = res.outputs[0]._postprocess_result if self.postproc_worker_enabled else post_processor(res, args) - for pp_res in pp_results: - yield pp_res - yield "data: [DONE]\n\n" - nvtx_mark("generation ends") + try: + async for res in promise: + pp_results = res.outputs[0]._postprocess_result if self.postproc_worker_enabled else post_processor(res, args) + for pp_res in pp_results: + for choice in pp_res.choices: + if choice.finish_reason is not None: + did_complete = True + + pp_res_json = pp_res.model_dump_json(exclude_unset=True) + yield f"data: {pp_res_json}\n\n" + yield f"data: [DONE]\n\n" + nvtx_mark("generation ends") + finally: + if not did_complete: + promise.abort() async def create_chat_response( promise: RequestOutput, postproc_params: PostprocParams, disaggregated_params: Optional[LlmDisaggregatedParams] = None) -> ChatCompletionResponse: @@ -246,6 +331,7 @@ async def create_chat_response( chat_response.prompt_token_ids = promise.prompt_token_ids return chat_response + promise: Optional[RequestOutput] = None try: check_multiple_response(request.n, self.llm.args.backend) conversation: List[ConversationMessage] = [] @@ -297,7 +383,6 @@ async def create_chat_response( lora_request=request.lora_request, disaggregated_params=disaggregated_params ) - asyncio.create_task(self.await_disconnected(raw_request, promise)) if not self.postproc_worker_enabled: postproc_args.tokenizer = self.tokenizer postproc_args.num_prompt_tokens = len(promise.prompt_token_ids) @@ -312,9 +397,14 @@ async def create_chat_response( except CppExecutorError: # If internal executor error is raised, shutdown the server signal.raise_signal(signal.SIGINT) + except asyncio.CancelledError: + if promise is not None: + promise.abort() + return self.create_error_response("cancelled") except Exception as e: return self.create_error_response(str(e)) + @cancel_on_disconnect(CompletionRequest) async def openai_completion(self, request: CompletionRequest, raw_request: Request) -> Response: def merge_promises( @@ -343,16 +433,27 @@ async def consumer(): return consumer() async def create_completion_generator( - generator: AsyncIterator[Tuple[RequestOutput, Optional[PostprocParams]]]): - async for request_output, postproc_params in generator: - if not self.postproc_worker_enabled: - post_processor, args = postproc_params.post_processor, postproc_params.postproc_args - pp_result = post_processor(request_output, args) - else: - pp_result = request_output.outputs[0]._postprocess_result - for pp_res in pp_result: - yield pp_res - yield "data: [DONE]\n\n" + generator: AsyncIterator[Tuple[RequestOutput, Optional[PostprocParams]]], + promises: List[RequestOutput]): + did_complete = False + try: + async for request_output, postproc_params in generator: + if not self.postproc_worker_enabled: + post_processor, args = postproc_params.post_processor, postproc_params.postproc_args + pp_result = post_processor(request_output, args) + else: + pp_result = request_output.outputs[0]._postprocess_result + for pp_res in pp_result: + for choice in pp_res.choices: + if choice.finish_reason is not None: + did_complete = True + pp_res_json = pp_res.model_dump_json(exclude_unset=False) + yield f"data: {pp_res_json}\n\n" + yield f"data: [DONE]\n\n" + finally: + if not did_complete: + for promise in promises: + promise.abort() async def create_completion_response( generator: AsyncIterator[Tuple[RequestOutput, Optional[PostprocParams]]], disaggregated_params: Optional[LlmDisaggregatedParams] = None) -> CompletionResponse: @@ -418,7 +519,6 @@ async def create_completion_response( lora_request=request.lora_request, disaggregated_params=disaggregated_params ) - asyncio.create_task(self.await_disconnected(raw_request, promise)) if not self.postproc_worker_enabled: postproc_args.tokenizer = self.tokenizer postproc_args.num_prompt_tokens = len(promise.prompt_token_ids) @@ -427,8 +527,7 @@ async def create_completion_response( generator = merge_promises(promises, postproc_params_collection) if request.stream: - response_generator = create_completion_generator( - generator) + response_generator = create_completion_generator(generator, promises) return StreamingResponse(content=response_generator, media_type="text/event-stream") else: @@ -438,6 +537,10 @@ async def create_completion_response( except CppExecutorError: # If internal executor error is raised, shutdown the server signal.raise_signal(signal.SIGINT) + except asyncio.CancelledError: + for promise in promises: + promise.abort() + return self.create_error_response("cancelled") except Exception as e: traceback.print_exc() return self.create_error_response(str(e)) diff --git a/tensorrt_llm/serve/postprocess_handlers.py b/tensorrt_llm/serve/postprocess_handlers.py index 321ff6cc9060..ead864e6eaa0 100644 --- a/tensorrt_llm/serve/postprocess_handlers.py +++ b/tensorrt_llm/serve/postprocess_handlers.py @@ -97,12 +97,12 @@ def apply_reasoning_parser(args: ChatPostprocArgs, output_index: int, text: str, @nvtx_range_debug("chat_stream_post_processor") -def chat_stream_post_processor(rsp: GenerationResultBase, args: ChatPostprocArgs) -> List[str]: +def chat_stream_post_processor(rsp: GenerationResultBase, args: ChatPostprocArgs) -> List[ChatCompletionStreamResponse]: def yield_first_chat(num_tokens: int, idx: int, role: str = None, - content: str = None): + content: str = None) -> ChatCompletionStreamResponse: choice_data = ChatCompletionResponseStreamChoice(index=idx, delta=DeltaMessage( role=role, @@ -114,10 +114,9 @@ def yield_first_chat(num_tokens: int, chunk.usage = UsageInfo(prompt_tokens=num_tokens, total_tokens=num_tokens, completion_tokens=0) - data = chunk.model_dump_json(exclude_none=True) - return data + return chunk - res: List[str] = [] + res: List[ChatCompletionStreamResponse] = [] finish_reason_sent = [False] * args.num_choices prompt_tokens = args.num_prompt_tokens if stream_option := args.stream_options: @@ -128,9 +127,9 @@ def yield_first_chat(num_tokens: int, include_continuous_usage = False if args.first_iteration: for i in range(args.num_choices): - res.append(f"data: {yield_first_chat(prompt_tokens, i, role=args.role)} \n\n") + res.append(yield_first_chat(prompt_tokens, i, role=args.role)) if args.echo and args.last_message_content: - res.append(f"data: {yield_first_chat(prompt_tokens, i, content=args.last_message_content)} \n\n") + res.append(yield_first_chat(prompt_tokens, i, content=args.last_message_content)) args.first_iteration = False for output in rsp.outputs: @@ -174,8 +173,7 @@ def yield_first_chat(num_tokens: int, chunk.usage = UsageInfo(prompt_tokens=prompt_tokens, completion_tokens=output.length, total_tokens=output.length + prompt_tokens) - data = chunk.model_dump_json(exclude_none=True) - res.append(f"data: {data}\n\n") + res.append(chunk) if include_usage and rsp._done: completion_tokens = sum(output.length @@ -188,8 +186,7 @@ def yield_first_chat(num_tokens: int, final_usage_chunk = ChatCompletionStreamResponse( choices=[], model=args.model, usage=final_usage) - final_usage_data = final_usage_chunk.model_dump_json() - res.append(f"data: {final_usage_data}\n\n") + res.append(final_usage_chunk) return res @@ -271,8 +268,8 @@ def from_request(cls, request: CompletionRequest): @nvtx_range_debug("completion_stream_post_processor") -def completion_stream_post_processor(rsp: DetokenizedGenerationResultBase, args: CompletionPostprocArgs) -> List[str]: - res: List[str] = [] +def completion_stream_post_processor(rsp: DetokenizedGenerationResultBase, args: CompletionPostprocArgs) -> List[CompletionStreamResponse]: + res: List[CompletionStreamResponse] = [] prompt_tokens = args.num_prompt_tokens if stream_option := args.stream_options: include_usage = stream_option.include_usage @@ -296,8 +293,7 @@ def completion_stream_post_processor(rsp: DetokenizedGenerationResultBase, args: chunk.usage = UsageInfo(prompt_tokens=prompt_tokens, completion_tokens=output.length, total_tokens=output.length + prompt_tokens) - data = chunk.model_dump_json(exclude_unset=False) - res.append(f"data: {data}\n\n") + res.append(chunk) if include_usage and rsp._done: completion_tokens = sum(output.length @@ -308,10 +304,9 @@ def completion_stream_post_processor(rsp: DetokenizedGenerationResultBase, args: total_tokens=prompt_tokens + completion_tokens, ) - final_usage_chunk = ChatCompletionStreamResponse( + final_usage_chunk = CompletionStreamResponse( choices=[], model=args.model, usage=final_usage) - final_usage_data = final_usage_chunk.model_dump_json() - res.append(f"data: {final_usage_data}\n\n") + res.append(final_usage_chunk) args.first_iteration = False return res From f78d0c1685a3306505c25dfce95127173523b0f2 Mon Sep 17 00:00:00 2001 From: Patrick Reiter Horn Date: Fri, 27 Jun 2025 16:24:34 -0700 Subject: [PATCH 2/2] Reimplement metrics endpoint with stats about requests Signed-off-by: Patrick Reiter Horn --- tensorrt_llm/_torch/pyexecutor/py_executor.py | 38 ++++++++- tensorrt_llm/serve/openai_server.py | 81 ++++++++++++++++++- 2 files changed, 116 insertions(+), 3 deletions(-) diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index ff0eff2b9d4d..bc5d96fed60c 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -36,7 +36,15 @@ LlmResponse, executor_request_to_llm_request) from .model_engine import ModelEngine from .sampler import Sampler, SampleState, SampleStateTensors, TorchSampler -from .scheduler import RequestScheduler, ScheduledRequests +from .scheduler import ScheduledRequests +from collections import defaultdict +import array +import json + +PROM_METRICS_FILENAME = '/dev/shm/prom_metrics.json' + +prom_metrics = defaultdict(float) +prom_metrics_file = None # Environment variable to specify iteration ranges for profiling start/stop. # Format: "start1-stop1,start2-stop2,..." or single iterations "iter1,iter2,..." @@ -487,6 +495,7 @@ def profile_step(): f"trace saved to {torch_trace_path}") torch.cuda.cudart().cudaProfilerStop() enabled = False + last_start_time = start_time if start_time is not None and self.print_log and self.dist.rank == 0: end_time = time.time() @@ -514,6 +523,33 @@ def profile_step(): enabled = True start_time = time.time() + if last_start_time is not None and self.dist.rank == 0: + iter_states = self.model_engine.iter_states + total_running = iter_states['num_ctx_requests'] + iter_states['num_generation_tokens'] / (1 + self.model_engine.max_draft_len) + prom_metrics["num_requests_running"] = total_running + prom_metrics["num_requests_swapped"] = total_running - len(self.active_requests) + prom_metrics["iteration_tokens_total_sum"] += iter_states['num_ctx_tokens'] + iter_states['num_generation_tokens'] + prom_metrics["iteration_tokens_total_count"] += 1 + prom_metrics["time_per_output_token_seconds_sum"] += (start_time - last_start_time) + prom_metrics["time_per_output_token_seconds_count"] += 1 + prom_metrics["prompt_tokens_total"] += iter_states['num_ctx_tokens'] + prom_metrics["request_prompt_tokens_total_sum"] += iter_states['num_ctx_tokens'] + prom_metrics["request_prompt_tokens_total_count"] += 1 + prom_metrics["generation_tokens_total"] += iter_states['num_generation_tokens'] + prom_metrics["request_generation_tokens_total_sum"] += iter_states['num_generation_tokens'] + prom_metrics["request_generation_tokens_total_count"] += 1 + global prom_metrics_file + try: + if prom_metrics_file is None: + prom_metrics_file = os.open(PROM_METRICS_FILENAME, + os.O_RDWR|os.O_CREAT) + os.pwrite(prom_metrics_file, ( + json.dumps(list(prom_metrics.keys())).encode('UTF-8') + + b'\0' + + array.array('d',prom_metrics.values()).tobytes()), 0) + except: + traceback.print_exc() + try: yield profile_step finally: diff --git a/tensorrt_llm/serve/openai_server.py b/tensorrt_llm/serve/openai_server.py index a02a4eb846dc..6d9de3cce63b 100644 --- a/tensorrt_llm/serve/openai_server.py +++ b/tensorrt_llm/serve/openai_server.py @@ -45,6 +45,21 @@ chat_stream_post_processor, completion_response_post_processor, completion_stream_post_processor) from tensorrt_llm.version import __version__ as VERSION +from tensorrt_llm._torch.pyexecutor.py_executor import PROM_METRICS_FILENAME + +from collections import defaultdict +import array +import json +import os +import traceback + +prom_metrics_file = None +prom_metrics = defaultdict(float, { + "num_requests_running": 0, + "num_requests_waiting": 0, + "prompt_tokens_total": 0, + "generation_tokens_total": 0, +}) from .._utils import nvtx_mark @@ -212,8 +227,9 @@ def register_routes(self): self.app.add_api_route("/health_generate", self.health_generate, methods=["GET"]) self.app.add_api_route("/version", self.version, methods=["GET"]) self.app.add_api_route("/v1/models", self.get_model, methods=["GET"]) - # TODO: the metrics endpoint only reports iteration stats, not the runtime stats for now - self.app.add_api_route("/metrics", self.get_iteration_stats, methods=["GET"]) + # TODO: the metrics endpoint only reports runtime stats, not iteration stats + self.app.add_api_route("/metrics", self.metrics, methods=["GET"]) + self.app.add_api_route("/metrics/", self.metrics, methods=["GET"]) # TODO: workaround before ETCD support self.app.add_api_route("/kv_cache_events", self.get_kv_cache_events, methods=["POST"]) self.app.add_api_route("/v1/completions", @@ -264,10 +280,49 @@ async def version(self) -> JSONResponse: ver = {"version": VERSION} return JSONResponse(content=ver) + async def metrics(self) -> Response: + global prom_metrics_file + bufs = None + try: + if prom_metrics_file is None: + prom_metrics_file = os.open(PROM_METRICS_FILENAME, + os.O_RDWR|os.O_CREAT|os.O_TRUNC) + bufs = os.pread(prom_metrics_file, 65536, 0).split(b'\0', 1) + if len(bufs) >= 2: + keybuf, valbuf = bufs + key_list = json.loads(keybuf.decode('UTF-8')) + value_list = array.array('d') + value_list.frombytes(valbuf) + for key, value in zip(key_list, value_list): + prom_metrics[key] = value + except: + print(bufs) + traceback.print_exc() + + all_requests_done = ( + prom_metrics["request_completed_total"] + + prom_metrics["request_cancelled_total"] + + prom_metrics["request_failed_total"]) + # NOTE: metrics do not update if the other thread is not running any requests. + # Make sure to zero out running and waiting in this case. + if prom_metrics["request_started_total"] == all_requests_done: + prom_metrics["num_requests_running"] = 0 + + # Detect number of requests not being processed by the TensorRT-LLM engine. + prom_metrics["num_requests_waiting"] = max(0, prom_metrics["request_started_total"] - ( + prom_metrics["num_requests_running"] + all_requests_done)) + + resp = '' + for metric_key, metric_val in prom_metrics.items(): + separator = ',' if '{' in metric_key else '{' + resp += f'pytrtllm:{metric_key}{separator}model_name="{self.model}"}} {float(metric_val)}\n' + return Response(status_code=200, content=resp) + async def get_model(self) -> JSONResponse: model_list = ModelList(data=[ModelCard(id=self.model)]) return JSONResponse(content=model_list.model_dump()) + # FIXME: Currently unused async def get_iteration_stats(self) -> JSONResponse: stats = [] async for stat in self.llm.get_stats_async(2): @@ -308,6 +363,8 @@ async def chat_stream_generator( for choice in pp_res.choices: if choice.finish_reason is not None: did_complete = True + prom_metrics["request_completed_total"] += 1 + prom_metrics[f"request_success_total{{finished_reason=\"{choice.finish_reason}\""] += 1 pp_res_json = pp_res.model_dump_json(exclude_unset=True) yield f"data: {pp_res_json}\n\n" @@ -315,6 +372,7 @@ async def chat_stream_generator( nvtx_mark("generation ends") finally: if not did_complete: + prom_metrics["request_cancelled_total"] += 1 promise.abort() async def create_chat_response( @@ -326,11 +384,17 @@ async def create_chat_response( post_processor, args = postproc_params.post_processor, postproc_params.postproc_args chat_response = post_processor(promise, args) + for choice in chat_response.choices: + if choice.finish_reason is not None: + prom_metrics["request_completed_total"] += 1 + prom_metrics[f"request_success_total{{finished_reason=\"{choice.finish_reason}\""] += 1 + # Add prompt_tokens_ids to the response if disaggregated_params and disaggregated_params.request_type and disaggregated_params.request_type == "context_only": chat_response.prompt_token_ids = promise.prompt_token_ids return chat_response + prom_metrics["request_started_total"] += 1 promise: Optional[RequestOutput] = None try: check_multiple_response(request.n, self.llm.args.backend) @@ -398,10 +462,12 @@ async def create_chat_response( # If internal executor error is raised, shutdown the server signal.raise_signal(signal.SIGINT) except asyncio.CancelledError: + prom_metrics["request_cancelled_total"] += 1 if promise is not None: promise.abort() return self.create_error_response("cancelled") except Exception as e: + prom_metrics["request_failed_total"] += 1 return self.create_error_response(str(e)) @cancel_on_disconnect(CompletionRequest) @@ -447,11 +513,14 @@ async def create_completion_generator( for choice in pp_res.choices: if choice.finish_reason is not None: did_complete = True + prom_metrics["request_completed_total"] += 1 + prom_metrics[f"request_success_total{{finished_reason=\"{choice.finish_reason}\""] += 1 pp_res_json = pp_res.model_dump_json(exclude_unset=False) yield f"data: {pp_res_json}\n\n" yield f"data: [DONE]\n\n" finally: if not did_complete: + prom_metrics["request_cancelled_total"] += 1 for promise in promises: promise.abort() @@ -468,6 +537,11 @@ async def create_completion_response( else: pp_result = request_output.outputs[0]._postprocess_result + for choice in pp_result.choices: + if choice.finish_reason is not None: + prom_metrics["request_completed_total"] += 1 + prom_metrics[f"request_success_total{{finished_reason=\"{choice.finish_reason}\""] += 1 + choices, usage = pp_result.choices, pp_result.usage all_choices.extend(choices) num_prompt_tokens += usage.prompt_tokens @@ -489,6 +563,7 @@ async def create_completion_response( ) return response + prom_metrics["request_started_total"] += 1 try: check_multiple_response(request.n, self.llm.args.backend) if isinstance(request.prompt, str) or \ @@ -538,11 +613,13 @@ async def create_completion_response( # If internal executor error is raised, shutdown the server signal.raise_signal(signal.SIGINT) except asyncio.CancelledError: + prom_metrics["request_cancelled_total"] += 1 for promise in promises: promise.abort() return self.create_error_response("cancelled") except Exception as e: traceback.print_exc() + prom_metrics["request_failed_total"] += 1 return self.create_error_response(str(e)) async def __call__(self, host, port):