Ouzhang's picture
Add files using upload-large-folder tool
d91766b verified
Raw
History Blame Contribute Delete
9.62 kB
from __future__ import annotations
from dataclasses import dataclass, field
from diffulex.logger import get_logger
from diffulex.server.protocol import (
ChatInput,
PromptInput,
ServingAbort,
ServingBufferSnapshot,
ServingCommand,
ServingDelta,
ServingError,
ServingEvent,
ServingGenerate,
ServingReply,
)
from diffulex.utils.output import decode_token_ids_robust
logger = get_logger(__name__)
SUPPORTED_STREAM_MODES = {"block_append", "denoise"}
@dataclass
class ServingRequestState:
rid: str
engine_req_id: int
stream: bool = False
stream_mode: str = "denoise"
emitted_token_count: int = 0
emitted_text_len: int = 0
@dataclass
class ServingState:
requests: dict[int, ServingRequestState] = field(default_factory=dict)
class DiffulexAsyncEngineMixin:
"""Serving-only owner-loop helpers for DiffulexEngine.
The async HTTP frontend should enqueue ServingCommand objects and call
run_serving_tick() from exactly one engine owner thread. The mixin keeps
scheduler/model_runner mutations synchronous and serialized.
"""
def init_serving_state(self) -> None:
if not hasattr(self, "serving_state"):
self.serving_state = ServingState()
def run_serving_tick(self, commands: list[ServingCommand]) -> list[ServingEvent]:
self.init_serving_state()
events: list[ServingEvent] = []
for command in commands:
if isinstance(command, ServingGenerate):
event = self.add_serving_request(command)
if event is not None:
events.append(event)
elif isinstance(command, ServingAbort):
self.abort_serving_request(command)
else:
raise TypeError(f"Unsupported serving command: {type(command)!r}")
if not self.is_finished():
events.extend(self.step_serving_requests())
return events
def add_serving_request(self, command: ServingGenerate) -> ServingError | None:
if command.stream and command.stream_mode not in SUPPORTED_STREAM_MODES:
return ServingError(command.rid, f"Unsupported stream_mode: {command.stream_mode}")
try:
prompt = self.prepare_serving_prompt(command)
engine_req_id = self.add_request(prompt, command.sampling_params)
except Exception as exc:
return ServingError(command.rid, str(exc))
self.serving_state.requests[engine_req_id] = ServingRequestState(
rid=command.rid,
engine_req_id=engine_req_id,
stream=command.stream,
stream_mode=command.stream_mode,
)
return None
def prepare_serving_prompt(self, command: ServingGenerate) -> str | list[int]:
if isinstance(command.input, PromptInput):
return command.input.prompt
if isinstance(command.input, ChatInput):
return self.render_chat_prompt_for_serving(command.input.messages)
raise TypeError(f"Unsupported serving input: {type(command.input)!r}")
def abort_serving_request(self, command: ServingAbort) -> None:
for engine_req_id, state in list(self.serving_state.requests.items()):
if state.rid == command.rid:
del self.serving_state.requests[engine_req_id]
self.abort_request(engine_req_id)
break
def step_serving_requests(self) -> list[ServingEvent]:
reqs, _ = self.step()
events: list[ServingEvent] = []
for req in reqs:
state = self.serving_state.requests.get(req.req_id)
if state is None:
continue
if state.stream:
events.extend(self.build_stream_events(state, req))
if not req.is_finished:
continue
del self.serving_state.requests[req.req_id]
events.append(self.build_serving_reply(state.rid, req))
return events
def build_stream_events(self, state: ServingRequestState, req) -> list[ServingEvent]:
if state.stream_mode == "block_append":
event = self.build_block_append_delta(state, req)
return [event] if event is not None else []
if state.stream_mode == "denoise":
event = self.build_denoise_snapshot(state, req)
return [event] if event is not None else []
return [ServingError(state.rid, f"Unsupported stream_mode: {state.stream_mode}")]
def build_block_append_delta(self, state: ServingRequestState, req) -> ServingDelta | None:
token_ids = self.stable_generated_token_ids(req)
text = decode_token_ids_robust(self.tokenizer, token_ids)
if len(token_ids) <= state.emitted_token_count and len(text) <= state.emitted_text_len:
return None
delta = ServingDelta(
rid=state.rid,
token_offset=state.emitted_token_count,
text=text[state.emitted_text_len :],
token_ids=token_ids[state.emitted_token_count :],
nfe=int(req.nfe),
finished=req.is_finished,
)
state.emitted_token_count = len(token_ids)
state.emitted_text_len = len(text)
return delta
def build_denoise_snapshot(self, state: ServingRequestState, req) -> ServingBufferSnapshot | None:
snapshot = self.current_buffer_snapshot(req)
if snapshot is None:
return None
absolute_start, absolute_end, token_ids = snapshot
prefix_len = self.prompt_len(req)
return ServingBufferSnapshot(
rid=state.rid,
token_offset=max(0, absolute_start - prefix_len),
absolute_start=absolute_start,
absolute_end=absolute_end,
text=decode_token_ids_robust(self.tokenizer, token_ids),
token_ids=token_ids,
nfe=int(req.nfe),
finished=req.is_finished,
)
def stable_generated_token_ids(self, req) -> list[int]:
if req.is_finished:
token_ids = list(req.truncated_response)
return self.trim_at_first_mask_token(token_ids, req)
if not req.is_multi_block:
return []
buffer = req.dllm_block_buffer
if buffer is None:
return []
prefix_len = self.prompt_len(req)
stable_abs_end = max(prefix_len, min(buffer.first_running_block.start, len(req.token_ids)))
token_ids = list(req.token_ids[prefix_len:stable_abs_end])
return self.trim_at_first_mask_token(token_ids, req)
def trim_at_first_mask_token(self, token_ids: list[int], req) -> list[int]:
mask_token_id = self.mask_token_id(req)
if mask_token_id is None or mask_token_id not in token_ids:
return token_ids
return token_ids[: token_ids.index(mask_token_id)]
def drop_mask_tokens(self, token_ids: list[int], req) -> list[int]:
mask_token_id = self.mask_token_id(req)
if mask_token_id is None:
return token_ids
return [token_id for token_id in token_ids if token_id != mask_token_id]
def mask_token_id(self, req) -> int | None:
req_mask = req.mask_token_id if req.is_multi_block else None
if req_mask is not None:
return int(req_mask)
tokenizer_mask = getattr(self.tokenizer, "mask_token_id", None)
return int(tokenizer_mask) if tokenizer_mask is not None else None
def current_buffer_snapshot(self, req) -> tuple[int, int, list[int]] | None:
prefix_len = self.prompt_len(req)
if req.is_multi_block:
buffer = req.dllm_block_buffer
if buffer is None:
return None
absolute_start = prefix_len
absolute_end = min(len(req.token_ids), buffer.last_running_block.end)
if absolute_end <= absolute_start:
return None
return absolute_start, absolute_end, list(req.token_ids[absolute_start:absolute_end])
token_ids = list(req.truncated_response)
if not token_ids:
return None
return prefix_len, prefix_len + len(token_ids), token_ids
def prompt_len(self, req) -> int:
return int(req.prefix_len if req.is_multi_block else req.num_prompt_tokens)
def build_serving_reply(self, rid: str, req) -> ServingReply:
token_ids = self.drop_mask_tokens(list(req.truncated_response), req)
full_token_ids = self.drop_mask_tokens(list(req.full_response), req)
eos = getattr(self.tokenizer, "eos_token", None) or ""
raw_text = decode_token_ids_robust(self.tokenizer, token_ids)
text = raw_text.split(eos)[0] if eos else raw_text
full_text = decode_token_ids_robust(self.tokenizer, full_token_ids)
return ServingReply(
rid=rid,
text=text,
token_ids=token_ids,
nfe=int(req.nfe),
finish_reason=req.completion_reason,
full_text=full_text,
full_token_ids=full_token_ids,
)
def render_chat_prompt_for_serving(self, messages: list[dict[str, str]]) -> str:
tokenizer = self.tokenizer
if hasattr(tokenizer, "apply_chat_template"):
try:
return tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
except Exception:
logger.warning("Tokenizer chat template failed; using plain chat fallback", exc_info=True)
return "\n".join(f"{m.get('role', 'user')}: {m.get('content', '')}" for m in messages) + "\nassistant:"