diff --git a/lmdeploy/serve/anthropic/adapter.py b/lmdeploy/serve/anthropic/adapter.py index cf85f777c7..b65d7f13ea 100644 --- a/lmdeploy/serve/anthropic/adapter.py +++ b/lmdeploy/serve/anthropic/adapter.py @@ -29,7 +29,7 @@ def get_model_list(server_context) -> list[str]: """Return available model names from the server context.""" model_names = [server_context.async_engine.model_name] - cfg = server_context.async_engine.backend_config + cfg = server_context.engine_config model_names += getattr(cfg, 'adapters', None) or [] return model_names diff --git a/lmdeploy/serve/anthropic/endpoints/messages.py b/lmdeploy/serve/anthropic/endpoints/messages.py index 4a20b2274e..a4421adbdc 100644 --- a/lmdeploy/serve/anthropic/endpoints/messages.py +++ b/lmdeploy/serve/anthropic/endpoints/messages.py @@ -44,7 +44,7 @@ def _is_tool_choice_auto(tool_choice): def _validate_extended_outputs(request: MessagesRequest, server_context): - engine_config = server_context.get_engine_config() + engine_config = server_context.engine_config logprobs_mode = engine_config.logprobs_mode if request.return_logprob and logprobs_mode is None: return create_error_response( @@ -186,7 +186,7 @@ async def create_message(request: MessagesRequest, raw_request: Request): ) request_id = f'msg_{shortuuid.random()}' - session_mgr = server_context.get_session_manager() + session_mgr = server_context.session_manager if request.stream: return StreamingResponse( diff --git a/lmdeploy/serve/openai/api_server.py b/lmdeploy/serve/openai/api_server.py index 9fe802ca1f..00d9aa52b7 100644 --- a/lmdeploy/serve/openai/api_server.py +++ b/lmdeploy/serve/openai/api_server.py @@ -3,97 +3,41 @@ from __future__ import annotations import asyncio -import json import os import re -import time -from collections.abc import AsyncGenerator -from contextlib import aclosing, asynccontextmanager -from functools import partial -from http import HTTPStatus +from contextlib import asynccontextmanager from typing import TYPE_CHECKING, Literal if TYPE_CHECKING: - from transformers import PreTrainedTokenizerBase from lmdeploy.serve.parsers import ResponseParser -import shortuuid import uvicorn -from fastapi import APIRouter, Depends, FastAPI, HTTPException, Request, status +from fastapi import FastAPI, HTTPException, Request, status from fastapi.encoders import jsonable_encoder from fastapi.exceptions import RequestValidationError from fastapi.middleware.cors import CORSMiddleware -from fastapi.responses import JSONResponse, Response, StreamingResponse +from fastapi.responses import JSONResponse from starlette.middleware.base import BaseHTTPMiddleware from starlette.routing import Mount from lmdeploy.archs import get_task from lmdeploy.messages import ( - GenerationConfig, - LogitsProcessor, PytorchEngineConfig, SpeculativeConfig, TurbomindEngineConfig, ) from lmdeploy.metrics.metrics_processor import metrics_processor from lmdeploy.model import ChatTemplateConfig -from lmdeploy.pytorch.disagg.config import DistServeEngineConfig -from lmdeploy.pytorch.disagg.conn.protocol import ( - DistServeCacheFreeRequest, - DistServeConnectionRequest, - DistServeDropConnectionRequest, - DistServeInitRequest, - MigrationRequest, -) from lmdeploy.serve.anthropic import create_anthropic_router from lmdeploy.serve.core import AsyncEngine, EngineHealthMonitor from lmdeploy.serve.core.generation_config import ( - build_generation_config, resolve_default_gen_config, ) -from lmdeploy.serve.openai.protocol import ( - AbortRequest, - ChatCompletionRequest, - ChatCompletionResponse, # noqa: E501 - ChatCompletionResponseChoice, - ChatCompletionResponseStreamChoice, - ChatCompletionStreamResponse, - ChatCompletionTokenLogprob, - ChatMessage, - ChoiceLogprobs, - CompletionRequest, - CompletionResponse, - CompletionResponseChoice, - CompletionResponseStreamChoice, - CompletionStreamResponse, - DeltaMessage, - DestroyWeightsUpdateGroupRequest, - EmbeddingsRequest, - EncodeRequest, - EncodeResponse, - ErrorResponse, - GenerateReqInput, - GenerateReqMetaOutput, - GenerateReqOutput, - InitWeightsUpdateGroupRequest, - LogProbs, - ModelCard, - ModelList, - ModelPermission, - PoolingRequest, - PoolingResponse, - PPLRequest, - PPLResponse, - TopLogprob, - UpdateParamsRequest, - UpdateWeightsFromDistributedRequest, - UsageInfo, -) +from lmdeploy.serve.openai.endpoints import create_openai_router from lmdeploy.serve.openai.responses import create_responses_router -from lmdeploy.serve.openai.utils import maybe_filter_parallel_tool_calls -from lmdeploy.serve.utils.request_cleanup import with_request_cleanup -from lmdeploy.serve.utils.server_utils import AuthenticationMiddleware, EngineSleepingMiddleware, validate_json_request +from lmdeploy.serve.openai.utils import get_model_list +from lmdeploy.serve.utils.server_utils import AuthenticationMiddleware, EngineSleepingMiddleware from lmdeploy.utils import get_logger if TYPE_CHECKING: @@ -104,22 +48,29 @@ logger = get_logger('lmdeploy') -class VariableInterface: - """A IO interface maintaining variables.""" - async_engine: AsyncEngine = None - health_monitor: EngineHealthMonitor | None = None - request_hosts = [] - # following are for registering to proxy server - proxy_url: str | None = None - api_server_url: str | None = None - allow_terminate_by_client: bool = False - enable_abort_handling: bool = False - response_parser_cls: type[ResponseParser] | None = None - default_gen_config: dict = {} - - @classmethod - def create_session(cls, user_session_id: int | None = None) -> Session: - session_mgr = cls.get_session_manager() +class ServerContext: + """Runtime dependencies and state shared by serving endpoints.""" + + def __init__(self): + self.async_engine: AsyncEngine | None = None + self.health_monitor: EngineHealthMonitor | None = None + self.proxy_url: str | None = None + self.api_server_url: str | None = None + self.allow_terminate_by_client = False + self.enable_abort_handling = False + self.response_parser_cls: type[ResponseParser] | None = None + self.default_gen_config: dict = {} + + @property + def engine_config(self): + return self.async_engine.backend_config + + @property + def session_manager(self): + return self.async_engine.session_mgr + + def create_session(self, user_session_id: int | None = None) -> Session: + session_mgr = self.session_manager if user_session_id is None or user_session_id == -1: # user doesn't input session_id, so we need to generate a new one session = session_mgr.get() @@ -129,1233 +80,22 @@ def create_session(cls, user_session_id: int | None = None) -> Session: session_id = session_mgr.map_user_session_id(user_session_id) session = session_mgr.get(session_id) # Stamp epoch for ``stop_all_session`` / ``abort_all`` coordination in ``AsyncEngine.generate``. - session.epoch = cls.async_engine.epoch + session.epoch = self.async_engine.epoch return session - @classmethod - def find_session(cls, user_session_id: int) -> Session | None: + def find_session(self, user_session_id: int) -> Session | None: """Find the session by user_session_id. Users cannot access inner session_id directly. """ - session_mgr = cls.get_session_manager() + session_mgr = self.session_manager session_id = session_mgr.user_session_id_map.get(user_session_id, None) if session_id is None: return None return session_mgr.get(session_id, create_if_not_exists=False) - @classmethod - def get_session_manager(cls): - return cls.async_engine.session_mgr - - @classmethod - def get_engine_config(cls): - return cls.async_engine.backend_config - - -def _build_serving_generation_config(request, **extra_kwargs) -> GenerationConfig: - """Build ``GenerationConfig`` with server and request sampling merge. - - Same-name request fields are extracted automatically; ``extra_kwargs`` is - for renamed, computed, or raw-json fields only. - """ - max_new_tokens = getattr(request, 'max_completion_tokens', None) - if max_new_tokens is None: - max_new_tokens = getattr(request, 'max_tokens', None) - return build_generation_config( - request, - VariableInterface.default_gen_config, - max_new_tokens=max_new_tokens, - **extra_kwargs, - ) - - -async def _with_request_cleanup(generator, result_generators, sessions): - """Yield from an API generator and cleanup when the HTTP task exits.""" - session_mgr = VariableInterface.get_session_manager() - async for item in with_request_cleanup(generator, result_generators, sessions, session_mgr): - yield item - - -router = APIRouter() -server_context = VariableInterface() - - -def get_model_list(): - """Available models. - - If it is a slora serving. The model list would be [model_name, adapter_name1, adapter_name2, ...] - """ - model_names = [VariableInterface.async_engine.model_name] - cfg = VariableInterface.async_engine.backend_config - model_names += getattr(cfg, 'adapters', None) or [] - return model_names - - -@router.get('/v1/models') -def available_models(): - """Show available models.""" - model_cards = [] - for model_name in get_model_list(): - model_cards.append(ModelCard(id=model_name, root=model_name, permission=[ModelPermission()])) - return ModelList(data=model_cards) - - -def create_error_response(status: HTTPStatus, message: str, error_type='invalid_request_error'): - """Create error response according to http status and message. - - Args: - status (HTTPStatus): HTTP status codes and reason phrases - message (str): error message - error_type (str): error type - """ - return JSONResponse(ErrorResponse(message=message, type=error_type, code=status.value).model_dump(), - status_code=status.value) - - -def check_request(request) -> JSONResponse | None: - """Check if a request is valid.""" - if hasattr(request, 'model') and request.model not in get_model_list(): - return create_error_response(HTTPStatus.NOT_FOUND, f'The model {request.model!r} does not exist.') - - # Import the appropriate check function based on request type - if isinstance(request, ChatCompletionRequest): - from .serving_chat_completion import check_request - check_func = check_request - elif isinstance(request, CompletionRequest): - from .serving_completion import check_request - check_func = check_request - elif isinstance(request, GenerateReqInput): - from .serving_generate import check_request - check_func = check_request - else: - # Define an async function that always returns success - def always_success(req, server_context): - return '' - - check_func = always_success - - error_msg = check_func(request, server_context) - if error_msg: - return create_error_response(HTTPStatus.BAD_REQUEST, error_msg) - return None - - -def _create_chat_completion_logprobs(tokenizer: PreTrainedTokenizerBase, - token_ids: list[int] | None = None, - logprobs: list[dict[int, float]] | None = None): - """Create openai LogProbs for chat.completion. - - Args: - tokenizer (PreTrainedTokenizerBase): tokenizer. - token_ids (list[int]): output token ids. - logprobs (list[dict[int, float]]): the top logprobs for each output - position. - Returns: - ChoiceLogprobs: logprob result. - """ - if token_ids is None or logprobs is None: - return None - - content: list[ChatCompletionTokenLogprob] = [] - for token_id, tops in zip(token_ids, logprobs): - item = ChatCompletionTokenLogprob(token='', bytes=[], logprob=0.0, top_logprobs=[]) - for top_id, prob in tops.items(): - token = tokenizer.convert_ids_to_tokens(top_id) - if isinstance(token, bytes): - _bytes = list(token) - token = token.decode('utf-8', errors='backslashreplace') - else: - _bytes = list(token.encode()) # token is str - if top_id == token_id: - item.token = token - item.bytes = _bytes - item.logprob = prob - else: - item.top_logprobs.append(TopLogprob(token=token, bytes=_bytes, logprob=prob)) - content.append(item) - return ChoiceLogprobs(content=content) - - -def _create_output_token_logprobs(token_ids: list[int] | None = None, - logprobs: list[dict[int, float]] | None = None): - """Create raw (logprob, token_id) pairs for output tokens.""" - if token_ids is None or logprobs is None: - return None - - output_token_logprobs = [] - for tok, tok_logprobs in zip(token_ids, logprobs): - output_token_logprobs.append((tok_logprobs[tok], tok)) - return output_token_logprobs or None - - -@router.get('/health') -async def health() -> JSONResponse: - """Health check.""" - monitor = VariableInterface.health_monitor - if monitor is None: - data = dict(status='unhealthy', - message='Engine health monitor is not initialized.') - return JSONResponse(jsonable_encoder(data), status_code=HTTPStatus.SERVICE_UNAVAILABLE) - data = monitor.snapshot() - if data['status'] == 'unhealthy': - data = await monitor.refresh_snapshot() - status_code = HTTPStatus.OK if data['status'] in ('healthy', 'sleeping') else HTTPStatus.SERVICE_UNAVAILABLE - return JSONResponse(jsonable_encoder(data), status_code=status_code) - - -@router.get('/terminate') -async def terminate(): - """Terminate server.""" - import signal - - if not VariableInterface.allow_terminate_by_client: - return create_error_response( - HTTPStatus.BAD_REQUEST, - 'The server can not be terminated. Please add --allow-terminate-by-client when start the server.') - os.kill(os.getpid(), signal.SIGTERM) - return Response(status_code=200) - - -# modified from https://github.com/vllm-project/vllm/blob/v0.5.4/vllm/entrypoints/openai/logits_processors.py#L51 # noqa -def logit_bias_logits_processor(logit_bias: dict[int, float] | dict[str, float], - tokenizer: PreTrainedTokenizerBase) -> LogitsProcessor: - try: - # Convert token_id to integer - # Clamp the bias between -100 and 100 per OpenAI API spec - clamped_logit_bias: dict[int, float] = { - int(token_id): min(100.0, max(-100.0, bias)) - for token_id, bias in logit_bias.items() - } - except ValueError as exc: - raise ValueError('Found token_id in logit_bias that is not ' - 'an integer or string representing an integer') from exc - - # Check if token_id is within the vocab size - for token_id, bias in clamped_logit_bias.items(): - if token_id < 0 or token_id >= tokenizer.vocab_size: - raise ValueError(f'token_id {token_id} in logit_bias contains ' - 'out-of-vocab token id') - - def _logit_bias_processor( - logit_bias, - token_ids, - logits, - ): - for token_id, bias in logit_bias.items(): - logits[token_id] = logits[token_id] + bias - return logits - - return partial(_logit_bias_processor, clamped_logit_bias) - - -@router.post('/v1/chat/completions', dependencies=[Depends(validate_json_request)]) -async def chat_completions_v1(request: ChatCompletionRequest, raw_request: Request = None): - """Completion API similar to OpenAI's API. - - Refer to https://platform.openai.com/docs/api-reference/chat/create - for the API specification. - - The request should be a JSON object with the following fields: - - - **model**: model name. Available from /v1/models. - - **messages**: string prompt or chat history in OpenAI format. Chat history example: - ``[{"role": "user", "content": "hi"}]``. - - **temperature** (float): to modulate the next token probability - - **top_p** (float): If set to float < 1, only the smallest set of most - probable tokens with probabilities that add up to top_p or higher - are kept for generation. - - **n** (int): How many chat completion choices to generate for each input - message. **Only support one here**. - - **stream**: whether to stream the results or not. Default to false. - - **stream_options**: Options for streaming response. Only set this when you - set stream: true. - - **max_completion_tokens** (int | None): output token nums. Default to None. - - **max_tokens** (int | None): output token nums. Default to None. - Deprecated: Use max_completion_tokens instead. - - **repetition_penalty** (float): The parameter for repetition penalty. - 1.0 means no penalty - - **stop** (str | list[str] | None): To stop generating further - tokens. Only accept stop words that's encoded to one token idex. - - **response_format** (dict | None): To generate response according to given - schema. Examples: - - .. code-block:: json - - { - "type": "json_schema", - "json_schema":{ - "name": "test", - "schema":{ - "properties":{ - "name":{"type":"string"} - }, - "required":["name"], - "type":"object" - } - } - } - - or ``{"type": "regex_schema", "regex_schema": "call me [A-Za-z]{1,10}"}`` - - **logit_bias** (dict): Bias to logits. Only supported in pytorch engine. - - **tools** (list): A list of tools the model may call. Currently, only - internlm2 functions are supported as a tool. Use this to specify a - list of functions for which the model can generate JSON inputs. - - **tool_choice** (str | object): Controls which (if any) tool is called by - the model. `none` means the model will not call any tool and instead - generates a message. Specifying a particular tool via - ``{"type": "function", "function": {"name": "my_function"}}`` - forces the model to call that tool. `auto` or `required` will put all - the tools informationto the model. - - Additional arguments supported by LMDeploy: - - - **top_k** (int): The number of the highest probability vocabulary - tokens to keep for top-k-filtering - - **ignore_eos** (bool): indicator for ignoring eos - - **skip_special_tokens** (bool): Whether or not to remove special tokens - in the decoding. Default to be True. - - **spaces_between_special_tokens** (bool): Whether or not to add spaces - around special tokens. The behavior of Fast tokenizers is to have - this to False. This is setup to True in slow tokenizers. - - **min_new_tokens** (int): To generate at least numbers of tokens. - - **min_p** (float): Minimum token probability, which will be scaled by the - probability of the most likely token. It must be a value between - 0 and 1. Typical values are in the 0.01-0.2 range, comparably - selective as setting `top_p` in the 0.99-0.8 range (use the - opposite of normal `top_p` values) - - **repetition_ngram_size** (int): N-gram length for repetition early stop - (PyTorch engine). ``0`` disables. - - **repetition_ngram_threshold** (int): How many times that n-gram must - repeat to trigger early stop. ``0`` disables. - - Currently we do not support the following features: - - - **presence_penalty** (replaced with repetition_penalty) - - **frequency_penalty** (replaced with repetition_penalty) - """ - error_check_ret = check_request(request) - if error_check_ret is not None: - return error_check_ret - session = VariableInterface.create_session(request.session_id) - - # Resolve input: messages has priority over input_ids/image_data - messages_empty = (request.messages is None - or request.messages == '' - or (isinstance(request.messages, list) and len(request.messages) == 0)) - resolved_input_ids = None - if messages_empty and request.input_ids is not None: - # /generate-style input: use input_ids (+ optional image_data) - resolved_input_ids = request.input_ids - if request.image_data is not None: - # Convert image_data to OpenAI multimodal content format - image_data = request.image_data - image_input = [] - if not isinstance(image_data, list): - image_data = [image_data] - for img in image_data: - if isinstance(img, str): - image_input.append(dict(type='image_url', image_url=dict(url=img))) - else: - image_input.append(dict(type='image_url', image_url=img)) - text_input = dict(type='text', text=request.input_ids) - request.messages = [dict(role='user', content=[text_input] + image_input)] - resolved_input_ids = None # image_data conversion takes over - else: - # input_ids only — engine requires messages=None - request.messages = None - - json_request = await raw_request.json() - migration_request = json_request.pop('migration_request', None) - with_cache = json_request.pop('with_cache', False) - preserve_cache = json_request.pop('preserve_cache', False) - if migration_request: - migration_request = MigrationRequest.model_validate(migration_request) - - model_name = request.model - adapter_name = None - if model_name != VariableInterface.async_engine.model_name: - adapter_name = model_name # got a adapter name - request_id = f'chatcmpl-{shortuuid.random()}' - created_time = int(time.time()) - - tokenizer = VariableInterface.async_engine.tokenizer.model.model - gen_logprobs, logits_processors = None, None - if request.logprobs: - gen_logprobs = request.top_logprobs or 1 - elif request.return_logprob: - gen_logprobs = 1 - if request.logit_bias is not None: - try: - logits_processors = [ - logit_bias_logits_processor(request.logit_bias, tokenizer) - ] - except Exception as e: - return create_error_response(HTTPStatus.BAD_REQUEST, str(e)) - - parser_cls = VariableInterface.response_parser_cls - if request.tool_choice != 'none' and request.tools: - if parser_cls is None or parser_cls.tool_parser_cls is None: - return create_error_response( - HTTPStatus.BAD_REQUEST, - 'Please launch the api_server with --tool-call-parser if you want to use tool.') - - parser_cls = VariableInterface.response_parser_cls - try: - response_parser = parser_cls(request) - except ValueError as e: - raise HTTPException(status_code=HTTPStatus.BAD_REQUEST, detail=str(e)) - # request is normalized and may be adjusted by the parser - # (e.g. GPT-OSS clears response_format and injects the schema into messages) - request = response_parser.request - - gen_config = _build_serving_generation_config( - request, - logprobs=gen_logprobs, - stop_words=request.stop, - logits_processors=logits_processors, - random_seed=request.seed, - migration_request=migration_request, - with_cache=with_cache, - preserve_cache=preserve_cache, - ) - - # text completion for string input or input_ids - do_preprocess = (False if isinstance(request.messages, str) - or resolved_input_ids is not None else request.do_preprocess) - chat_template_kwargs = request.chat_template_kwargs or {} - if request.enable_thinking is not None: - logger.warning('`enable_thinking` will be deprecated in the future, ' - 'please use `chat_template_kwargs` instead.') - if chat_template_kwargs.get('enable_thinking') is None: - chat_template_kwargs['enable_thinking'] = request.enable_thinking - else: - logger.warning('`enable_thinking` in `chat_template_kwargs` will override the value in request.') - - result_generator = VariableInterface.async_engine.generate( - request.messages, - session, - gen_config=gen_config, - tools=request.tools, - reasoning_effort=request.reasoning_effort, - stream_response=True, # always use stream to enable batching - do_preprocess=do_preprocess, - adapter_name=adapter_name, - chat_template_kwargs=chat_template_kwargs or None, - input_ids=resolved_input_ids, - media_io_kwargs=request.media_io_kwargs, - mm_processor_kwargs=request.mm_processor_kwargs) - include_usage = bool(request.stream_options and request.stream_options.include_usage) - - def create_stream_response_json(index: int, - delta_message: DeltaMessage, - finish_reason: str | None = None, - logprobs: ChoiceLogprobs | None = None, - output_token_logprobs: list[tuple[float, int]] | None = None, - routed_experts=None, - output_ids=None) -> dict: - choice_data = ChatCompletionResponseStreamChoice(index=index, - delta=delta_message, - finish_reason=finish_reason, - logprobs=logprobs, - output_token_logprobs=output_token_logprobs, - output_ids=output_ids, - routed_experts=routed_experts) - choice_data = maybe_filter_parallel_tool_calls(choice_data, request) - response = ChatCompletionStreamResponse( - id=request_id, - created=created_time, - model=model_name, - choices=[choice_data], - usage=None, - ) - response_dict = response.model_dump(mode='json', exclude_none=True) - if include_usage: - response_dict['usage'] = None - return response_dict - - def create_stream_usage_response_json(usage: UsageInfo) -> str: - response = ChatCompletionStreamResponse( - id=request_id, - created=created_time, - model=model_name, - choices=[], - usage=usage, - ) - return response.model_dump_json(exclude_none=True) - - async def completion_stream_generator() -> AsyncGenerator[str, None]: - streaming_tools = False - final_usage = None - async for res in result_generator: - logprobs = None - output_token_logprobs = None - if request.logprobs and res.logprobs: - logprobs = _create_chat_completion_logprobs(tokenizer, res.token_ids, res.logprobs) - if request.return_logprob: - output_token_logprobs = _create_output_token_logprobs(res.token_ids, res.logprobs) - if res.finish_reason and include_usage: - final_usage = UsageInfo.build( - prompt_tokens=res.input_token_len, - completion_tokens=res.generate_token_len, - cached_tokens=res.cached_tokens, - ) - delta_token_ids = res.token_ids if res.token_ids is not None else [] - stream_deltas = response_parser.stream_chunk( - res.response, - delta_token_ids - ) - if not stream_deltas: - # Parser may buffer partial protocol tags and emit no visible delta - # while the engine still produced new tokens (e.g. MTP batch). Do not - # drop those token ids; emit them once on a placeholder delta. - if res.finish_reason is None and not delta_token_ids: - continue - stream_deltas = [(DeltaMessage(role='assistant', content=''), False)] - should_validate_complete = ( - res.finish_reason in ('stop', 'length') - and (request.return_token_ids or request.return_routed_experts) - ) - if should_validate_complete and not response_parser.validate_complete(): - res.finish_reason = 'parse_error' - - for delta_index, (delta_message, tool_emitted) in enumerate(stream_deltas): - if tool_emitted: - streaming_tools = True - - is_last_delta = delta_index == len(stream_deltas) - 1 - # The chat parser may split one engine yield into multiple protocol deltas, - # so attach the engine-level metadata to the last parsed delta. - finish_reason = res.finish_reason if is_last_delta else None - chunk_logprobs = logprobs if is_last_delta else None - chunk_output_token_logprobs = output_token_logprobs if is_last_delta else None - - if (request.tool_choice != 'none' and response_parser.tool_parser is not None): - if finish_reason == 'stop' and streaming_tools is True: - finish_reason = 'tool_calls' - - # Only output routed_experts in the final chunk - routed_experts = res.routed_experts if finish_reason is not None else None - # Emit token ids once per engine yield on the last parsed delta, when - # accumulated delta text and token ids for this step are aligned. - stream_output_ids = delta_token_ids if (request.return_token_ids and is_last_delta) else None - - response_json = create_stream_response_json(index=0, - delta_message=delta_message, - finish_reason=finish_reason, - logprobs=chunk_logprobs, - output_token_logprobs=chunk_output_token_logprobs, - routed_experts=routed_experts, - output_ids=stream_output_ids) - if res.cache_block_ids is not None and is_last_delta: - response_json['cache_block_ids'] = res.cache_block_ids - response_json['remote_token_ids'] = res.token_ids - yield f'data: {json.dumps(response_json)}\n\n' - if final_usage is not None: - yield f'data: {create_stream_usage_response_json(final_usage)}\n\n' - yield 'data: [DONE]\n\n' - - # Streaming response - if request.stream: - stream_generator = _with_request_cleanup(completion_stream_generator(), [result_generator], [session]) - return StreamingResponse(stream_generator, media_type='text/event-stream') - - # Non-streaming response - final_logprobs = [] - final_token_ids = [] - final_res = None - text = '' - cache_block_ids = [] - remote_token_ids = [] - async with aclosing(_with_request_cleanup(result_generator, [result_generator], [session])) as generator: - async for res in generator: - if await raw_request.is_disconnected(): - # Abort the request if the client disconnects. - await session.async_abort() - return create_error_response(HTTPStatus.BAD_REQUEST, 'Client disconnected') - final_res = res - text += res.response - if res.token_ids: - final_token_ids.extend(res.token_ids) - if res.logprobs: - final_logprobs.extend(res.logprobs) - cache_block_ids.append(res.cache_block_ids) - remote_token_ids.append(res.token_ids) - - tool_calls = None - reasoning_content = None - - try: - raw_text = text - text, tool_calls, reasoning_content = response_parser.parse_complete(text, final_token_ids) - should_validate_complete = ( - final_res.finish_reason in ('stop', 'length') - and (request.return_token_ids or request.return_routed_experts) - ) - if should_validate_complete and not response_parser.validate_complete(raw_text): - final_res.finish_reason = 'parse_error' - if isinstance(tool_calls, list) and len(tool_calls): - if final_res.finish_reason == 'stop': - final_res.finish_reason = 'tool_calls' - - except Exception as e: - logger.error(f'Failed to parse {text}. Exception: {e}.') - return create_error_response(HTTPStatus.BAD_REQUEST, 'Failed to parse fc related info to json format!') - - message = ChatMessage(role='assistant', - content=text, - tool_calls=tool_calls, - reasoning_content=reasoning_content) - - logprobs = None - if request.logprobs and len(final_logprobs): - logprobs = _create_chat_completion_logprobs(tokenizer, final_token_ids, final_logprobs) - output_token_logprobs = None - if request.return_logprob and len(final_logprobs): - output_token_logprobs = _create_output_token_logprobs(final_token_ids, final_logprobs) - - assert final_res is not None - choices = [] - choice_data = ChatCompletionResponseChoice( - index=0, - message=message, - logprobs=logprobs, - output_token_logprobs=output_token_logprobs, - finish_reason=final_res.finish_reason, - output_ids=final_token_ids if request.return_token_ids else None, - routed_experts=final_res.routed_experts if request.return_routed_experts else None, - ) - choice_data = maybe_filter_parallel_tool_calls(choice_data, request) - choices.append(choice_data) - - if with_cache: - cache_block_ids = cache_block_ids[0] - remote_token_ids = [remote_token_ids[0][-1]] - - usage = UsageInfo.build( - prompt_tokens=final_res.input_token_len, - completion_tokens=final_res.generate_token_len, - cached_tokens=final_res.cached_tokens, - ) - response = ChatCompletionResponse( - id=request_id, - created=created_time, - model=model_name, - choices=choices, - usage=usage, - ).model_dump() - - if with_cache: - response['cache_block_ids'] = cache_block_ids - response['remote_token_ids'] = remote_token_ids - - return response - - -@router.post('/v1/completions', dependencies=[Depends(validate_json_request)]) -async def completions_v1(request: CompletionRequest, raw_request: Request = None): - """Completion API similar to OpenAI's API. - - Go to https://platform.openai.com/docs/api-reference/completions/create - for the API specification. - - The request should be a JSON object with the following fields: - - - **model** (str): model name. Available from /v1/models. - - **prompt** (str): the input prompt. - - **suffix** (str): The suffix that comes after a completion of inserted text. - - **max_completion_tokens** (int | None): output token nums. Default to None. - - **max_tokens** (int | None): output token nums. Default to 16. - Deprecated: Use max_completion_tokens instead. - - **temperature** (float): to modulate the next token probability - - **top_p** (float): If set to float < 1, only the smallest set of most - probable tokens with probabilities that add up to top_p or higher - are kept for generation. - - **n** (int): How many chat completion choices to generate for each input - message. **Only support one here**. - - **stream**: whether to stream the results or not. Default to false. - - **stream_options**: Options for streaming response. Only set this when you - set stream: true. - - **repetition_penalty** (float): The parameter for repetition penalty. - 1.0 means no penalty - - **user** (str): A unique identifier representing your end-user. - - **stop** (str | list[str] | None): To stop generating further - tokens. Only accept stop words that's encoded to one token idex. - - Additional arguments supported by LMDeploy: - - - **ignore_eos** (bool): indicator for ignoring eos - - **skip_special_tokens** (bool): Whether or not to remove special tokens - in the decoding. Default to be True. - - **spaces_between_special_tokens** (bool): Whether or not to add spaces - around special tokens. The behavior of Fast tokenizers is to have - this to False. This is setup to True in slow tokenizers. - - **top_k** (int): The number of the highest probability vocabulary - tokens to keep for top-k-filtering - - **min_p** (float): Minimum token probability, which will be scaled by the - probability of the most likely token. It must be a value between - 0 and 1. Typical values are in the 0.01-0.2 range, comparably - selective as setting `top_p` in the 0.99-0.8 range (use the - opposite of normal `top_p` values) - - **repetition_ngram_size** (int): N-gram length for repetition early stop - (PyTorch engine). ``0`` disables. - - **repetition_ngram_threshold** (int): How many times that n-gram must - repeat to trigger early stop. ``0`` disables. - - Currently we do not support the following features: - - - **logprobs** (not supported yet) - - **presence_penalty** (replaced with repetition_penalty) - - **frequency_penalty** (replaced with repetition_penalty) - """ - error_check_ret = check_request(request) - if error_check_ret is not None: - return error_check_ret - json_request = await raw_request.json() - migration_request = json_request.pop('migration_request', None) - with_cache = json_request.pop('with_cache', False) - preserve_cache = json_request.pop('preserve_cache', False) - if migration_request: - migration_request = MigrationRequest.model_validate(migration_request) - - model_name = request.model - adapter_name = None - if model_name != VariableInterface.async_engine.model_name: - adapter_name = model_name # got a adapter name - request_id = str(request.session_id) - created_time = int(time.time()) - sessions = [] - if isinstance(request.prompt, str): - request.prompt = [request.prompt] - sessions.append(VariableInterface.create_session(request.session_id)) - elif isinstance(request.prompt, list): - for _ in request.prompt: - sessions.append(VariableInterface.create_session()) - if isinstance(request.stop, str): - request.stop = [request.stop] - - gen_config = _build_serving_generation_config( - request, - stop_words=request.stop, - random_seed=request.seed, - migration_request=migration_request, - with_cache=with_cache, - preserve_cache=preserve_cache, - ) - generators = [] - for prompt, session in zip(request.prompt, sessions): - result_generator = VariableInterface.async_engine.generate( - prompt, - session, - gen_config=gen_config, - stream_response=True, # always use stream to enable batching - do_preprocess=False, - adapter_name=adapter_name) - generators.append(result_generator) - - include_usage = bool(request.stream_options and request.stream_options.include_usage) - - def create_stream_response_json(index: int, - text: str, - finish_reason: str | None = None, - logprobs: LogProbs | None = None) -> dict: - choice_data = CompletionResponseStreamChoice(index=index, - text=text, - finish_reason=finish_reason, - logprobs=logprobs) - response = CompletionStreamResponse( - id=request_id, - created=created_time, - model=model_name, - choices=[choice_data], - usage=None, - ) - response_json = response.model_dump(mode='json', exclude_none=True) - if include_usage: - response_json['usage'] = None - return response_json - - def create_stream_usage_response_json(usage: UsageInfo) -> dict: - response = CompletionStreamResponse( - id=request_id, - created=created_time, - model=model_name, - choices=[], - usage=usage, - ) - return response.model_dump(mode='json', exclude_none=True) - - async def completion_stream_generator() -> AsyncGenerator[str, None]: - # First chunk with role - final_usage = UsageInfo() if include_usage else None - for generator in generators: - async for res in generator: - logprobs = None - if request.logprobs and res.logprobs: - raise ValueError('logprobs is removed') - if res.finish_reason and final_usage is not None: - final_res = res - total_tokens = sum([final_res.input_token_len, final_res.generate_token_len]) - final_usage.prompt_tokens += final_res.input_token_len - final_usage.completion_tokens += final_res.generate_token_len - final_usage.total_tokens += total_tokens - response_json = create_stream_response_json(index=0, - text=res.response, - finish_reason=res.finish_reason, - logprobs=logprobs) - if res.cache_block_ids is not None: - response_json['cache_block_ids'] = res.cache_block_ids - response_json['remote_token_ids'] = res.token_ids - yield f'data: {json.dumps(response_json)}\n\n' - if final_usage is not None: - yield f'data: {json.dumps(create_stream_usage_response_json(final_usage))}\n\n' - yield 'data: [DONE]\n\n' - - # Streaming response - if request.stream: - stream_generator = _with_request_cleanup(completion_stream_generator(), generators, sessions) - return StreamingResponse(stream_generator, media_type='text/event-stream') - - # Non-streaming response - usage = UsageInfo() - choices = [None] * len(generators) - cache_block_ids = [] - remote_token_ids = [] - - async def _inner_call(i, generator, session): - nonlocal cache_block_ids, remote_token_ids - final_logprobs = [] - final_token_ids = [] - final_res = None - text = '' - async with aclosing(_with_request_cleanup(generator, [generator], [session])) as cleanup_generator: - async for res in cleanup_generator: - if await raw_request.is_disconnected(): - # Abort the request if the client disconnects. - await VariableInterface.async_engine.stop_session(request.session_id) - return create_error_response(HTTPStatus.BAD_REQUEST, 'Client disconnected') - final_res = res - text += res.response - cache_block_ids.append(res.cache_block_ids) - remote_token_ids.append(res.token_ids) - if res.token_ids: - final_token_ids.extend(res.token_ids) - if res.logprobs: - final_logprobs.extend(res.logprobs) - - logprobs = None - assert final_res is not None - choice_data = CompletionResponseChoice(index=i, - text=text, - finish_reason=final_res.finish_reason, - logprobs=logprobs) - choices[i] = choice_data - - if with_cache: - cache_block_ids = cache_block_ids[0] - remote_token_ids = [remote_token_ids[0][-1]] - - total_tokens = sum([final_res.input_token_len, final_res.generate_token_len]) - usage.prompt_tokens += final_res.input_token_len - usage.completion_tokens += final_res.generate_token_len - usage.total_tokens += total_tokens - - await asyncio.gather(*[_inner_call(i, generators[i], sessions[i]) for i in range(len(generators))]) - - response = CompletionResponse( - id=request_id, - created=created_time, - model=model_name, - choices=choices, - usage=usage, - ).model_dump() - - if with_cache: - response['cache_block_ids'] = cache_block_ids - response['remote_token_ids'] = remote_token_ids - - return response - - -@router.post('/generate', dependencies=[Depends(validate_json_request)]) -async def generate(request: GenerateReqInput, raw_request: Request = None): - error_check_ret = check_request(request) - if error_check_ret is not None: - return error_check_ret - - session = VariableInterface.create_session(request.session_id) - - prompt = request.prompt - input_ids = request.input_ids - image_data = request.image_data - if image_data is not None: - # convert to openai format - image_input = [] - if not isinstance(image_data, list): - image_data = [image_data] - for img in image_data: - if isinstance(img, str): - image_input.append(dict(type='image_url', image_url=dict(url=img))) - else: - image_input.append(dict(type='image_url', image_url=img)) - text_input = dict(type='text', text=prompt if prompt else input_ids) - prompt = [dict(role='user', content=[text_input] + image_input)] - input_ids = None - - gen_config = _build_serving_generation_config( - request, - logprobs=1 if request.return_logprob else None, - stop_words=request.stop, - ) - - result_generator = VariableInterface.async_engine.generate( - messages=prompt, - session_id=session, - input_ids=input_ids, - gen_config=gen_config, - stream_response=True, # always use stream to enable batching - do_preprocess=False, - media_io_kwargs=request.media_io_kwargs, - mm_processor_kwargs=request.mm_processor_kwargs) - - def create_generate_response_json(res, text, output_ids, logprobs, finish_reason, routed_experts=None): - # only output router experts in last chunk - routed_experts = None if finish_reason is None else routed_experts - meta = GenerateReqMetaOutput(finish_reason=dict(type=finish_reason) if finish_reason else None, - output_token_logprobs=logprobs or None, - prompt_tokens=res.input_token_len, - routed_experts=routed_experts, - completion_tokens=res.generate_token_len) - - response = GenerateReqOutput(text=text, output_ids=output_ids, meta_info=meta, routed_experts=routed_experts) - return response.model_dump_json() - - async def generate_stream_generator(): - async for res in result_generator: - text = res.response or '' - output_ids = res.token_ids - routed_experts = res.routed_experts - logprobs = [] - if res.logprobs: - for tok, tok_logprobs in zip(res.token_ids, res.logprobs): - logprobs.append((tok_logprobs[tok], tok)) - response_json = create_generate_response_json(res, - text, - output_ids, - logprobs, - res.finish_reason, - routed_experts=routed_experts) - yield f'data: {response_json}\n\n' - yield 'data: [DONE]\n\n' - - if request.stream: - stream_generator = _with_request_cleanup(generate_stream_generator(), [result_generator], [session]) - return StreamingResponse(stream_generator, media_type='text/event-stream') - - response = None - - async def _inner_call(): - text = '' - output_ids = [] - logprobs = [] - async with aclosing(_with_request_cleanup(result_generator, [result_generator], [session])) as generator: - async for res in generator: - if await raw_request.is_disconnected(): - # Abort the request if the client disconnects. - await session.async_abort() - return create_error_response(HTTPStatus.BAD_REQUEST, 'Client disconnected') - text += res.response or '' - output_ids.extend(res.token_ids or []) - logprobs.extend(res.logprobs or []) - - output_token_logprobs = [] - if len(logprobs) and len(output_ids): - for tok, tok_logprobs in zip(output_ids, logprobs): - output_token_logprobs.append((tok_logprobs[tok], tok)) - - nonlocal response - meta = GenerateReqMetaOutput(finish_reason=dict(type=res.finish_reason) if res.finish_reason else None, - output_token_logprobs=output_token_logprobs or None, - prompt_tokens=res.input_token_len, - routed_experts=res.routed_experts, - completion_tokens=res.generate_token_len) - response = GenerateReqOutput(text=text, output_ids=output_ids, meta_info=meta) - - await _inner_call() - return response - - -@router.post('/v1/embeddings', tags=['unsupported']) -async def create_embeddings(request: EmbeddingsRequest, raw_request: Request = None): - """Creates embeddings for the text.""" - return create_error_response(HTTPStatus.BAD_REQUEST, 'Unsupported by turbomind.') - - -@router.post('/v1/encode', dependencies=[Depends(validate_json_request)]) -async def encode(request: EncodeRequest, raw_request: Request = None): - """Encode prompts. - - The request should be a JSON object with the following fields: - - - **input**: the prompt to be encoded. In str or list[str] format. - - **do_preprocess**: whether do preprocess or not. Default to False. - - **add_bos**: Whether to add a BOS token when encoding. Default to True. - """ - - def encode(prompt: str, do_preprocess: bool, add_bos: bool): - if do_preprocess: - prompt = VariableInterface.async_engine.chat_template.get_prompt(prompt) - input_ids = VariableInterface.async_engine.tokenizer.encode(prompt, add_bos=add_bos) - return input_ids - - if isinstance(request.input, str): - encoded = encode(request.input, request.do_preprocess, request.add_bos) - return EncodeResponse(input_ids=encoded, length=len(encoded)) - else: - encoded, length = [], [] - for prompt in request.input: - ids = encode(prompt, request.do_preprocess, request.add_bos) - encoded.append(ids) - length.append(len(ids)) - return EncodeResponse(input_ids=encoded, length=length) - - -@router.post('/pooling', dependencies=[Depends(validate_json_request)]) -async def pooling(request: PoolingRequest, raw_request: Request = None): - """Pooling prompts for reward model. - - In vLLM documentation, https://docs.vllm.ai/en/latest/serving/openai_compatible_server.html#pooling-api_1, - the input format of Pooling API is the same as Embeddings API. - - Go to https://platform.openai.com/docs/api-reference/embeddings/create - for the Embeddings API specification. - - The request should be a JSON object with the following fields: - - - **model** (str): model name. Available from /v1/models. - - **input** (list[int] | list[list[int]] | str | list[str]): input text to be embed - """ - - async_engine = VariableInterface.async_engine - - request_input = request.input - model_name = request.model or async_engine.model_name - - # Normalize all inputs to be a batch (List[List[int]]) - if isinstance(request_input, str): - input_ids = [async_engine.tokenizer.encode(request_input)] - elif isinstance(request_input, list): - if not request_input: - return create_error_response(HTTPStatus.BAD_REQUEST, 'Input list cannot be empty.') - if isinstance(request_input[0], str): # list[str] - input_ids = [async_engine.tokenizer.encode(p) for p in request_input] - elif isinstance(request_input[0], int): # list[int] - input_ids = [request_input] - elif isinstance(request_input[0], list): # list[list[int]] - input_ids = request_input - else: - return create_error_response(HTTPStatus.BAD_REQUEST, 'Input list contains an invalid type.') - else: - return create_error_response(HTTPStatus.BAD_REQUEST, 'Input must be a string or a list.') - - batch_scores = await async_engine.async_get_reward_score(input_ids) - prompt_tokens = sum(len(ids) for ids in input_ids) - usage = UsageInfo.build(prompt_tokens=prompt_tokens, completion_tokens=0, cached_tokens=0) - - data = [] - for i, score in enumerate(batch_scores): - data.append({ - 'index': i, - 'object': 'pooling', - 'data': score, - }) - - response = PoolingResponse(model=model_name, data=data, usage=usage) - return response.model_dump() - - -@router.post('/get_ppl', dependencies=[Depends(validate_json_request)]) -async def get_ppl(request: PPLRequest, raw_request: Request = None): - """Get the perplexity (mean cross-entropy loss) of the input prompt. - - The request should be a JSON object with the following fields: - - - **input** (str | list[int]): the input to score, either raw text or token - ids. Text is tokenized with ``tokenizer.encode`` (no chat template is - applied). - """ - async_engine = VariableInterface.async_engine - - request_input = request.input - - # pydantic already validated `input` as `str | list[int]`; text -> - # tokenizer.encode, otherwise the token ids are used as-is - if isinstance(request_input, str): - input_ids = async_engine.tokenizer.encode(request_input) - else: - input_ids = request_input - if not input_ids: - return create_error_response(HTTPStatus.BAD_REQUEST, 'Input must not be empty.') - - try: - ppl = await async_engine.async_get_ppl(input_ids) - except ValueError as e: - return create_error_response(HTTPStatus.BAD_REQUEST, str(e)) - - response = PPLResponse(ppl=ppl) - return response.model_dump() - - -@router.post('/update_weights', dependencies=[Depends(validate_json_request)]) -def update_params(request: UpdateParamsRequest, raw_request: Request = None): - """Update weights for the model.""" - VariableInterface.async_engine.engine.update_params(request) - return JSONResponse(content=None) - - -def _check_pytorch_backend_for_disagg_weight_update(): - """Disaggregated weight-update endpoints are PyTorch-backend only for - now.""" - backend = getattr(VariableInterface.async_engine, 'backend', None) - if backend != 'pytorch': - return create_error_response( - HTTPStatus.NOT_IMPLEMENTED, - f'Disaggregated weight-update endpoints require backend="pytorch", got {backend!r}.') - return None - - -@router.post('/init_weights_update_group', dependencies=[Depends(validate_json_request)]) -async def init_weights_update_group(request: InitWeightsUpdateGroupRequest, raw_request: Request = None): - """Initialize the torch.distributed process group used by an external - trainer to broadcast weights into this rollout engine.""" - err = _check_pytorch_backend_for_disagg_weight_update() - if err is not None: - return err - success, message = await VariableInterface.async_engine.engine.init_weights_update_group(request) - content = {'success': success, 'message': message} - return JSONResponse(content=content, status_code=200 if success else HTTPStatus.BAD_REQUEST) - - -@router.post('/update_weights_from_distributed', dependencies=[Depends(validate_json_request)]) -async def update_weights_from_distributed(request: UpdateWeightsFromDistributedRequest, raw_request: Request = None): - """Receive a bucket of weights through a previously initialized weights- - update group and load them into the running model.""" - err = _check_pytorch_backend_for_disagg_weight_update() - if err is not None: - return err - success, message = await VariableInterface.async_engine.engine.update_weights_from_distributed(request) - content = {'success': success, 'message': message} - return JSONResponse(content=content, status_code=200 if success else HTTPStatus.BAD_REQUEST) - - -@router.post('/destroy_weights_update_group', dependencies=[Depends(validate_json_request)]) -async def destroy_weights_update_group(request: DestroyWeightsUpdateGroupRequest, raw_request: Request = None): - """Tear down a previously initialized weights-update group.""" - err = _check_pytorch_backend_for_disagg_weight_update() - if err is not None: - return err - success, message = await VariableInterface.async_engine.engine.destroy_weights_update_group(request) - content = {'success': success, 'message': message} - return JSONResponse(content=content, status_code=200 if success else HTTPStatus.BAD_REQUEST) - - -@router.post('/sleep', dependencies=[Depends(validate_json_request)]) -async def sleep(raw_request: Request = None): - level = raw_request.query_params.get('level', '1') - try: - level = int(level) - except (TypeError, ValueError): - return create_error_response(HTTPStatus.BAD_REQUEST, 'The "level" query parameter must be an integer.') - if level not in (1, 2): - return create_error_response(HTTPStatus.BAD_REQUEST, 'The "level" query parameter must be 1 or 2.') - async_engine = VariableInterface.async_engine - await async_engine.sleep(level) - return Response(status_code=200) - - -@router.post('/wakeup', dependencies=[Depends(validate_json_request)]) -async def wakeup(raw_request: Request = None): - tags = raw_request.query_params.getlist('tags') - tags = tags or None - VariableInterface.async_engine.wakeup(tags) - return Response(status_code=200) - - -@router.get('/is_sleeping') -async def is_sleeping(): - is_sleeping = VariableInterface.async_engine.is_sleeping - return JSONResponse(content={'is_sleeping': is_sleeping}) - - -""" PD Disaggregation API Begin """ - - -@router.get('/distserve/engine_info') -async def engine_info(): - engine_config = VariableInterface.async_engine.backend_config - - response = DistServeEngineConfig(tp_size=engine_config.tp, - dp_size=engine_config.dp, - pp_size=None, - ep_size=engine_config.ep, - dp_rank=engine_config.dp_rank, - block_size=engine_config.block_size, - num_cpu_blocks=engine_config.num_cpu_blocks, - num_gpu_blocks=engine_config.num_gpu_blocks) - - return response.model_dump_json() - - -@router.post('/distserve/p2p_initialize') -async def p2p_initialize(init_request: DistServeInitRequest): - return VariableInterface.async_engine.p2p_initialize(init_request) - - -@router.post('/distserve/p2p_connect') -async def p2p_connect(conn_request: DistServeConnectionRequest): - return VariableInterface.async_engine.p2p_connect(conn_request) - - -@router.post('/distserve/p2p_drop_connect') -async def p2p_drop_connect(drop_conn_request: DistServeDropConnectionRequest): - return VariableInterface.async_engine.p2p_drop_connect(drop_conn_request) - - -@router.post('/distserve/free_cache') -async def free_cache(cache_free_request: DistServeCacheFreeRequest) -> JSONResponse: - session_id = cache_free_request.remote_session_id - VariableInterface.async_engine.free_cache(session_id) - return {'status': 'SUCCESS'} - - -""" PD Disaggregation API End """ - - -@router.post('/abort_request') -async def abort_request(request: AbortRequest, raw_request: Request = None): - """Abort an ongoing request.""" - if not VariableInterface.enable_abort_handling: - return Response( - status_code=501, - content='This server does not support abort requests. Enable with --enable-abort-handling flag.') - - if request.abort_all: - await VariableInterface.async_engine.stop_all_session() - else: - session = VariableInterface.find_session(request.session_id) - if session is None: - return create_error_response(HTTPStatus.BAD_REQUEST, f'Session {request.session_id} not found.') - await session.async_abort() - session_mgr = VariableInterface.get_session_manager() - session_mgr.remove(session) - return Response(status_code=200) - - -@router.post('/v1/chat/interactive', dependencies=[Depends(validate_json_request)], include_in_schema=False) -async def chat_interactive_v1(request, raw_request: Request = None): - return create_error_response( - HTTPStatus.BAD_REQUEST, 'v1/chat/interactive is deprecated, please launch server with --enable-prefix-cache ' - 'and use /v1/chat/completions instead.') +server_context = ServerContext() +router = create_openai_router(server_context) def handle_torchrun(): @@ -1368,42 +108,55 @@ def dummy_get_device_id(): from lmdeploy.vl.model.utils import _set_func # the replacement can't be recovered - _set_func('mmengine.logging.logger._get_device_id', dummy_get_device_id) + _set_func('mmengine.logging.logger._get_device_id', + dummy_get_device_id) @router.on_event('startup') async def startup_event(): - async_engine = VariableInterface.async_engine + async_engine = server_context.async_engine async_engine.start_loop(asyncio.get_running_loop(), use_async_api=True) - if VariableInterface.proxy_url is None: + if server_context.proxy_url is None: return elif getattr(async_engine.engine, 'is_dummy', False): logger.info('Dummy node started') return try: import requests - engine_config = VariableInterface.async_engine.backend_config - engine_role = engine_config.role.value if hasattr(engine_config, 'role') else 1 - url = f'{VariableInterface.proxy_url}/nodes/add' - data = {'url': VariableInterface.api_server_url, 'status': {'models': get_model_list(), 'role': engine_role}} - headers = {'accept': 'application/json', 'Content-Type': 'application/json'} + engine_config = server_context.engine_config + engine_role = engine_config.role.value if hasattr( + engine_config, 'role') else 1 + url = f'{server_context.proxy_url}/nodes/add' + data = { + 'url': server_context.api_server_url, + 'status': { + 'models': get_model_list(server_context), + 'role': engine_role + } + } + headers = { + 'accept': 'application/json', + 'Content-Type': 'application/json' + } response = requests.post(url, headers=headers, json=data) if response.status_code != 200: - raise HTTPException(status_code=response.status_code, detail=response.text) + raise HTTPException(status_code=response.status_code, + detail=response.text) except Exception as e: logger.error(f'Service registration failed: {e}') @router.on_event('shutdown') async def shutdown_event(): - async_engine = VariableInterface.async_engine + async_engine = server_context.async_engine if async_engine is not None: async_engine.close() -async def validation_exception_handler(request: Request, exc: RequestValidationError): +async def validation_exception_handler(request: Request, + exc: RequestValidationError): """Handler for RequestValidationError.""" return JSONResponse( status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, @@ -1426,23 +179,29 @@ async def dispatch(self, request: Request, call_next): return response -def set_parsers(reasoning_parser_name: str | None = None, tool_parser_name: str | None = None, **kwargs): +def set_parsers(context: ServerContext, + reasoning_parser_name: str | None = None, + tool_parser_name: str | None = None, + **kwargs): from lmdeploy.serve.parsers import ResponseParserManager name = 'default' - arch = VariableInterface.async_engine.arch + arch = context.async_engine.arch if arch == 'GptOssForCausalLM': name = 'gpt-oss' cls = ResponseParserManager.get(name) if cls is None: - raise ValueError(f'The response parser {name} is not in the parser list: ' - f'{ResponseParserManager.module_dict.keys()}') - tokenizer = VariableInterface.async_engine.tokenizer.model.model + raise ValueError( + f'The response parser {name} is not in the parser list: ' + f'{ResponseParserManager.module_dict.keys()}') + tokenizer = context.async_engine.tokenizer.model.model cls.set_parsers(reasoning_parser_name=reasoning_parser_name, tool_parser_name=tool_parser_name, tokenizer=tokenizer) - VariableInterface.response_parser_cls = cls + context.response_parser_cls = cls + -def mount_metrics(app: FastAPI, backend_config: PytorchEngineConfig | TurbomindEngineConfig): +def mount_metrics(app: FastAPI, + backend_config: PytorchEngineConfig | TurbomindEngineConfig): if not getattr(backend_config, 'enable_metrics', False): return @@ -1457,14 +216,17 @@ def mount_metrics(app: FastAPI, backend_config: PytorchEngineConfig | TurbomindE app.routes.append(metrics_route) -def create_lifespan_handler(backend_config: PytorchEngineConfig | TurbomindEngineConfig, async_engine: AsyncEngine): +def create_lifespan_handler(backend_config: PytorchEngineConfig + | TurbomindEngineConfig, + async_engine: AsyncEngine, + context: ServerContext): """Factory function to create a lifespan handler.""" @asynccontextmanager async def lifespan_handler(app: FastAPI): task = None health_monitor = EngineHealthMonitor(async_engine) - VariableInterface.health_monitor = health_monitor + context.health_monitor = health_monitor try: health_monitor.start() if getattr(backend_config, 'enable_metrics', False): @@ -1476,8 +238,10 @@ async def _force_log(): await asyncio.sleep(log_interval) # periodically update schedule metrics, as they change less frequently than iteration stats - schedule_metrics = await async_engine.get_schedule_metrics() - await metrics_processor.update_schedule_stats(schedule_metrics) + schedule_metrics = await async_engine.get_schedule_metrics( + ) + await metrics_processor.update_schedule_stats( + schedule_metrics) await async_engine.do_log_stats() @@ -1486,8 +250,8 @@ async def _force_log(): yield finally: await health_monitor.stop() - if VariableInterface.health_monitor is health_monitor: - VariableInterface.health_monitor = None + if context.health_monitor is health_monitor: + context.health_monitor = None if task: task.cancel() await metrics_processor.stop_metrics_handler() @@ -1498,7 +262,8 @@ async def _force_log(): def serve(model_path: str, model_name: str | None = None, backend: Literal['turbomind', 'pytorch'] = 'turbomind', - backend_config: PytorchEngineConfig | TurbomindEngineConfig | None = None, + backend_config: PytorchEngineConfig | TurbomindEngineConfig + | None = None, chat_template_config: ChatTemplateConfig | None = None, server_name: str = '0.0.0.0', server_port: int = 23333, @@ -1579,12 +344,13 @@ def serve(model_path: str, os.environ['TM_LOG_LEVEL'] = log_level logger.setLevel(log_level) - VariableInterface.allow_terminate_by_client = allow_terminate_by_client - VariableInterface.enable_abort_handling = enable_abort_handling + server_context.allow_terminate_by_client = allow_terminate_by_client + server_context.enable_abort_handling = enable_abort_handling from lmdeploy.serve.parsers import validate_parser_names - reasoning_parser, tool_call_parser = validate_parser_names(reasoning_parser, tool_call_parser) + reasoning_parser, tool_call_parser = validate_parser_names( + reasoning_parser, tool_call_parser) - VariableInterface.default_gen_config = resolve_default_gen_config( + server_context.default_gen_config = resolve_default_gen_config( generation_config, model_path, trust_remote_code, @@ -1606,30 +372,37 @@ def serve(model_path: str, # router replay if backend_config.enable_return_routed_experts: backend_config.enable_transfer_obj_ref = True - VariableInterface.async_engine = pipeline_class(model_path=model_path, - model_name=model_name, - backend=backend, - backend_config=backend_config, - chat_template_config=chat_template_config, - max_log_len=max_log_len, - trust_remote_code=trust_remote_code, - speculative_config=speculative_config, - allowed_media_domains=allowed_media_domains, - **kwargs) - set_parsers(reasoning_parser, tool_call_parser) + server_context.async_engine = pipeline_class( + model_path=model_path, + model_name=model_name, + backend=backend, + backend_config=backend_config, + chat_template_config=chat_template_config, + max_log_len=max_log_len, + trust_remote_code=trust_remote_code, + speculative_config=speculative_config, + allowed_media_domains=allowed_media_domains, + **kwargs) + set_parsers(server_context, reasoning_parser, tool_call_parser) # create FastAPI lifespan events - lifespan = create_lifespan_handler(backend_config, VariableInterface.async_engine) + lifespan = create_lifespan_handler(backend_config, + server_context.async_engine, + server_context) if disable_fastapi_docs: - app = FastAPI(docs_url=None, redoc_url=None, openapi_url=None, lifespan=lifespan) + app = FastAPI(docs_url=None, + redoc_url=None, + openapi_url=None, + lifespan=lifespan) else: app = FastAPI(docs_url='/', lifespan=lifespan) app.include_router(router) - app.include_router(create_anthropic_router(VariableInterface)) - app.include_router(create_responses_router(VariableInterface)) - app.add_exception_handler(RequestValidationError, validation_exception_handler) + app.include_router(create_anthropic_router(server_context)) + app.include_router(create_responses_router(server_context)) + app.add_exception_handler(RequestValidationError, + validation_exception_handler) mount_metrics(app, backend_config) if allow_origins: @@ -1645,28 +418,33 @@ def serve(model_path: str, app.add_middleware(AuthenticationMiddleware, tokens=tokens) def is_engine_sleeping() -> bool: - eng = VariableInterface.async_engine + eng = server_context.async_engine return eng is not None and eng.is_sleeping - app.add_middleware(EngineSleepingMiddleware, is_sleeping=is_engine_sleeping) + + app.add_middleware(EngineSleepingMiddleware, + is_sleeping=is_engine_sleeping) # set the maximum number of concurrent requests if max_concurrent_requests is not None: - app.add_middleware(ConcurrencyLimitMiddleware, max_concurrent_requests=max_concurrent_requests) + app.add_middleware(ConcurrencyLimitMiddleware, + max_concurrent_requests=max_concurrent_requests) if proxy_url is not None: - VariableInterface.proxy_url = proxy_url - VariableInterface.api_server_url = f'{http_or_https}://{server_name}:{server_port}' # noqa + server_context.proxy_url = proxy_url + server_context.api_server_url = f'{http_or_https}://{server_name}:{server_port}' # noqa for i in range(3): - print(f'HINT: Please open \033[93m\033[1m{http_or_https}://' - f'{server_name}:{server_port}\033[0m in a browser for detailed api' - ' usage!!!') + print( + f'HINT: Please open \033[93m\033[1m{http_or_https}://' + f'{server_name}:{server_port}\033[0m in a browser for detailed api' + ' usage!!!') uvicorn.run(app=app, host=server_name, port=server_port, log_level=os.getenv('UVICORN_LOG_LEVEL', 'info').lower(), ssl_keyfile=ssl_keyfile, ssl_certfile=ssl_certfile, - timeout_keep_alive=int(os.environ.get('UVICORN_TIMEOUT_KEEP_ALIVE', 5))) + timeout_keep_alive=int( + os.environ.get('UVICORN_TIMEOUT_KEEP_ALIVE', 5))) if __name__ == '__main__': diff --git a/lmdeploy/serve/openai/endpoints/__init__.py b/lmdeploy/serve/openai/endpoints/__init__.py new file mode 100644 index 0000000000..d3faeb7a17 --- /dev/null +++ b/lmdeploy/serve/openai/endpoints/__init__.py @@ -0,0 +1,6 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""OpenAI and LMDeploy endpoint router assembly.""" + +from .router import create_openai_router + +__all__ = ['create_openai_router'] diff --git a/lmdeploy/serve/openai/endpoints/auxiliary.py b/lmdeploy/serve/openai/endpoints/auxiliary.py new file mode 100644 index 0000000000..2f5eb747b2 --- /dev/null +++ b/lmdeploy/serve/openai/endpoints/auxiliary.py @@ -0,0 +1,153 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from __future__ import annotations + +from http import HTTPStatus + +from fastapi import APIRouter, Depends, Request + +from lmdeploy.serve.openai.protocol import ( + EmbeddingsRequest, + EncodeRequest, + EncodeResponse, + PoolingRequest, + PoolingResponse, + PPLRequest, + PPLResponse, + UsageInfo, +) +from lmdeploy.serve.openai.utils import create_error_response +from lmdeploy.serve.utils.server_utils import validate_json_request + + +def register(router: APIRouter, server_context) -> None: + + @router.post('/v1/embeddings', tags=['unsupported']) + async def create_embeddings(request: EmbeddingsRequest, + raw_request: Request = None): + """Creates embeddings for the text.""" + return create_error_response(HTTPStatus.BAD_REQUEST, + 'Unsupported by turbomind.') + + @router.post('/v1/encode', dependencies=[Depends(validate_json_request)]) + async def encode(request: EncodeRequest, raw_request: Request = None): + """Encode prompts. + + The request should be a JSON object with the following fields: + + - **input**: the prompt to be encoded. In str or list[str] format. + - **do_preprocess**: whether do preprocess or not. Default to False. + - **add_bos**: Whether to add a BOS token when encoding. Default to True. + """ + + def encode(prompt: str, do_preprocess: bool, add_bos: bool): + if do_preprocess: + prompt = server_context.async_engine.chat_template.get_prompt( + prompt) + input_ids = server_context.async_engine.tokenizer.encode( + prompt, add_bos=add_bos) + return input_ids + + if isinstance(request.input, str): + encoded = encode(request.input, request.do_preprocess, + request.add_bos) + return EncodeResponse(input_ids=encoded, length=len(encoded)) + else: + encoded, length = [], [] + for prompt in request.input: + ids = encode(prompt, request.do_preprocess, request.add_bos) + encoded.append(ids) + length.append(len(ids)) + return EncodeResponse(input_ids=encoded, length=length) + + @router.post('/pooling', dependencies=[Depends(validate_json_request)]) + async def pooling(request: PoolingRequest, raw_request: Request = None): + """Pooling prompts for reward model. + + In vLLM documentation, https://docs.vllm.ai/en/latest/serving/openai_compatible_server.html#pooling-api_1, + the input format of Pooling API is the same as Embeddings API. + + Go to https://platform.openai.com/docs/api-reference/embeddings/create + for the Embeddings API specification. + + The request should be a JSON object with the following fields: + + - **model** (str): model name. Available from /v1/models. + - **input** (list[int] | list[list[int]] | str | list[str]): input text to be embed + """ + + async_engine = server_context.async_engine + + request_input = request.input + model_name = request.model or async_engine.model_name + + # Normalize all inputs to be a batch (List[List[int]]) + if isinstance(request_input, str): + input_ids = [async_engine.tokenizer.encode(request_input)] + elif isinstance(request_input, list): + if not request_input: + return create_error_response(HTTPStatus.BAD_REQUEST, + 'Input list cannot be empty.') + if isinstance(request_input[0], str): # list[str] + input_ids = [ + async_engine.tokenizer.encode(p) for p in request_input + ] + elif isinstance(request_input[0], int): # list[int] + input_ids = [request_input] + elif isinstance(request_input[0], list): # list[list[int]] + input_ids = request_input + else: + return create_error_response( + HTTPStatus.BAD_REQUEST, + 'Input list contains an invalid type.') + else: + return create_error_response(HTTPStatus.BAD_REQUEST, + 'Input must be a string or a list.') + + batch_scores = await async_engine.async_get_reward_score(input_ids) + prompt_tokens = sum(len(ids) for ids in input_ids) + usage = UsageInfo.build(prompt_tokens=prompt_tokens, + completion_tokens=0, + cached_tokens=0) + + data = [] + for i, score in enumerate(batch_scores): + data.append({ + 'index': i, + 'object': 'pooling', + 'data': score, + }) + + response = PoolingResponse(model=model_name, data=data, usage=usage) + return response.model_dump() + + @router.post('/get_ppl', dependencies=[Depends(validate_json_request)]) + async def get_ppl(request: PPLRequest, raw_request: Request = None): + """Get the perplexity (mean cross-entropy loss) of the input prompt. + + The request should be a JSON object with the following fields: + + - **input** (str | list[int]): the input to score, either raw text or token + ids. Text is tokenized with ``tokenizer.encode`` (no chat template is + applied). + """ + async_engine = server_context.async_engine + + request_input = request.input + + # pydantic already validated `input` as `str | list[int]`; text -> + # tokenizer.encode, otherwise the token ids are used as-is + if isinstance(request_input, str): + input_ids = async_engine.tokenizer.encode(request_input) + else: + input_ids = request_input + if not input_ids: + return create_error_response(HTTPStatus.BAD_REQUEST, + 'Input must not be empty.') + + try: + ppl = await async_engine.async_get_ppl(input_ids) + except ValueError as e: + return create_error_response(HTTPStatus.BAD_REQUEST, str(e)) + + response = PPLResponse(ppl=ppl) + return response.model_dump() diff --git a/lmdeploy/serve/openai/endpoints/chat_completions.py b/lmdeploy/serve/openai/endpoints/chat_completions.py new file mode 100644 index 0000000000..6467e78d84 --- /dev/null +++ b/lmdeploy/serve/openai/endpoints/chat_completions.py @@ -0,0 +1,639 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from __future__ import annotations + +import json +import time +from collections.abc import AsyncGenerator +from contextlib import aclosing +from functools import partial +from http import HTTPStatus +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from transformers import PreTrainedTokenizerBase + + +import shortuuid +from fastapi import APIRouter, Depends, HTTPException, Request +from fastapi.responses import StreamingResponse + +from lmdeploy.messages import LogitsProcessor +from lmdeploy.pytorch.disagg.conn.protocol import MigrationRequest +from lmdeploy.serve.openai.endpoints.common import build_serving_generation_config, validate_request +from lmdeploy.serve.openai.protocol import ( + ChatCompletionRequest, + ChatCompletionResponse, # noqa: E501 + ChatCompletionResponseChoice, + ChatCompletionResponseStreamChoice, + ChatCompletionStreamResponse, + ChatCompletionTokenLogprob, + ChatMessage, + ChoiceLogprobs, + DeltaMessage, + TopLogprob, + UsageInfo, +) +from lmdeploy.serve.openai.utils import create_error_response, maybe_filter_parallel_tool_calls +from lmdeploy.serve.utils.request_cleanup import with_request_cleanup +from lmdeploy.serve.utils.server_utils import validate_json_request +from lmdeploy.utils import get_logger + +logger = get_logger('lmdeploy') + + +def check_request(request: ChatCompletionRequest, server_context) -> str: + engine_config = server_context.engine_config + session_manager = server_context.session_manager + try: + # Check logprobs settings + logprobs_mode = engine_config.logprobs_mode + logprobs = request.logprobs + top_logprobs = request.top_logprobs or 0 + return_logprob = request.return_logprob + if logprobs_mode is None: + if logprobs or top_logprobs > 0: + return ( + f'Logprobs({logprobs})/top_logprobs({top_logprobs}) requested ' + 'but not enabled logprobs_mode in engine configuration') + if return_logprob: + return ( + f'return_logprob({return_logprob}) requested ' + 'but not enabled logprobs_mode in engine configuration.') + if logprobs_mode is not None and (top_logprobs < 0 or + (not logprobs and top_logprobs > 0)): + return ( + f'Invalid logprobs({logprobs})/top_logprobs({top_logprobs}) requested ' + 'when logprobs_mode is enabled in engine configuration.') + except AttributeError: + pass + + if session_manager.has(request.session_id): + return f'The session_id {request.session_id!r} is occupied.' + + # check sampling settings + if request.n <= 0: + return f'The n {request.n!r} must be a positive int.' + if request.top_p is not None and not (0 < request.top_p <= 1): + return f'The top_p {request.top_p!r} must be in (0, 1].' + if request.top_k is not None and request.top_k < 0: + return f'The top_k {request.top_k!r} cannot be a negative integer.' + if request.temperature is not None and not (0 <= request.temperature <= 2): + return f'The temperature {request.temperature!r} must be in [0, 2]' + + # Validate input_ids and image_data constraints. + # messages has higher priority. input_ids and image_data are only used when + # messages is empty (None, '', or []). image_data requires input_ids. + messages_empty = (request.messages is None or request.messages == '' + or (isinstance(request.messages, list) + and len(request.messages) == 0)) + if not messages_empty: + # messages is active — input_ids and image_data must not be set + if request.input_ids is not None: + return 'input_ids cannot be used when messages is non-empty. messages takes priority.' + if request.image_data is not None: + return 'image_data cannot be used when messages is non-empty. messages takes priority.' + else: + # messages is empty — input_ids and image_data are the active inputs + if request.input_ids is not None and len(request.input_ids) == 0: + return 'The input_ids must not be an empty list.' + if request.image_data is not None and request.input_ids is None: + return 'image_data requires input_ids to be set when messages is empty.' + + if request.return_routed_experts and not engine_config.enable_return_routed_experts: + return ( + 'routed experts requested but not configured in engine configuration. ' + 'May start api_server with --enable-return-routed-experts flag.') + + return '' + + +def _create_chat_completion_logprobs(tokenizer: PreTrainedTokenizerBase, + token_ids: list[int] | None = None, + logprobs: list[dict[int, float]] + | None = None): + """Create openai LogProbs for chat.completion. + + Args: + tokenizer (PreTrainedTokenizerBase): tokenizer. + token_ids (list[int]): output token ids. + logprobs (list[dict[int, float]]): the top logprobs for each output + position. + Returns: + ChoiceLogprobs: logprob result. + """ + if token_ids is None or logprobs is None: + return None + + content: list[ChatCompletionTokenLogprob] = [] + for token_id, tops in zip(token_ids, logprobs): + item = ChatCompletionTokenLogprob(token='', + bytes=[], + logprob=0.0, + top_logprobs=[]) + for top_id, prob in tops.items(): + token = tokenizer.convert_ids_to_tokens(top_id) + if isinstance(token, bytes): + _bytes = list(token) + token = token.decode('utf-8', errors='backslashreplace') + else: + _bytes = list(token.encode()) # token is str + if top_id == token_id: + item.token = token + item.bytes = _bytes + item.logprob = prob + else: + item.top_logprobs.append( + TopLogprob(token=token, bytes=_bytes, logprob=prob)) + content.append(item) + return ChoiceLogprobs(content=content) + + +def _create_output_token_logprobs(token_ids: list[int] | None = None, + logprobs: list[dict[int, float]] + | None = None): + """Create raw (logprob, token_id) pairs for output tokens.""" + if token_ids is None or logprobs is None: + return None + + output_token_logprobs = [] + for tok, tok_logprobs in zip(token_ids, logprobs): + output_token_logprobs.append((tok_logprobs[tok], tok)) + return output_token_logprobs or None + + +# modified from https://github.com/vllm-project/vllm/blob/v0.5.4/vllm/entrypoints/openai/logits_processors.py#L51 # noqa +def logit_bias_logits_processor( + logit_bias: dict[int, float] | dict[str, float], + tokenizer: PreTrainedTokenizerBase) -> LogitsProcessor: + try: + # Convert token_id to integer + # Clamp the bias between -100 and 100 per OpenAI API spec + clamped_logit_bias: dict[int, float] = { + int(token_id): min(100.0, max(-100.0, bias)) + for token_id, bias in logit_bias.items() + } + except ValueError as exc: + raise ValueError( + 'Found token_id in logit_bias that is not ' + 'an integer or string representing an integer') from exc + + # Check if token_id is within the vocab size + for token_id, bias in clamped_logit_bias.items(): + if token_id < 0 or token_id >= tokenizer.vocab_size: + raise ValueError(f'token_id {token_id} in logit_bias contains ' + 'out-of-vocab token id') + + def _logit_bias_processor( + logit_bias, + token_ids, + logits, + ): + for token_id, bias in logit_bias.items(): + logits[token_id] = logits[token_id] + bias + return logits + + return partial(_logit_bias_processor, clamped_logit_bias) + + +def register(router: APIRouter, server_context) -> None: + + @router.post('/v1/chat/completions', + dependencies=[Depends(validate_json_request)]) + async def chat_completions_v1(request: ChatCompletionRequest, + raw_request: Request = None): + """Completion API similar to OpenAI's API. + + Refer to https://platform.openai.com/docs/api-reference/chat/create + for the API specification. + + The request should be a JSON object with the following fields: + + - **model**: model name. Available from /v1/models. + - **messages**: string prompt or chat history in OpenAI format. Chat history example: + ``[{"role": "user", "content": "hi"}]``. + - **temperature** (float): to modulate the next token probability + - **top_p** (float): If set to float < 1, only the smallest set of most + probable tokens with probabilities that add up to top_p or higher + are kept for generation. + - **n** (int): How many chat completion choices to generate for each input + message. **Only support one here**. + - **stream**: whether to stream the results or not. Default to false. + - **stream_options**: Options for streaming response. Only set this when you + set stream: true. + - **max_completion_tokens** (int | None): output token nums. Default to None. + - **max_tokens** (int | None): output token nums. Default to None. + Deprecated: Use max_completion_tokens instead. + - **repetition_penalty** (float): The parameter for repetition penalty. + 1.0 means no penalty + - **stop** (str | list[str] | None): To stop generating further + tokens. Only accept stop words that's encoded to one token idex. + - **response_format** (dict | None): To generate response according to given + schema. Examples: + + .. code-block:: json + + { + "type": "json_schema", + "json_schema":{ + "name": "test", + "schema":{ + "properties":{ + "name":{"type":"string"} + }, + "required":["name"], + "type":"object" + } + } + } + + or ``{"type": "regex_schema", "regex_schema": "call me [A-Za-z]{1,10}"}`` + - **logit_bias** (dict): Bias to logits. Only supported in pytorch engine. + - **tools** (list): A list of tools the model may call. Currently, only + internlm2 functions are supported as a tool. Use this to specify a + list of functions for which the model can generate JSON inputs. + - **tool_choice** (str | object): Controls which (if any) tool is called by + the model. `none` means the model will not call any tool and instead + generates a message. Specifying a particular tool via + ``{"type": "function", "function": {"name": "my_function"}}`` + forces the model to call that tool. `auto` or `required` will put all + the tools informationto the model. + + Additional arguments supported by LMDeploy: + + - **top_k** (int): The number of the highest probability vocabulary + tokens to keep for top-k-filtering + - **ignore_eos** (bool): indicator for ignoring eos + - **skip_special_tokens** (bool): Whether or not to remove special tokens + in the decoding. Default to be True. + - **spaces_between_special_tokens** (bool): Whether or not to add spaces + around special tokens. The behavior of Fast tokenizers is to have + this to False. This is setup to True in slow tokenizers. + - **min_new_tokens** (int): To generate at least numbers of tokens. + - **min_p** (float): Minimum token probability, which will be scaled by the + probability of the most likely token. It must be a value between + 0 and 1. Typical values are in the 0.01-0.2 range, comparably + selective as setting `top_p` in the 0.99-0.8 range (use the + opposite of normal `top_p` values) + - **repetition_ngram_size** (int): N-gram length for repetition early stop + (PyTorch engine). ``0`` disables. + - **repetition_ngram_threshold** (int): How many times that n-gram must + repeat to trigger early stop. ``0`` disables. + + Currently we do not support the following features: + + - **presence_penalty** (replaced with repetition_penalty) + - **frequency_penalty** (replaced with repetition_penalty) + """ + error_check_ret = validate_request(request, server_context, + check_request) + if error_check_ret is not None: + return error_check_ret + session = server_context.create_session(request.session_id) + + # Resolve input: messages has priority over input_ids/image_data + messages_empty = (request.messages is None or request.messages == '' + or (isinstance(request.messages, list) + and len(request.messages) == 0)) + resolved_input_ids = None + if messages_empty and request.input_ids is not None: + # /generate-style input: use input_ids (+ optional image_data) + resolved_input_ids = request.input_ids + if request.image_data is not None: + # Convert image_data to OpenAI multimodal content format + image_data = request.image_data + image_input = [] + if not isinstance(image_data, list): + image_data = [image_data] + for img in image_data: + if isinstance(img, str): + image_input.append( + dict(type='image_url', image_url=dict(url=img))) + else: + image_input.append( + dict(type='image_url', image_url=img)) + text_input = dict(type='text', text=request.input_ids) + request.messages = [ + dict(role='user', content=[text_input] + image_input) + ] + resolved_input_ids = None # image_data conversion takes over + else: + # input_ids only — engine requires messages=None + request.messages = None + + json_request = await raw_request.json() + migration_request = json_request.pop('migration_request', None) + with_cache = json_request.pop('with_cache', False) + preserve_cache = json_request.pop('preserve_cache', False) + if migration_request: + migration_request = MigrationRequest.model_validate( + migration_request) + + model_name = request.model + adapter_name = None + if model_name != server_context.async_engine.model_name: + adapter_name = model_name # got a adapter name + request_id = f'chatcmpl-{shortuuid.random()}' + created_time = int(time.time()) + + tokenizer = server_context.async_engine.tokenizer.model.model + gen_logprobs, logits_processors = None, None + if request.logprobs: + gen_logprobs = request.top_logprobs or 1 + elif request.return_logprob: + gen_logprobs = 1 + if request.logit_bias is not None: + try: + logits_processors = [ + logit_bias_logits_processor(request.logit_bias, tokenizer) + ] + except Exception as e: + return create_error_response(HTTPStatus.BAD_REQUEST, str(e)) + + parser_cls = server_context.response_parser_cls + if request.tool_choice != 'none' and request.tools: + if parser_cls is None or parser_cls.tool_parser_cls is None: + return create_error_response( + HTTPStatus.BAD_REQUEST, + 'Please launch the api_server with --tool-call-parser if you want to use tool.' + ) + + parser_cls = server_context.response_parser_cls + try: + response_parser = parser_cls(request) + except ValueError as e: + raise HTTPException(status_code=HTTPStatus.BAD_REQUEST, + detail=str(e)) + # request is normalized and may be adjusted by the parser + # (e.g. GPT-OSS clears response_format and injects the schema into messages) + request = response_parser.request + + gen_config = build_serving_generation_config( + request, + server_context, + logprobs=gen_logprobs, + stop_words=request.stop, + logits_processors=logits_processors, + random_seed=request.seed, + migration_request=migration_request, + with_cache=with_cache, + preserve_cache=preserve_cache, + ) + + # text completion for string input or input_ids + do_preprocess = (False if isinstance(request.messages, str) + or resolved_input_ids is not None else + request.do_preprocess) + chat_template_kwargs = request.chat_template_kwargs or {} + if request.enable_thinking is not None: + logger.warning( + '`enable_thinking` will be deprecated in the future, ' + 'please use `chat_template_kwargs` instead.') + if chat_template_kwargs.get('enable_thinking') is None: + chat_template_kwargs[ + 'enable_thinking'] = request.enable_thinking + else: + logger.warning( + '`enable_thinking` in `chat_template_kwargs` will override the value in request.' + ) + + result_generator = server_context.async_engine.generate( + request.messages, + session, + gen_config=gen_config, + tools=request.tools, + reasoning_effort=request.reasoning_effort, + stream_response=True, # always use stream to enable batching + do_preprocess=do_preprocess, + adapter_name=adapter_name, + chat_template_kwargs=chat_template_kwargs or None, + input_ids=resolved_input_ids, + media_io_kwargs=request.media_io_kwargs, + mm_processor_kwargs=request.mm_processor_kwargs) + include_usage = bool(request.stream_options + and request.stream_options.include_usage) + + def create_stream_response_json( + index: int, + delta_message: DeltaMessage, + finish_reason: str | None = None, + logprobs: ChoiceLogprobs | None = None, + output_token_logprobs: list[tuple[float, int]] | None = None, + routed_experts=None, + output_ids=None) -> dict: + choice_data = ChatCompletionResponseStreamChoice( + index=index, + delta=delta_message, + finish_reason=finish_reason, + logprobs=logprobs, + output_token_logprobs=output_token_logprobs, + output_ids=output_ids, + routed_experts=routed_experts) + choice_data = maybe_filter_parallel_tool_calls( + choice_data, request) + response = ChatCompletionStreamResponse( + id=request_id, + created=created_time, + model=model_name, + choices=[choice_data], + usage=None, + ) + response_dict = response.model_dump(mode='json', exclude_none=True) + if include_usage: + response_dict['usage'] = None + return response_dict + + def create_stream_usage_response_json(usage: UsageInfo) -> str: + response = ChatCompletionStreamResponse( + id=request_id, + created=created_time, + model=model_name, + choices=[], + usage=usage, + ) + return response.model_dump_json(exclude_none=True) + + async def completion_stream_generator() -> AsyncGenerator[str, None]: + streaming_tools = False + final_usage = None + async for res in result_generator: + logprobs = None + output_token_logprobs = None + if request.logprobs and res.logprobs: + logprobs = _create_chat_completion_logprobs( + tokenizer, res.token_ids, res.logprobs) + if request.return_logprob: + output_token_logprobs = _create_output_token_logprobs( + res.token_ids, res.logprobs) + if res.finish_reason and include_usage: + final_usage = UsageInfo.build( + prompt_tokens=res.input_token_len, + completion_tokens=res.generate_token_len, + cached_tokens=res.cached_tokens, + ) + delta_token_ids = res.token_ids if res.token_ids is not None else [] + stream_deltas = response_parser.stream_chunk( + res.response, delta_token_ids) + if not stream_deltas: + # Parser may buffer partial protocol tags and emit no visible delta + # while the engine still produced new tokens (e.g. MTP batch). Do not + # drop those token ids; emit them once on a placeholder delta. + if res.finish_reason is None and not delta_token_ids: + continue + stream_deltas = [(DeltaMessage(role='assistant', + content=''), False)] + should_validate_complete = (res.finish_reason + in ('stop', 'length') and + (request.return_token_ids + or request.return_routed_experts)) + if should_validate_complete and not response_parser.validate_complete( + ): + res.finish_reason = 'parse_error' + + for delta_index, (delta_message, + tool_emitted) in enumerate(stream_deltas): + if tool_emitted: + streaming_tools = True + + is_last_delta = delta_index == len(stream_deltas) - 1 + # The chat parser may split one engine yield into multiple protocol deltas, + # so attach the engine-level metadata to the last parsed delta. + finish_reason = res.finish_reason if is_last_delta else None + chunk_logprobs = logprobs if is_last_delta else None + chunk_output_token_logprobs = output_token_logprobs if is_last_delta else None + + if (request.tool_choice != 'none' + and response_parser.tool_parser is not None): + if finish_reason == 'stop' and streaming_tools is True: + finish_reason = 'tool_calls' + + # Only output routed_experts in the final chunk + routed_experts = res.routed_experts if finish_reason is not None else None + # Emit token ids once per engine yield on the last parsed delta, when + # accumulated delta text and token ids for this step are aligned. + stream_output_ids = delta_token_ids if ( + request.return_token_ids and is_last_delta) else None + + response_json = create_stream_response_json( + index=0, + delta_message=delta_message, + finish_reason=finish_reason, + logprobs=chunk_logprobs, + output_token_logprobs=chunk_output_token_logprobs, + routed_experts=routed_experts, + output_ids=stream_output_ids) + if res.cache_block_ids is not None and is_last_delta: + response_json['cache_block_ids'] = res.cache_block_ids + response_json['remote_token_ids'] = res.token_ids + yield f'data: {json.dumps(response_json)}\n\n' + if final_usage is not None: + yield f'data: {create_stream_usage_response_json(final_usage)}\n\n' + yield 'data: [DONE]\n\n' + + # Streaming response + if request.stream: + stream_generator = with_request_cleanup( + completion_stream_generator(), [result_generator], [session], + server_context.session_manager) + return StreamingResponse(stream_generator, + media_type='text/event-stream') + + # Non-streaming response + final_logprobs = [] + final_token_ids = [] + final_res = None + text = '' + cache_block_ids = [] + remote_token_ids = [] + async with aclosing( + with_request_cleanup( + result_generator, [result_generator], [session], + server_context.session_manager)) as generator: + async for res in generator: + if await raw_request.is_disconnected(): + # Abort the request if the client disconnects. + await session.async_abort() + return create_error_response(HTTPStatus.BAD_REQUEST, + 'Client disconnected') + final_res = res + text += res.response + if res.token_ids: + final_token_ids.extend(res.token_ids) + if res.logprobs: + final_logprobs.extend(res.logprobs) + cache_block_ids.append(res.cache_block_ids) + remote_token_ids.append(res.token_ids) + + tool_calls = None + reasoning_content = None + + try: + raw_text = text + text, tool_calls, reasoning_content = response_parser.parse_complete( + text, final_token_ids) + should_validate_complete = ( + final_res.finish_reason in ('stop', 'length') and + (request.return_token_ids or request.return_routed_experts)) + if should_validate_complete and not response_parser.validate_complete( + raw_text): + final_res.finish_reason = 'parse_error' + if isinstance(tool_calls, list) and len(tool_calls): + if final_res.finish_reason == 'stop': + final_res.finish_reason = 'tool_calls' + + except Exception as e: + logger.error(f'Failed to parse {text}. Exception: {e}.') + return create_error_response( + HTTPStatus.BAD_REQUEST, + 'Failed to parse fc related info to json format!') + + message = ChatMessage(role='assistant', + content=text, + tool_calls=tool_calls, + reasoning_content=reasoning_content) + + logprobs = None + if request.logprobs and len(final_logprobs): + logprobs = _create_chat_completion_logprobs( + tokenizer, final_token_ids, final_logprobs) + output_token_logprobs = None + if request.return_logprob and len(final_logprobs): + output_token_logprobs = _create_output_token_logprobs( + final_token_ids, final_logprobs) + + assert final_res is not None + choices = [] + choice_data = ChatCompletionResponseChoice( + index=0, + message=message, + logprobs=logprobs, + output_token_logprobs=output_token_logprobs, + finish_reason=final_res.finish_reason, + output_ids=final_token_ids if request.return_token_ids else None, + routed_experts=final_res.routed_experts + if request.return_routed_experts else None, + ) + choice_data = maybe_filter_parallel_tool_calls(choice_data, request) + choices.append(choice_data) + + if with_cache: + cache_block_ids = cache_block_ids[0] + remote_token_ids = [remote_token_ids[0][-1]] + + usage = UsageInfo.build( + prompt_tokens=final_res.input_token_len, + completion_tokens=final_res.generate_token_len, + cached_tokens=final_res.cached_tokens, + ) + response = ChatCompletionResponse( + id=request_id, + created=created_time, + model=model_name, + choices=choices, + usage=usage, + ).model_dump() + + if with_cache: + response['cache_block_ids'] = cache_block_ids + response['remote_token_ids'] = remote_token_ids + + return response diff --git a/lmdeploy/serve/openai/endpoints/common.py b/lmdeploy/serve/openai/endpoints/common.py new file mode 100644 index 0000000000..79e3e7eace --- /dev/null +++ b/lmdeploy/serve/openai/endpoints/common.py @@ -0,0 +1,37 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""Shared helpers for OpenAI-compatible endpoint handlers.""" + +from http import HTTPStatus + +from lmdeploy.messages import GenerationConfig +from lmdeploy.serve.core.generation_config import build_generation_config +from lmdeploy.serve.openai.utils import create_error_response, get_model_list + + +def build_serving_generation_config(request, server_context, + **extra_kwargs) -> GenerationConfig: + """Merge request sampling values with server generation defaults.""" + max_new_tokens = getattr(request, 'max_completion_tokens', None) + if max_new_tokens is None: + max_new_tokens = getattr(request, 'max_tokens', None) + return build_generation_config( + request, + server_context.default_gen_config, + max_new_tokens=max_new_tokens, + **extra_kwargs, + ) + + +def validate_request(request, server_context, request_validator): + """Validate the selected model and endpoint-specific request contract.""" + if hasattr( + request, + 'model') and request.model not in get_model_list(server_context): + return create_error_response( + HTTPStatus.NOT_FOUND, + f'The model {request.model!r} does not exist.') + + error_message = request_validator(request, server_context) + if error_message: + return create_error_response(HTTPStatus.BAD_REQUEST, error_message) + return None diff --git a/lmdeploy/serve/openai/endpoints/completions.py b/lmdeploy/serve/openai/endpoints/completions.py new file mode 100644 index 0000000000..38ebc01d10 --- /dev/null +++ b/lmdeploy/serve/openai/endpoints/completions.py @@ -0,0 +1,312 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from __future__ import annotations + +import asyncio +import json +import time +from collections.abc import AsyncGenerator +from contextlib import aclosing +from http import HTTPStatus + +from fastapi import APIRouter, Depends, Request +from fastapi.responses import StreamingResponse + +from lmdeploy.pytorch.disagg.conn.protocol import MigrationRequest +from lmdeploy.serve.openai.endpoints.common import build_serving_generation_config, validate_request +from lmdeploy.serve.openai.protocol import ( + CompletionRequest, + CompletionResponse, + CompletionResponseChoice, + CompletionResponseStreamChoice, + CompletionStreamResponse, + LogProbs, + UsageInfo, +) +from lmdeploy.serve.openai.utils import create_error_response +from lmdeploy.serve.utils.request_cleanup import with_request_cleanup +from lmdeploy.serve.utils.server_utils import validate_json_request + + +def check_request(request: CompletionRequest, server_context) -> str: + engine_config = server_context.engine_config + session_manager = server_context.session_manager + try: + # Check logprobs settings + logprobs_mode = engine_config.logprobs_mode + logprobs = request.logprobs or 0 + if logprobs > 0 and logprobs_mode is None: + return f'logprobs({logprobs}) requested but not enabled logprobs_mode in engine configuration.' + if logprobs_mode is not None and logprobs < 0: + return 'logprobs must be non-negative when logprobs_mode is enabled in engine configuration.' + except AttributeError: + pass + + if session_manager.has(request.session_id): + return f'The session_id {request.session_id!r} is occupied.' + + # check sampling settings + if request.n <= 0: + return f'The n {request.n!r} must be a positive int.' + if request.top_p is not None and not (0 < request.top_p <= 1): + return f'The top_p {request.top_p!r} must be in (0, 1].' + if request.top_k is not None and request.top_k < 0: + return f'The top_k {request.top_k!r} cannot be a negative integer.' + if request.temperature is not None and not (0 <= request.temperature <= 2): + return f'The temperature {request.temperature!r} must be in [0, 2]' + + return '' + + +def register(router: APIRouter, server_context) -> None: + + @router.post('/v1/completions', + dependencies=[Depends(validate_json_request)]) + async def completions_v1(request: CompletionRequest, + raw_request: Request = None): + """Completion API similar to OpenAI's API. + + Go to https://platform.openai.com/docs/api-reference/completions/create + for the API specification. + + The request should be a JSON object with the following fields: + + - **model** (str): model name. Available from /v1/models. + - **prompt** (str): the input prompt. + - **suffix** (str): The suffix that comes after a completion of inserted text. + - **max_completion_tokens** (int | None): output token nums. Default to None. + - **max_tokens** (int | None): output token nums. Default to 16. + Deprecated: Use max_completion_tokens instead. + - **temperature** (float): to modulate the next token probability + - **top_p** (float): If set to float < 1, only the smallest set of most + probable tokens with probabilities that add up to top_p or higher + are kept for generation. + - **n** (int): How many chat completion choices to generate for each input + message. **Only support one here**. + - **stream**: whether to stream the results or not. Default to false. + - **stream_options**: Options for streaming response. Only set this when you + set stream: true. + - **repetition_penalty** (float): The parameter for repetition penalty. + 1.0 means no penalty + - **user** (str): A unique identifier representing your end-user. + - **stop** (str | list[str] | None): To stop generating further + tokens. Only accept stop words that's encoded to one token idex. + + Additional arguments supported by LMDeploy: + + - **ignore_eos** (bool): indicator for ignoring eos + - **skip_special_tokens** (bool): Whether or not to remove special tokens + in the decoding. Default to be True. + - **spaces_between_special_tokens** (bool): Whether or not to add spaces + around special tokens. The behavior of Fast tokenizers is to have + this to False. This is setup to True in slow tokenizers. + - **top_k** (int): The number of the highest probability vocabulary + tokens to keep for top-k-filtering + - **min_p** (float): Minimum token probability, which will be scaled by the + probability of the most likely token. It must be a value between + 0 and 1. Typical values are in the 0.01-0.2 range, comparably + selective as setting `top_p` in the 0.99-0.8 range (use the + opposite of normal `top_p` values) + - **repetition_ngram_size** (int): N-gram length for repetition early stop + (PyTorch engine). ``0`` disables. + - **repetition_ngram_threshold** (int): How many times that n-gram must + repeat to trigger early stop. ``0`` disables. + + Currently we do not support the following features: + + - **logprobs** (not supported yet) + - **presence_penalty** (replaced with repetition_penalty) + - **frequency_penalty** (replaced with repetition_penalty) + """ + error_check_ret = validate_request(request, server_context, + check_request) + if error_check_ret is not None: + return error_check_ret + json_request = await raw_request.json() + migration_request = json_request.pop('migration_request', None) + with_cache = json_request.pop('with_cache', False) + preserve_cache = json_request.pop('preserve_cache', False) + if migration_request: + migration_request = MigrationRequest.model_validate( + migration_request) + + model_name = request.model + adapter_name = None + if model_name != server_context.async_engine.model_name: + adapter_name = model_name # got a adapter name + request_id = str(request.session_id) + created_time = int(time.time()) + sessions = [] + if isinstance(request.prompt, str): + request.prompt = [request.prompt] + sessions.append(server_context.create_session(request.session_id)) + elif isinstance(request.prompt, list): + for _ in request.prompt: + sessions.append(server_context.create_session()) + if isinstance(request.stop, str): + request.stop = [request.stop] + + gen_config = build_serving_generation_config( + request, + server_context, + stop_words=request.stop, + random_seed=request.seed, + migration_request=migration_request, + with_cache=with_cache, + preserve_cache=preserve_cache, + ) + generators = [] + for prompt, session in zip(request.prompt, sessions): + result_generator = server_context.async_engine.generate( + prompt, + session, + gen_config=gen_config, + stream_response=True, # always use stream to enable batching + do_preprocess=False, + adapter_name=adapter_name) + generators.append(result_generator) + + include_usage = bool(request.stream_options + and request.stream_options.include_usage) + + def create_stream_response_json( + index: int, + text: str, + finish_reason: str | None = None, + logprobs: LogProbs | None = None) -> dict: + choice_data = CompletionResponseStreamChoice( + index=index, + text=text, + finish_reason=finish_reason, + logprobs=logprobs) + response = CompletionStreamResponse( + id=request_id, + created=created_time, + model=model_name, + choices=[choice_data], + usage=None, + ) + response_json = response.model_dump(mode='json', exclude_none=True) + if include_usage: + response_json['usage'] = None + return response_json + + def create_stream_usage_response_json(usage: UsageInfo) -> dict: + response = CompletionStreamResponse( + id=request_id, + created=created_time, + model=model_name, + choices=[], + usage=usage, + ) + return response.model_dump(mode='json', exclude_none=True) + + async def completion_stream_generator() -> AsyncGenerator[str, None]: + # First chunk with role + final_usage = UsageInfo() if include_usage else None + for generator in generators: + async for res in generator: + logprobs = None + if request.logprobs and res.logprobs: + raise ValueError('logprobs is removed') + if res.finish_reason and final_usage is not None: + final_res = res + total_tokens = sum([ + final_res.input_token_len, + final_res.generate_token_len + ]) + final_usage.prompt_tokens += final_res.input_token_len + final_usage.completion_tokens += final_res.generate_token_len + final_usage.total_tokens += total_tokens + response_json = create_stream_response_json( + index=0, + text=res.response, + finish_reason=res.finish_reason, + logprobs=logprobs) + if res.cache_block_ids is not None: + response_json['cache_block_ids'] = res.cache_block_ids + response_json['remote_token_ids'] = res.token_ids + yield f'data: {json.dumps(response_json)}\n\n' + if final_usage is not None: + yield f'data: {json.dumps(create_stream_usage_response_json(final_usage))}\n\n' + yield 'data: [DONE]\n\n' + + # Streaming response + if request.stream: + stream_generator = with_request_cleanup( + completion_stream_generator(), generators, sessions, + server_context.session_manager) + return StreamingResponse(stream_generator, + media_type='text/event-stream') + + # Non-streaming response + usage = UsageInfo() + choices = [None] * len(generators) + cache_block_ids = [] + remote_token_ids = [] + + async def _inner_call(i, generator, session): + nonlocal cache_block_ids, remote_token_ids + final_logprobs = [] + final_token_ids = [] + final_res = None + text = '' + async with aclosing( + with_request_cleanup(generator, [generator], [session], + server_context.session_manager) + ) as cleanup_generator: + async for res in cleanup_generator: + if await raw_request.is_disconnected(): + # Abort the request if the client disconnects. + await session.async_abort() + return create_error_response(HTTPStatus.BAD_REQUEST, + 'Client disconnected') + final_res = res + text += res.response + cache_block_ids.append(res.cache_block_ids) + remote_token_ids.append(res.token_ids) + if res.token_ids: + final_token_ids.extend(res.token_ids) + if res.logprobs: + final_logprobs.extend(res.logprobs) + + logprobs = None + assert final_res is not None + choice_data = CompletionResponseChoice( + index=i, + text=text, + finish_reason=final_res.finish_reason, + logprobs=logprobs) + choices[i] = choice_data + + if with_cache: + cache_block_ids = cache_block_ids[0] + remote_token_ids = [remote_token_ids[0][-1]] + + total_tokens = sum( + [final_res.input_token_len, final_res.generate_token_len]) + usage.prompt_tokens += final_res.input_token_len + usage.completion_tokens += final_res.generate_token_len + usage.total_tokens += total_tokens + + call_results = await asyncio.gather(*[ + _inner_call(i, generators[i], sessions[i]) + for i in range(len(generators)) + ]) + error_response = next( + (result for result in call_results if result is not None), None) + if error_response is not None: + return error_response + + response = CompletionResponse( + id=request_id, + created=created_time, + model=model_name, + choices=choices, + usage=usage, + ).model_dump() + + if with_cache: + response['cache_block_ids'] = cache_block_ids + response['remote_token_ids'] = remote_token_ids + + return response diff --git a/lmdeploy/serve/openai/endpoints/distserve.py b/lmdeploy/serve/openai/endpoints/distserve.py new file mode 100644 index 0000000000..da88fb83d5 --- /dev/null +++ b/lmdeploy/serve/openai/endpoints/distserve.py @@ -0,0 +1,53 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from __future__ import annotations + +from fastapi import APIRouter +from fastapi.responses import JSONResponse + +from lmdeploy.pytorch.disagg.config import DistServeEngineConfig +from lmdeploy.pytorch.disagg.conn.protocol import ( + DistServeCacheFreeRequest, + DistServeConnectionRequest, + DistServeDropConnectionRequest, + DistServeInitRequest, +) + + +def register(router: APIRouter, server_context) -> None: + """PD Disaggregation API Begin.""" + + @router.get('/distserve/engine_info') + async def engine_info(): + engine_config = server_context.engine_config + + response = DistServeEngineConfig( + tp_size=engine_config.tp, + dp_size=engine_config.dp, + pp_size=None, + ep_size=engine_config.ep, + dp_rank=engine_config.dp_rank, + block_size=engine_config.block_size, + num_cpu_blocks=engine_config.num_cpu_blocks, + num_gpu_blocks=engine_config.num_gpu_blocks) + + return response.model_dump_json() + + @router.post('/distserve/p2p_initialize') + async def p2p_initialize(init_request: DistServeInitRequest): + return server_context.async_engine.p2p_initialize(init_request) + + @router.post('/distserve/p2p_connect') + async def p2p_connect(conn_request: DistServeConnectionRequest): + return server_context.async_engine.p2p_connect(conn_request) + + @router.post('/distserve/p2p_drop_connect') + async def p2p_drop_connect( + drop_conn_request: DistServeDropConnectionRequest): + return server_context.async_engine.p2p_drop_connect(drop_conn_request) + + @router.post('/distserve/free_cache') + async def free_cache( + cache_free_request: DistServeCacheFreeRequest) -> JSONResponse: + session_id = cache_free_request.remote_session_id + server_context.async_engine.free_cache(session_id) + return {'status': 'SUCCESS'} diff --git a/lmdeploy/serve/openai/endpoints/generate.py b/lmdeploy/serve/openai/endpoints/generate.py new file mode 100644 index 0000000000..23d5bb342e --- /dev/null +++ b/lmdeploy/serve/openai/endpoints/generate.py @@ -0,0 +1,188 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from __future__ import annotations + +from contextlib import aclosing +from http import HTTPStatus + +from fastapi import APIRouter, Depends, Request +from fastapi.responses import StreamingResponse + +from lmdeploy.serve.openai.endpoints.common import build_serving_generation_config, validate_request +from lmdeploy.serve.openai.protocol import ( + GenerateReqInput, + GenerateReqMetaOutput, + GenerateReqOutput, +) +from lmdeploy.serve.openai.utils import create_error_response +from lmdeploy.serve.utils.request_cleanup import with_request_cleanup +from lmdeploy.serve.utils.server_utils import validate_json_request + + +def check_request(request: GenerateReqInput, server_context) -> str: + engine_config = server_context.engine_config + session_manager = server_context.session_manager + try: + # Check logprobs settings + logprobs_mode = engine_config.logprobs_mode + return_logprob = request.return_logprob + if logprobs_mode is None and return_logprob: + return f'return_logprob({return_logprob}) requested but not enabled logprobs_mode in engine configuration.' + except AttributeError: + pass + + if (request.prompt is not None) ^ (request.input_ids is None): + return 'You must specify exactly one of prompt or input_ids' + + if request.prompt is not None and request.prompt == '': + return 'The prompt must not be an empty string' + + if request.input_ids is not None and len(request.input_ids) == 0: + return 'The input_ids must not be an empty list' + + if request.max_tokens is not None and request.max_tokens <= 0: + return f'The max_tokens {request.max_tokens!r} must be a positive integer.' + + if session_manager.has(request.session_id): + return f'The session_id {request.session_id!r} is occupied.' + + # check sampling settings + if request.top_p is not None and not (0 < request.top_p <= 1): + return f'The top_p {request.top_p!r} must be in (0, 1].' + if request.top_k is not None and request.top_k < 0: + return f'The top_k {request.top_k!r} cannot be a negative integer.' + if request.temperature is not None and not (0 <= request.temperature <= 2): + return f'The temperature {request.temperature!r} must be in [0, 2]' + + return '' + + +def register(router: APIRouter, server_context) -> None: + + @router.post('/generate', dependencies=[Depends(validate_json_request)]) + async def generate(request: GenerateReqInput, raw_request: Request = None): + error_check_ret = validate_request(request, server_context, + check_request) + if error_check_ret is not None: + return error_check_ret + + session = server_context.create_session(request.session_id) + + prompt = request.prompt + input_ids = request.input_ids + image_data = request.image_data + if image_data is not None: + # convert to openai format + image_input = [] + if not isinstance(image_data, list): + image_data = [image_data] + for img in image_data: + if isinstance(img, str): + image_input.append( + dict(type='image_url', image_url=dict(url=img))) + else: + image_input.append(dict(type='image_url', image_url=img)) + text_input = dict(type='text', + text=prompt if prompt else input_ids) + prompt = [dict(role='user', content=[text_input] + image_input)] + input_ids = None + + gen_config = build_serving_generation_config( + request, + server_context, + logprobs=1 if request.return_logprob else None, + stop_words=request.stop, + ) + + result_generator = server_context.async_engine.generate( + messages=prompt, + session_id=session, + input_ids=input_ids, + gen_config=gen_config, + stream_response=True, # always use stream to enable batching + do_preprocess=False, + media_io_kwargs=request.media_io_kwargs, + mm_processor_kwargs=request.mm_processor_kwargs) + + def create_generate_response_json(res, + text, + output_ids, + logprobs, + finish_reason, + routed_experts=None): + # only output router experts in last chunk + routed_experts = None if finish_reason is None else routed_experts + meta = GenerateReqMetaOutput( + finish_reason=dict( + type=finish_reason) if finish_reason else None, + output_token_logprobs=logprobs or None, + prompt_tokens=res.input_token_len, + routed_experts=routed_experts, + completion_tokens=res.generate_token_len) + + response = GenerateReqOutput(text=text, + output_ids=output_ids, + meta_info=meta, + routed_experts=routed_experts) + return response.model_dump_json() + + async def generate_stream_generator(): + async for res in result_generator: + text = res.response or '' + output_ids = res.token_ids + routed_experts = res.routed_experts + logprobs = [] + if res.logprobs: + for tok, tok_logprobs in zip(res.token_ids, res.logprobs): + logprobs.append((tok_logprobs[tok], tok)) + response_json = create_generate_response_json( + res, + text, + output_ids, + logprobs, + res.finish_reason, + routed_experts=routed_experts) + yield f'data: {response_json}\n\n' + yield 'data: [DONE]\n\n' + + if request.stream: + stream_generator = with_request_cleanup( + generate_stream_generator(), [result_generator], [session], + server_context.session_manager) + return StreamingResponse(stream_generator, + media_type='text/event-stream') + + async def _inner_call(): + text = '' + output_ids = [] + logprobs = [] + async with aclosing( + with_request_cleanup( + result_generator, [result_generator], [session], + server_context.session_manager)) as generator: + async for res in generator: + if await raw_request.is_disconnected(): + # Abort the request if the client disconnects. + await session.async_abort() + return create_error_response(HTTPStatus.BAD_REQUEST, + 'Client disconnected') + text += res.response or '' + output_ids.extend(res.token_ids or []) + logprobs.extend(res.logprobs or []) + + output_token_logprobs = [] + if len(logprobs) and len(output_ids): + for tok, tok_logprobs in zip(output_ids, logprobs): + output_token_logprobs.append((tok_logprobs[tok], tok)) + + meta = GenerateReqMetaOutput( + finish_reason=dict( + type=res.finish_reason) if res.finish_reason else None, + output_token_logprobs=output_token_logprobs or None, + prompt_tokens=res.input_token_len, + routed_experts=res.routed_experts, + completion_tokens=res.generate_token_len) + return GenerateReqOutput(text=text, + output_ids=output_ids, + meta_info=meta) + + return await _inner_call() diff --git a/lmdeploy/serve/openai/endpoints/management.py b/lmdeploy/serve/openai/endpoints/management.py new file mode 100644 index 0000000000..6726ee17ae --- /dev/null +++ b/lmdeploy/serve/openai/endpoints/management.py @@ -0,0 +1,175 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from __future__ import annotations + +import os +from http import HTTPStatus + +from fastapi import APIRouter, Depends, Request +from fastapi.encoders import jsonable_encoder +from fastapi.responses import JSONResponse, Response + +from lmdeploy.serve.openai.protocol import ( + AbortRequest, + DestroyWeightsUpdateGroupRequest, + InitWeightsUpdateGroupRequest, + UpdateParamsRequest, + UpdateWeightsFromDistributedRequest, +) +from lmdeploy.serve.openai.utils import create_error_response +from lmdeploy.serve.utils.server_utils import validate_json_request + + +def register(router: APIRouter, server_context) -> None: + + @router.get('/health') + async def health() -> JSONResponse: + """Health check.""" + monitor = server_context.health_monitor + if monitor is None: + data = dict(status='unhealthy', + message='Engine health monitor is not initialized.') + return JSONResponse(jsonable_encoder(data), + status_code=HTTPStatus.SERVICE_UNAVAILABLE) + data = monitor.snapshot() + if data['status'] == 'unhealthy': + data = await monitor.refresh_snapshot() + status_code = HTTPStatus.OK if data['status'] in ( + 'healthy', 'sleeping') else HTTPStatus.SERVICE_UNAVAILABLE + return JSONResponse(jsonable_encoder(data), status_code=status_code) + + @router.get('/terminate') + async def terminate(): + """Terminate server.""" + import signal + + if not server_context.allow_terminate_by_client: + return create_error_response( + HTTPStatus.BAD_REQUEST, + 'The server can not be terminated. Please add --allow-terminate-by-client when start the server.' + ) + os.kill(os.getpid(), signal.SIGTERM) + return Response(status_code=200) + + @router.post('/update_weights', + dependencies=[Depends(validate_json_request)]) + def update_params(request: UpdateParamsRequest, + raw_request: Request = None): + """Update weights for the model.""" + server_context.async_engine.engine.update_params(request) + return JSONResponse(content=None) + + def _check_pytorch_backend_for_disagg_weight_update(): + """Disaggregated weight-update endpoints are PyTorch-backend only for + now.""" + backend = getattr(server_context.async_engine, 'backend', None) + if backend != 'pytorch': + return create_error_response( + HTTPStatus.NOT_IMPLEMENTED, + f'Disaggregated weight-update endpoints require backend="pytorch", got {backend!r}.' + ) + return None + + @router.post('/init_weights_update_group', + dependencies=[Depends(validate_json_request)]) + async def init_weights_update_group(request: InitWeightsUpdateGroupRequest, + raw_request: Request = None): + """Initialize the torch.distributed process group used by an external + trainer to broadcast weights into this rollout engine.""" + err = _check_pytorch_backend_for_disagg_weight_update() + if err is not None: + return err + success, message = await server_context.async_engine.engine.init_weights_update_group( + request) + content = {'success': success, 'message': message} + return JSONResponse( + content=content, + status_code=200 if success else HTTPStatus.BAD_REQUEST) + + @router.post( + '/update_weights_from_distributed', + dependencies=[Depends(validate_json_request)], + description=('Receive a bucket of weights through a previously initialized weights-\n' + 'update group and load them into the running model.')) + async def update_weights_from_distributed( + request: UpdateWeightsFromDistributedRequest, + raw_request: Request = None): + """Receive a bucket of weights through a previously initialized + weights- update group and load them into the running model.""" + err = _check_pytorch_backend_for_disagg_weight_update() + if err is not None: + return err + success, message = await server_context.async_engine.engine.update_weights_from_distributed( + request) + content = {'success': success, 'message': message} + return JSONResponse( + content=content, + status_code=200 if success else HTTPStatus.BAD_REQUEST) + + @router.post('/destroy_weights_update_group', + dependencies=[Depends(validate_json_request)]) + async def destroy_weights_update_group( + request: DestroyWeightsUpdateGroupRequest, + raw_request: Request = None): + """Tear down a previously initialized weights-update group.""" + err = _check_pytorch_backend_for_disagg_weight_update() + if err is not None: + return err + success, message = await server_context.async_engine.engine.destroy_weights_update_group( + request) + content = {'success': success, 'message': message} + return JSONResponse( + content=content, + status_code=200 if success else HTTPStatus.BAD_REQUEST) + + @router.post('/sleep', dependencies=[Depends(validate_json_request)]) + async def sleep(raw_request: Request = None): + level = raw_request.query_params.get('level', '1') + try: + level = int(level) + except (TypeError, ValueError): + return create_error_response( + HTTPStatus.BAD_REQUEST, + 'The "level" query parameter must be an integer.') + if level not in (1, 2): + return create_error_response( + HTTPStatus.BAD_REQUEST, + 'The "level" query parameter must be 1 or 2.') + async_engine = server_context.async_engine + await async_engine.sleep(level) + return Response(status_code=200) + + @router.post('/wakeup', dependencies=[Depends(validate_json_request)]) + async def wakeup(raw_request: Request = None): + tags = raw_request.query_params.getlist('tags') + tags = tags or None + server_context.async_engine.wakeup(tags) + return Response(status_code=200) + + @router.get('/is_sleeping') + async def is_sleeping(): + is_sleeping = server_context.async_engine.is_sleeping + return JSONResponse(content={'is_sleeping': is_sleeping}) + + @router.post('/abort_request') + async def abort_request(request: AbortRequest, + raw_request: Request = None): + """Abort an ongoing request.""" + if not server_context.enable_abort_handling: + return Response( + status_code=501, + content= + 'This server does not support abort requests. Enable with --enable-abort-handling flag.' + ) + + if request.abort_all: + await server_context.async_engine.stop_all_session() + else: + session = server_context.find_session(request.session_id) + if session is None: + return create_error_response( + HTTPStatus.BAD_REQUEST, + f'Session {request.session_id} not found.') + await session.async_abort() + session_mgr = server_context.session_manager + session_mgr.remove(session) + return Response(status_code=200) diff --git a/lmdeploy/serve/openai/endpoints/models.py b/lmdeploy/serve/openai/endpoints/models.py new file mode 100644 index 0000000000..6dcb3b81d1 --- /dev/null +++ b/lmdeploy/serve/openai/endpoints/models.py @@ -0,0 +1,25 @@ +# Copyright (c) OpenMMLab. All rights reserved. +from __future__ import annotations + +from fastapi import APIRouter + +from lmdeploy.serve.openai.protocol import ( + ModelCard, + ModelList, + ModelPermission, +) +from lmdeploy.serve.openai.utils import get_model_list + + +def register(router: APIRouter, server_context) -> None: + + @router.get('/v1/models') + def available_models(): + """Show available models.""" + model_cards = [] + for model_name in get_model_list(server_context): + model_cards.append( + ModelCard(id=model_name, + root=model_name, + permission=[ModelPermission()])) + return ModelList(data=model_cards) diff --git a/lmdeploy/serve/openai/endpoints/router.py b/lmdeploy/serve/openai/endpoints/router.py new file mode 100644 index 0000000000..2a2af77cd2 --- /dev/null +++ b/lmdeploy/serve/openai/endpoints/router.py @@ -0,0 +1,19 @@ +# Copyright (c) OpenMMLab. All rights reserved. +"""Router assembly for OpenAI-compatible and LMDeploy endpoints.""" + +from fastapi import APIRouter + +from . import auxiliary, chat_completions, completions, distserve, generate, management, models + + +def create_openai_router(server_context) -> APIRouter: + """Create a router containing all legacy API-server endpoints.""" + router = APIRouter() + models.register(router, server_context) + management.register(router, server_context) + chat_completions.register(router, server_context) + completions.register(router, server_context) + generate.register(router, server_context) + auxiliary.register(router, server_context) + distserve.register(router, server_context) + return router diff --git a/lmdeploy/serve/openai/responses/serving.py b/lmdeploy/serve/openai/responses/serving.py index c1d4f4811f..fd04bc8768 100644 --- a/lmdeploy/serve/openai/responses/serving.py +++ b/lmdeploy/serve/openai/responses/serving.py @@ -35,7 +35,7 @@ def __init__(self, server_context): def _get_model_list(self) -> list[str]: model_names = [self.server_context.async_engine.model_name] - cfg = self.server_context.async_engine.backend_config + cfg = self.server_context.engine_config model_names += getattr(cfg, 'adapters', None) or [] return model_names @@ -117,7 +117,7 @@ async def create_response(self, request: ResponsesRequest, raw_request: Request) session, result_generator = self._generate(model_name, parsed_request, gen_config) created_time = int(time.time()) - session_mgr = self.server_context.async_engine.session_mgr + session_mgr = self.server_context.session_manager if request.stream: stream_generator = stream_response( diff --git a/lmdeploy/serve/openai/serving_chat_completion.py b/lmdeploy/serve/openai/serving_chat_completion.py deleted file mode 100644 index 0270dd8546..0000000000 --- a/lmdeploy/serve/openai/serving_chat_completion.py +++ /dev/null @@ -1,68 +0,0 @@ -# Copyright (c) OpenMMLab. All rights reserved. -from typing import TYPE_CHECKING - -from .protocol import ChatCompletionRequest - -if TYPE_CHECKING: - from .api_server import VariableInterface - - -def check_request(request: ChatCompletionRequest, server_context: 'VariableInterface') -> str: - engine_config = server_context.get_engine_config() - session_manager = server_context.get_session_manager() - try: - # Check logprobs settings - logprobs_mode = engine_config.logprobs_mode - logprobs = request.logprobs - top_logprobs = request.top_logprobs or 0 - return_logprob = request.return_logprob - if logprobs_mode is None: - if logprobs or top_logprobs > 0: - return (f'Logprobs({logprobs})/top_logprobs({top_logprobs}) requested ' - 'but not enabled logprobs_mode in engine configuration') - if return_logprob: - return (f'return_logprob({return_logprob}) requested ' - 'but not enabled logprobs_mode in engine configuration.') - if logprobs_mode is not None and (top_logprobs < 0 or (not logprobs and top_logprobs > 0)): - return (f'Invalid logprobs({logprobs})/top_logprobs({top_logprobs}) requested ' - 'when logprobs_mode is enabled in engine configuration.') - except AttributeError: - pass - - if session_manager.has(request.session_id): - return f'The session_id {request.session_id!r} is occupied.' - - # check sampling settings - if request.n <= 0: - return f'The n {request.n!r} must be a positive int.' - if request.top_p is not None and not (0 < request.top_p <= 1): - return f'The top_p {request.top_p!r} must be in (0, 1].' - if request.top_k is not None and request.top_k < 0: - return f'The top_k {request.top_k!r} cannot be a negative integer.' - if request.temperature is not None and not (0 <= request.temperature <= 2): - return f'The temperature {request.temperature!r} must be in [0, 2]' - - # Validate input_ids and image_data constraints. - # messages has higher priority. input_ids and image_data are only used when - # messages is empty (None, '', or []). image_data requires input_ids. - messages_empty = (request.messages is None - or request.messages == '' - or (isinstance(request.messages, list) and len(request.messages) == 0)) - if not messages_empty: - # messages is active — input_ids and image_data must not be set - if request.input_ids is not None: - return 'input_ids cannot be used when messages is non-empty. messages takes priority.' - if request.image_data is not None: - return 'image_data cannot be used when messages is non-empty. messages takes priority.' - else: - # messages is empty — input_ids and image_data are the active inputs - if request.input_ids is not None and len(request.input_ids) == 0: - return 'The input_ids must not be an empty list.' - if request.image_data is not None and request.input_ids is None: - return 'image_data requires input_ids to be set when messages is empty.' - - if request.return_routed_experts and not engine_config.enable_return_routed_experts: - return ('routed experts requested but not configured in engine configuration. ' - 'May start api_server with --enable-return-routed-experts flag.') - - return '' diff --git a/lmdeploy/serve/openai/serving_completion.py b/lmdeploy/serve/openai/serving_completion.py deleted file mode 100644 index 4e0b8d1f65..0000000000 --- a/lmdeploy/serve/openai/serving_completion.py +++ /dev/null @@ -1,37 +0,0 @@ -# Copyright (c) OpenMMLab. All rights reserved. -from typing import TYPE_CHECKING - -from .protocol import CompletionRequest - -if TYPE_CHECKING: - from .api_server import VariableInterface - - -def check_request(request: CompletionRequest, server_context: 'VariableInterface') -> str: - engine_config = server_context.get_engine_config() - session_manager = server_context.get_session_manager() - try: - # Check logprobs settings - logprobs_mode = engine_config.logprobs_mode - logprobs = request.logprobs or 0 - if logprobs > 0 and logprobs_mode is None: - return f'logprobs({logprobs}) requested but not enabled logprobs_mode in engine configuration.' - if logprobs_mode is not None and logprobs < 0: - return 'logprobs must be non-negative when logprobs_mode is enabled in engine configuration.' - except AttributeError: - pass - - if session_manager.has(request.session_id): - return f'The session_id {request.session_id!r} is occupied.' - - # check sampling settings - if request.n <= 0: - return f'The n {request.n!r} must be a positive int.' - if request.top_p is not None and not (0 < request.top_p <= 1): - return f'The top_p {request.top_p!r} must be in (0, 1].' - if request.top_k is not None and request.top_k < 0: - return f'The top_k {request.top_k!r} cannot be a negative integer.' - if request.temperature is not None and not (0 <= request.temperature <= 2): - return f'The temperature {request.temperature!r} must be in [0, 2]' - - return '' diff --git a/lmdeploy/serve/openai/serving_generate.py b/lmdeploy/serve/openai/serving_generate.py deleted file mode 100644 index 10e9dde597..0000000000 --- a/lmdeploy/serve/openai/serving_generate.py +++ /dev/null @@ -1,45 +0,0 @@ -# Copyright (c) OpenMMLab. All rights reserved. -from typing import TYPE_CHECKING - -from .protocol import GenerateReqInput - -if TYPE_CHECKING: - from .api_server import VariableInterface - - -def check_request(request: GenerateReqInput, server_context: 'VariableInterface') -> str: - engine_config = server_context.get_engine_config() - session_manager = server_context.get_session_manager() - try: - # Check logprobs settings - logprobs_mode = engine_config.logprobs_mode - return_logprob = request.return_logprob - if logprobs_mode is None and return_logprob: - return f'return_logprob({return_logprob}) requested but not enabled logprobs_mode in engine configuration.' - except AttributeError: - pass - - if (request.prompt is not None) ^ (request.input_ids is None): - return 'You must specify exactly one of prompt or input_ids' - - if request.prompt is not None and request.prompt == '': - return 'The prompt must not be an empty string' - - if request.input_ids is not None and len(request.input_ids) == 0: - return 'The input_ids must not be an empty list' - - if request.max_tokens is not None and request.max_tokens <= 0: - return f'The max_tokens {request.max_tokens!r} must be a positive integer.' - - if session_manager.has(request.session_id): - return f'The session_id {request.session_id!r} is occupied.' - - # check sampling settings - if request.top_p is not None and not (0 < request.top_p <= 1): - return f'The top_p {request.top_p!r} must be in (0, 1].' - if request.top_k is not None and request.top_k < 0: - return f'The top_k {request.top_k!r} cannot be a negative integer.' - if request.temperature is not None and not (0 <= request.temperature <= 2): - return f'The temperature {request.temperature!r} must be in [0, 2]' - - return '' diff --git a/lmdeploy/serve/openai/utils.py b/lmdeploy/serve/openai/utils.py index 9a72005193..0bf054b3a3 100644 --- a/lmdeploy/serve/openai/utils.py +++ b/lmdeploy/serve/openai/utils.py @@ -1,11 +1,15 @@ # Copyright (c) OpenMMLab. All rights reserved. +from http import HTTPStatus from typing import Any, TypeVar +from fastapi.responses import JSONResponse + from lmdeploy.serve.openai.protocol import ( ChatCompletionRequest, ChatCompletionResponseChoice, ChatCompletionResponseStreamChoice, + ErrorResponse, ) _ToolCallT = TypeVar('_ToolCallT') @@ -13,6 +17,25 @@ ChatCompletionResponseStreamChoice) +def get_model_list(server_context) -> list[str]: + """Return the model and adapter names exposed by a server.""" + model_names = [server_context.async_engine.model_name] + cfg = server_context.engine_config + model_names += getattr(cfg, 'adapters', None) or [] + return model_names + + +def create_error_response( + status: HTTPStatus, + message: str, + error_type: str = 'invalid_request_error') -> JSONResponse: + """Create an OpenAI-compatible error response.""" + payload = ErrorResponse(message=message, + type=error_type, + code=status.value) + return JSONResponse(payload.model_dump(), status_code=status.value) + + def filter_parallel_tool_calls(tool_calls: list[_ToolCallT] | None, parallel_tool_calls: bool | None) -> list[_ToolCallT] | None: """Filter to the first tool call only when parallel_tool_calls is false.""" diff --git a/lmdeploy/serve/proxy/proxy.py b/lmdeploy/serve/proxy/proxy.py index bb5231bb44..5414f77dcb 100644 --- a/lmdeploy/serve/proxy/proxy.py +++ b/lmdeploy/serve/proxy/proxy.py @@ -25,7 +25,6 @@ from lmdeploy.pytorch.disagg.conn.protocol import MigrationProtocol, MigrationRequest from lmdeploy.pytorch.disagg.conn.proxy_conn import PDConnectionPool from lmdeploy.pytorch.disagg.messages import PDConnectionMessage -from lmdeploy.serve.openai.api_server import create_error_response from lmdeploy.serve.openai.protocol import ( ChatCompletionRequest, CompletionRequest, @@ -33,6 +32,7 @@ ModelList, ModelPermission, ) +from lmdeploy.serve.openai.utils import create_error_response from lmdeploy.serve.proxy.utils import AIOHTTP_TIMEOUT, LATENCY_DEQUE_LEN, ErrorCodes, RoutingStrategy, err_msg from lmdeploy.serve.utils.server_utils import validate_json_request from lmdeploy.utils import get_logger diff --git a/tests/test_lmdeploy/serve/anthropic/test_endpoints.py b/tests/test_lmdeploy/serve/anthropic/test_endpoints.py index f1449c0307..019c736a98 100644 --- a/tests/test_lmdeploy/serve/anthropic/test_endpoints.py +++ b/tests/test_lmdeploy/serve/anthropic/test_endpoints.py @@ -142,18 +142,20 @@ def __init__( logprobs_mode=logprobs_mode, enable_return_routed_experts=enable_return_routed_experts, ) + self.async_engine.session_mgr = self.session_mgr self.default_gen_config = {} self.response_parser_cls = response_parser_cls - def create_session(self, _session_id: int | None = None): - return self.session_mgr.get() - - def get_session_manager(self): - return self.session_mgr - - def get_engine_config(self): + @property + def engine_config(self): return self.async_engine.backend_config + @property + def session_manager(self): + return self.async_engine.session_mgr + + def create_session(self, _session_id: int | None = None): + return self.session_mgr.get() class _FakeRawRequest: diff --git a/tests/test_lmdeploy/serve/openai/responses/conftest.py b/tests/test_lmdeploy/serve/openai/responses/conftest.py index 9cb5fded74..36ecb2972c 100644 --- a/tests/test_lmdeploy/serve/openai/responses/conftest.py +++ b/tests/test_lmdeploy/serve/openai/responses/conftest.py @@ -67,6 +67,14 @@ def __init__(self): self.default_gen_config = {} self.sessions = [] + @property + def engine_config(self): + return self.async_engine.backend_config + + @property + def session_manager(self): + return self.async_engine.session_mgr + def create_session(self, session_id): session = FakeSession(session_id) self.sessions.append(session) diff --git a/tests/test_lmdeploy/serve/test_generation_config.py b/tests/test_lmdeploy/serve/test_generation_config.py index 2c39ba5fb5..075785623e 100644 --- a/tests/test_lmdeploy/serve/test_generation_config.py +++ b/tests/test_lmdeploy/serve/test_generation_config.py @@ -1,5 +1,6 @@ # Copyright (c) OpenMMLab. All rights reserved. import warnings +from types import SimpleNamespace from unittest.mock import patch from pydantic import BaseModel, ConfigDict @@ -11,9 +12,9 @@ merge_gen_config, resolve_default_gen_config, ) +from lmdeploy.serve.openai.endpoints.generate import check_request as check_generate_request from lmdeploy.serve.openai.protocol import ChatCompletionRequest, CompletionRequest, GenerateReqInput from lmdeploy.serve.openai.responses import ResponsesRequest -from lmdeploy.serve.openai.serving_generate import check_request as check_generate_request _DEFAULTS = GenerationConfig() @@ -30,11 +31,19 @@ def has(self, session_id): class _FakeServerContext: - def get_engine_config(self): - return _FakeEngineConfig() + def __init__(self): + self.async_engine = SimpleNamespace( + backend_config=_FakeEngineConfig(), + session_mgr=_FakeSessionManager(), + ) - def get_session_manager(self): - return _FakeSessionManager() + @property + def engine_config(self): + return self.async_engine.backend_config + + @property + def session_manager(self): + return self.async_engine.session_mgr def test_merge_gen_config_priority(): diff --git a/tests/test_lmdeploy/serve/test_health_monitor.py b/tests/test_lmdeploy/serve/test_health_monitor.py index 83e1f09a03..9d82bf6f4f 100644 --- a/tests/test_lmdeploy/serve/test_health_monitor.py +++ b/tests/test_lmdeploy/serve/test_health_monitor.py @@ -70,7 +70,8 @@ def test_concurrent_refresh_snapshot_serializes_probes(): async def _run_health_endpoint_refreshes_cached_unhealthy_snapshot(): - from lmdeploy.serve.openai.api_server import VariableInterface, health + from lmdeploy.serve.openai.api_server import ServerContext + from lmdeploy.serve.openai.endpoints import create_openai_router class _FakeMonitor: @@ -85,12 +86,11 @@ async def refresh_snapshot(self): return dict(status='healthy', message='fresh probe succeeded') monitor = _FakeMonitor() - original_monitor = VariableInterface.health_monitor - VariableInterface.health_monitor = monitor - try: - response = await health() - finally: - VariableInterface.health_monitor = original_monitor + context = ServerContext() + context.health_monitor = monitor + router = create_openai_router(context) + health = next(route.endpoint for route in router.routes if route.path == '/health') + response = await health() assert response.status_code == 200 assert json.loads(response.body) == dict(status='healthy', message='fresh probe succeeded') @@ -102,7 +102,8 @@ def test_health_endpoint_refreshes_cached_unhealthy_snapshot(): async def _run_health_endpoint_does_not_refresh_initializing_snapshot(): - from lmdeploy.serve.openai.api_server import VariableInterface, health + from lmdeploy.serve.openai.api_server import ServerContext + from lmdeploy.serve.openai.endpoints import create_openai_router class _FakeMonitor: @@ -117,12 +118,11 @@ async def refresh_snapshot(self): return dict(status='healthy', message='fresh probe succeeded') monitor = _FakeMonitor() - original_monitor = VariableInterface.health_monitor - VariableInterface.health_monitor = monitor - try: - response = await health() - finally: - VariableInterface.health_monitor = original_monitor + context = ServerContext() + context.health_monitor = monitor + router = create_openai_router(context) + health = next(route.endpoint for route in router.routes if route.path == '/health') + response = await health() assert response.status_code == 503 assert json.loads(response.body) == dict(status='initializing', message='Engine health monitor is starting.') diff --git a/tests/test_lmdeploy/serve/test_session_cleanup.py b/tests/test_lmdeploy/serve/test_session_cleanup.py index 8b8b1b0f3f..497bce901e 100644 --- a/tests/test_lmdeploy/serve/test_session_cleanup.py +++ b/tests/test_lmdeploy/serve/test_session_cleanup.py @@ -82,53 +82,6 @@ def test_idle_async_close_removes_session_and_allows_reuse(): asyncio.run(_run_idle_async_close_removes_session_and_allows_reuse()) -async def _run_api_wrapper_cleanup_after_cancel(): - from lmdeploy.serve.openai.api_server import VariableInterface, _with_request_cleanup - - session_mgr = SessionManager() - session = session_mgr.get(260606) - closed = asyncio.Event() - never = asyncio.Event() - - class _FakeAsyncEngine: - pass - - async def result_generator(): - try: - yield 'engine' - await never.wait() - finally: - closed.set() - - async def response_generator(): - async for item in result: - yield item - - result = result_generator() - - origin_async_engine = VariableInterface.async_engine - VariableInterface.async_engine = _FakeAsyncEngine() - VariableInterface.async_engine.session_mgr = session_mgr - try: - wrapped = _with_request_cleanup(response_generator(), [result], [session]) - assert await wrapped.__anext__() == 'engine' - - next_task = asyncio.create_task(wrapped.__anext__()) - await asyncio.sleep(0) - next_task.cancel() - with suppress(asyncio.CancelledError): - await next_task - - await asyncio.wait_for(closed.wait(), timeout=1) - assert session_mgr.sessions == {} - finally: - VariableInterface.async_engine = origin_async_engine - - -def test_api_wrapper_cleanup_runs_after_cancel(): - asyncio.run(_run_api_wrapper_cleanup_after_cancel()) - - async def _run_request_cleanup_removes_unstarted_generator_session(): from lmdeploy.serve.utils.request_cleanup import with_request_cleanup