Optimize streaming detokenizer updates (#24659)
Co-authored-by: Alex Nails <alex.nails@radixark.ai>
This commit is contained in:
@@ -68,6 +68,22 @@ class DecodeStatus:
|
|||||||
read_offset: int
|
read_offset: int
|
||||||
# Offset that's sent to tokenizer for incremental update.
|
# Offset that's sent to tokenizer for incremental update.
|
||||||
sent_offset: int = 0
|
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):
|
class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
||||||
@@ -188,42 +204,61 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
|
|||||||
space_list: List[bool],
|
space_list: List[bool],
|
||||||
) -> List[str]:
|
) -> List[str]:
|
||||||
"""Batch decode with grouping by (skip_special_tokens, spaces_between_special_tokens)."""
|
"""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):
|
if not getattr(self.tokenizer, "is_fast", False):
|
||||||
return [
|
decoded = [
|
||||||
decode_without_hf_kwargs(self.tokenizer, ids, skip)
|
decode_without_hf_kwargs(self.tokenizer, ids, skip)
|
||||||
for ids, skip in zip(ids_list, skip_list)
|
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
|
decoded = [""] * len(ids_list)
|
||||||
first_skip, first_space = skip_list[0], space_list[0]
|
for (skip, space), indices in groups.items():
|
||||||
if all(
|
group_decoded = self.tokenizer.batch_decode(
|
||||||
s == first_skip and sp == first_space
|
[ids_list[idx] for idx in indices],
|
||||||
for s, sp in zip(skip_list, space_list)
|
skip_special_tokens=skip,
|
||||||
):
|
spaces_between_special_tokens=space,
|
||||||
return self.tokenizer.batch_decode(
|
)
|
||||||
ids_list,
|
for idx, text in zip(indices, group_decoded):
|
||||||
skip_special_tokens=first_skip,
|
decoded[idx] = text
|
||||||
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
|
|
||||||
|
|
||||||
|
if keep_idx is None:
|
||||||
|
return decoded
|
||||||
|
results = [""] * n
|
||||||
|
for i, text in zip(keep_idx, decoded):
|
||||||
|
results[i] = text
|
||||||
return results
|
return results
|
||||||
|
|
||||||
def _decode_batch_token_id_output(self, recv_obj: BatchTokenIDOutput):
|
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]) :]
|
new_text = read_texts[i][len(surr_texts[i]) :]
|
||||||
if recv_obj.finished_reasons[i] is None:
|
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("�"):
|
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.surr_offset = s.read_offset
|
||||||
s.read_offset = len(s.decode_ids)
|
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:
|
else:
|
||||||
new_text = find_printable_text(new_text)
|
# Incomplete UTF-8: emit the printable prefix only; do not
|
||||||
else:
|
# commit (token offsets stay so the next iteration retries
|
||||||
if rid in self.decode_status:
|
# with more tokens).
|
||||||
del self.decode_status[rid]
|
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(
|
output_str = self.trim_matched_stop(
|
||||||
s.decoded_text + new_text,
|
s.get_decoded_text() + new_text,
|
||||||
recv_obj.finished_reasons[i],
|
recv_obj.finished_reasons[i],
|
||||||
recv_obj.no_stop_trim[i],
|
recv_obj.no_stop_trim[i],
|
||||||
)
|
)
|
||||||
|
|
||||||
# Incrementally send text.
|
|
||||||
incremental_output = output_str[s.sent_offset :]
|
incremental_output = output_str[s.sent_offset :]
|
||||||
s.sent_offset = len(output_str)
|
s.sent_offset = len(output_str)
|
||||||
output_strs.append(incremental_output)
|
output_strs.append(incremental_output)
|
||||||
|
|||||||
Reference in New Issue
Block a user