From 9d3ccfdffc0cd95b442d9e94e0b13a180faa783d Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Sat, 18 Jul 2026 13:05:36 +0800 Subject: [PATCH] fix: incremental decode to avoid U+FFFD in streaming - StreamDecoder buffers incomplete multi-byte sequences - Task.decode_new_token replaces per-token decode in scheduler - flush_remaining emits final buffered text on task finish --- astrai/inference/core/scheduler.py | 10 +++-- astrai/inference/core/task.py | 62 ++++++++++++++++++++++++++++++ 2 files changed, 68 insertions(+), 4 deletions(-) diff --git a/astrai/inference/core/scheduler.py b/astrai/inference/core/scheduler.py index 3b8c8fa..70c55e1 100644 --- a/astrai/inference/core/scheduler.py +++ b/astrai/inference/core/scheduler.py @@ -154,13 +154,15 @@ class InferenceScheduler: for t, ntok in zip(valid, next_tokens): t.output_ids.append(ntok) t.output_tokens += 1 - self._task_mgr.invoke_callback( - t.task_id, - self._task_mgr.tokenizer.decode([ntok]), - ) + new_text = t.decode_new_token(self._task_mgr.tokenizer) + if new_text: + self._task_mgr.invoke_callback(t.task_id, new_text) for t in valid: if t.is_finished(stop_ids): + remaining = t.flush_remaining(self._task_mgr.tokenizer) + if remaining: + self._task_mgr.invoke_callback(t.task_id, remaining) self._task_mgr.invoke_callback(t.task_id, STOP) except Exception as e: diff --git a/astrai/inference/core/task.py b/astrai/inference/core/task.py index 8006567..7af11ca 100644 --- a/astrai/inference/core/task.py +++ b/astrai/inference/core/task.py @@ -13,6 +13,40 @@ logger = logging.getLogger(__name__) STOP = object() +class StreamDecoder: + """Incremental decoder for byte-level BPE streaming. + + Byte-level BPE may split a single Unicode character (e.g. em-dash, + smart quotes) across multiple tokens. Decoding such a token in + isolation produces U+FFFD (replacement char). This decoder + accumulates token IDs and only emits text once the trailing + characters are complete, buffering incomplete multi-byte sequences + until the next token arrives. + """ + + __slots__ = ("_tokenizer", "_ids", "_emitted") + + def __init__(self, tokenizer: AutoTokenizer): + self._tokenizer = tokenizer + self._ids: List[int] = [] + self._emitted: str = "" + + def push(self, token_id: int) -> str: + """Append a token ID and return newly completed text. + + Returns "" while a multi-byte character is still incomplete. + """ + self._ids.append(token_id) + full = self._tokenizer.decode(self._ids, skip_special_tokens=True) + if full.endswith("\ufffd"): + return "" + if len(full) > len(self._emitted): + diff = full[len(self._emitted) :] + self._emitted = full + return diff + return "" + + class TaskStatus(Enum): """Task lifecycle states.""" @@ -51,6 +85,34 @@ class Task: self.output_tokens: int = 0 self.arrival_time = time.time() self.finish_time: Optional[float] = None + self._decoder: Optional[StreamDecoder] = None + + def decode_new_token(self, tokenizer: AutoTokenizer) -> str: + """Decode the last appended output token, buffering incomplete + multi-byte sequences across calls. + + Lazily creates a :class:`StreamDecoder` on first use. + """ + if self._decoder is None: + self._decoder = StreamDecoder(tokenizer) + return self._decoder.push(self.output_ids[-1]) + + def flush_remaining(self, tokenizer: AutoTokenizer) -> str: + """Emit any text still buffered in the decoder. + + Called when generation terminates (max_tokens reached, stop + sequence, or external removal) to avoid dropping a final + incomplete-looking fragment that is actually complete when + adjacent to the stop token. + """ + if self._decoder is None or not self.output_ids: + return "" + full = tokenizer.decode(self.output_ids, skip_special_tokens=True) + if len(full) > len(self._decoder._emitted): + diff = full[len(self._decoder._emitted) :] + self._decoder._emitted = full + return diff + return "" @property def next_pos(self) -> int: