Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
108 changes: 88 additions & 20 deletions lmdeploy/serve/openai/api_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
import time
from collections.abc import AsyncGenerator
from contextlib import aclosing, asynccontextmanager
from dataclasses import dataclass
from functools import partial
from http import HTTPStatus
from typing import TYPE_CHECKING, Literal
Expand Down Expand Up @@ -289,6 +290,62 @@ def _create_output_token_logprobs(token_ids: list[int] | None = None,
return output_token_logprobs or None


@dataclass
class _StreamTokenMetadata:
"""Token metadata buffered across parser steps with no visible delta."""

token_ids: list[int] | None
logprobs: list[dict[int, float]] | None

@classmethod
def from_result(
cls,
token_ids: list[int] | None,
logprobs: list[dict[int, float]] | None,
*,
keep_token_ids: bool,
keep_logprobs: bool,
) -> _StreamTokenMetadata:
if keep_logprobs and logprobs is not None and len(token_ids or []) != len(logprobs):
raise ValueError('Token ids and logprobs must have the same length.')
return cls(
list(token_ids or []) if keep_token_ids else None,
list(logprobs or []) if keep_logprobs else None,
)

def extend(self, other: _StreamTokenMetadata) -> None:
if self.token_ids is not None:
assert other.token_ids is not None
self.token_ids.extend(other.token_ids)
if self.logprobs is not None:
assert other.logprobs is not None
self.logprobs.extend(other.logprobs)

def pop_with(self, current: _StreamTokenMetadata) -> _StreamTokenMetadata:
merged = _StreamTokenMetadata(
token_ids=None if self.token_ids is None else self.token_ids + (current.token_ids or []),
logprobs=None if self.logprobs is None else self.logprobs + (current.logprobs or []),
)
if self.token_ids is not None:
self.token_ids.clear()
if self.logprobs is not None:
self.logprobs.clear()
return merged

def output_ids(self, enabled: bool | None) -> list[int] | None:
return self.token_ids if enabled else None

def chat_logprobs(self, tokenizer: PreTrainedTokenizerBase, enabled: bool | None) -> ChoiceLogprobs | None:
if not enabled or not self.token_ids or not self.logprobs:
return None
return _create_chat_completion_logprobs(tokenizer, self.token_ids, self.logprobs)

def output_token_logprobs(self, enabled: bool | None) -> list[tuple[float, int]] | None:
if not enabled:
return None
return _create_output_token_logprobs(self.token_ids, self.logprobs)


@router.get('/health')
async def health() -> JSONResponse:
"""Health check."""
Expand Down Expand Up @@ -587,31 +644,39 @@ def create_stream_usage_response_json(usage: UsageInfo) -> str:
async def completion_stream_generator() -> AsyncGenerator[str, None]:
streaming_tools = False
final_usage = None
keep_logprobs = bool(request.logprobs or request.return_logprob)
keep_token_ids = bool(request.return_token_ids or keep_logprobs)
pending_token_metadata = _StreamTokenMetadata(
token_ids=[] if keep_token_ids else None,
logprobs=[] if keep_logprobs else 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 []
raw_token_ids = res.token_ids or []
current_token_metadata = _StreamTokenMetadata.from_result(
res.token_ids,
res.logprobs,
keep_token_ids=keep_token_ids,
keep_logprobs=keep_logprobs,
)
stream_deltas = response_parser.stream_chunk(
res.response,
delta_token_ids
raw_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:
pending_token_metadata.extend(current_token_metadata)
if res.finish_reason is None:
continue
stream_deltas = [(DeltaMessage(role='assistant', content=''), False)]
current_token_metadata = _StreamTokenMetadata(
token_ids=[] if keep_token_ids else None,
logprobs=[] if keep_logprobs else None,
)
stream_deltas = [(DeltaMessage(role='assistant'), False)]
should_validate_complete = (
res.finish_reason in ('stop', 'length')
and (request.return_token_ids or request.return_routed_experts)
Expand All @@ -627,19 +692,22 @@ async def completion_stream_generator() -> AsyncGenerator[str, None]:
# 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 is_last_delta:
stream_token_metadata = pending_token_metadata.pop_with(current_token_metadata)
chunk_logprobs = stream_token_metadata.chat_logprobs(tokenizer, request.logprobs)
chunk_output_token_logprobs = stream_token_metadata.output_token_logprobs(request.return_logprob)
stream_output_ids = stream_token_metadata.output_ids(request.return_token_ids)
else:
chunk_logprobs = None
chunk_output_token_logprobs = None
stream_output_ids = 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,
Expand All @@ -649,7 +717,7 @@ async def completion_stream_generator() -> AsyncGenerator[str, None]:
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
response_json['remote_token_ids'] = raw_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'
Expand Down
19 changes: 3 additions & 16 deletions lmdeploy/serve/parsers/response_parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -177,7 +177,7 @@ def stream_chunk(self,
Returns:
A list of ``(delta_message, tool_calls_emitted)`` pairs. Return
``[]`` when this engine step produces no visible delta (for example
while buffering a partial protocol tag).
while buffering protocol syntax or tool-call payload).
"""
raise NotImplementedError

Expand Down Expand Up @@ -322,7 +322,7 @@ def stream_chunk(
from this stream step. Multiple entries may be returned when one
engine chunk contains reasoning, content, and tool-call segments.
Return ``[]`` when this engine step produces no visible delta (for
example while buffering a partial protocol tag).
example while buffering protocol syntax or tool-call payload).
"""
# Special-case: some backends emit a leading empty delta (no text, no
# tokens) before any actual content. Tests treat this as a visible empty
Expand Down Expand Up @@ -555,8 +555,7 @@ def _consume_tool(self) -> tuple[list[DeltaToolCall], bool]:
emit = self._pending
self._pending = ''
out = self.tool_parser.decode_tool_incremental(added_text=emit, final=False)
if (self.profile.tool_payload_format == 'json'
and self._is_complete_json_object(self.tool_parser._tool_payload)):
if self.profile.tool_payload_format == 'json' and self.tool_parser._payload_closed:
out.extend(self.tool_parser.decode_tool_incremental(added_text='', final=True))
self.tool_parser.finish_tool_call()
self._mode = self.MODE_PLAIN
Comment on lines 557 to 561
Expand Down Expand Up @@ -735,15 +734,3 @@ def _longest_open_tag_prefix_suffix(text: str, tags: list[str]) -> int:
best = k
break
return best

@staticmethod
def _is_complete_json_object(payload: str) -> bool:
payload = payload.strip()
if not payload:
return False
decoder = json.JSONDecoder()
try:
obj, end = decoder.raw_decode(payload)
except json.JSONDecodeError:
return False
return isinstance(obj, dict) and end == len(payload)
2 changes: 2 additions & 0 deletions lmdeploy/serve/parsers/tool_parser/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
from .glm47_tool_parser import Glm47ToolParser
from .internlm2_tool_parser import Internlm2ToolParser
from .interns2preview_tool_parser import InternS2PreviewToolParser
from .json_tool_parser import JsonToolParser
from .llama3_tool_parser import Llama3JsonToolParser
from .qwen2d5_tool_parser import Qwen2d5ToolParser
from .qwen3_tool_parser import Qwen3ToolParser
Expand All @@ -14,6 +15,7 @@
__all__ = [
'ToolParser',
'ToolParserManager',
'JsonToolParser',
'XmlToolParser',
'DeepSeekV32ToolParser',
'DeepSeekV4ToolParser',
Expand Down
Loading
Loading