From e5b8e3a66aa6052d86905869a7dd7c816c8401f7 Mon Sep 17 00:00:00 2001 From: inkcherry Date: Wed, 3 Jun 2026 14:44:21 +0800 Subject: [PATCH] Optimize streaming detokenizer updates (#24659) Co-authored-by: Alex Nails --- .../srt/managers/detokenizer_manager.py | 124 ++++++++++++------ 1 file changed, 85 insertions(+), 39 deletions(-) diff --git a/python/sglang/srt/managers/detokenizer_manager.py b/python/sglang/srt/managers/detokenizer_manager.py index 5f67f1db1..065546088 100644 --- a/python/sglang/srt/managers/detokenizer_manager.py +++ b/python/sglang/srt/managers/detokenizer_manager.py @@ -68,6 +68,22 @@ class DecodeStatus: read_offset: int # Offset that's sent to tokenizer for incremental update. sent_offset: int = 0 + decoded_text_len: int = dataclasses.field(init=False) + decoded_text_chunks: List[str] = dataclasses.field(default_factory=list) + + def __post_init__(self): + self.decoded_text_len = len(self.decoded_text) + + def append_decoded_text(self, text: str): + if text: + self.decoded_text_chunks.append(text) + self.decoded_text_len += len(text) + + def get_decoded_text(self) -> str: + if self.decoded_text_chunks: + self.decoded_text += "".join(self.decoded_text_chunks) + self.decoded_text_chunks.clear() + return self.decoded_text class DetokenizerManager(MultiHttpWorkerDetokenizerMixin): @@ -188,42 +204,61 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin): space_list: List[bool], ) -> List[str]: """Batch decode with grouping by (skip_special_tokens, spaces_between_special_tokens).""" + n = len(ids_list) + if n == 0: + return [] + + # Empty token spans decode to "" but tokenizer.batch_decode (and the + # slow per-row decode_without_hf_kwargs path) still pays per-row + # overhead; under high-concurrency streaming this adds up. Filter + # empties out, decode the rest, then scatter back. + keep_idx: Optional[List[int]] = None + if not all(ids_list): + keep_idx = [i for i, ids in enumerate(ids_list) if ids] + if not keep_idx: + return [""] * n + ids_list = [ids_list[i] for i in keep_idx] + skip_list = [skip_list[i] for i in keep_idx] + space_list = [space_list[i] for i in keep_idx] if not getattr(self.tokenizer, "is_fast", False): - return [ + decoded = [ decode_without_hf_kwargs(self.tokenizer, ids, skip) for ids, skip in zip(ids_list, skip_list) ] + else: + # fast path: all rows share the same (skip, space) flags. + first_skip, first_space = skip_list[0], space_list[0] + if all( + s == first_skip and sp == first_space + for s, sp in zip(skip_list, space_list) + ): + decoded = self.tokenizer.batch_decode( + ids_list, + skip_special_tokens=first_skip, + spaces_between_special_tokens=first_space, + ) + else: + # Group indices by (skip, space) tuple and decode each group. + groups: Dict[Tuple[bool, bool], List[int]] = defaultdict(list) + for idx, (skip, space) in enumerate(zip(skip_list, space_list)): + groups[(skip, space)].append(idx) - # fast path - first_skip, first_space = skip_list[0], space_list[0] - if all( - s == first_skip and sp == first_space - for s, sp in zip(skip_list, space_list) - ): - return self.tokenizer.batch_decode( - ids_list, - skip_special_tokens=first_skip, - spaces_between_special_tokens=first_space, - ) - - # Group indices by (skip, space) tuple - groups: Dict[Tuple[bool, bool], List[int]] - groups = defaultdict(list) - for idx, (skip, space) in enumerate(zip(skip_list, space_list)): - groups[(skip, space)].append(idx) - - # Decode each group and collect results - results: List[str] = [""] * len(ids_list) - for (skip, space), indices in groups.items(): - decoded = self.tokenizer.batch_decode( - [ids_list[idx] for idx in indices], - skip_special_tokens=skip, - spaces_between_special_tokens=space, - ) - for idx, text in zip(indices, decoded): - results[idx] = text + decoded = [""] * len(ids_list) + for (skip, space), indices in groups.items(): + group_decoded = self.tokenizer.batch_decode( + [ids_list[idx] for idx in indices], + skip_special_tokens=skip, + spaces_between_special_tokens=space, + ) + for idx, text in zip(indices, group_decoded): + decoded[idx] = text + if keep_idx is None: + return decoded + results = [""] * n + for i, text in zip(keep_idx, decoded): + results[i] = text return results def _decode_batch_token_id_output(self, recv_obj: BatchTokenIDOutput): @@ -306,25 +341,36 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin): ) new_text = read_texts[i][len(surr_texts[i]) :] if recv_obj.finished_reasons[i] is None: - # Streaming chunk: update the decode status + # Streaming. Invariant: sent_offset >= decoded_text_len. The + # gap (`pending`) is "printable but uncommitted" text emitted + # in a prior "�" recovery step; we skip it from this step's + # emission so we don't double-send. + pending = s.sent_offset - s.decoded_text_len if new_text and not new_text.endswith("�"): - s.decoded_text += new_text + # Clean text: commit to decoded_text and advance offsets. + s.append_decoded_text(new_text) s.surr_offset = s.read_offset s.read_offset = len(s.decode_ids) - new_text = "" + s.sent_offset = s.decoded_text_len + output_strs.append(new_text[pending:] if pending else new_text) else: - new_text = find_printable_text(new_text) - else: - if rid in self.decode_status: - del self.decode_status[rid] + # Incomplete UTF-8: emit the printable prefix only; do not + # commit (token offsets stay so the next iteration retries + # with more tokens). + printable = find_printable_text(new_text) + s.sent_offset = s.decoded_text_len + len(printable) + output_strs.append(printable[pending:] if pending else printable) + continue + if rid in self.decode_status: + del self.decode_status[rid] + + # Finished: materialize once, trim the matched stop, emit the tail. output_str = self.trim_matched_stop( - s.decoded_text + new_text, + s.get_decoded_text() + new_text, recv_obj.finished_reasons[i], recv_obj.no_stop_trim[i], ) - - # Incrementally send text. incremental_output = output_str[s.sent_offset :] s.sent_offset = len(output_str) output_strs.append(incremental_output)