[tokenizer] eliminate O(n²) copy in non-incremental streaming (#22567)

Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Alex Nails
2026-04-11 23:05:36 -07:00
committed by GitHub
co-authored by Claude Opus 4.6
parent 45472d70cc
commit c6fd9a00c7
+43 -13
View File
@@ -1268,6 +1268,17 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
else: else:
out = out_list[-1] out = out_list[-1]
# Resolve deferred text for non-incremental streaming.
# _handle_batch_output sets "text": None on intermediate chunks
# to avoid O(n) string rebuild per step (O(n^2) total).
if (
is_stream
and not incremental_stream
and "text" in out
and out["text"] is None
):
out["text"] = state.get_text()
if finished: if finished:
# For non-streaming cases, response has not been sent yet (`response_sent_to_client_time` has not been set yet). # For non-streaming cases, response has not been sent yet (`response_sent_to_client_time` has not been set yet).
# Record response sent time right before we log finished results and metrics. # Record response sent time right before we log finished results and metrics.
@@ -1708,15 +1719,26 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
output_token_ids = delta_output_ids output_token_ids = delta_output_ids
_slice_streaming_output_meta_info(meta_info, output_offset) _slice_streaming_output_meta_info(meta_info, output_offset)
state.last_output_offset = len(state.output_ids) state.last_output_offset = len(state.output_ids)
output_text = delta_text out_dict = {
"text": delta_text,
"output_ids": output_token_ids,
"meta_info": meta_info,
}
elif state.finished:
out_dict = {
"text": state.get_text(),
"output_ids": state.output_ids.copy(),
"meta_info": meta_info,
}
else: else:
output_token_ids = state.output_ids.copy() # Non-incremental intermediate: pass reference (no
output_text = state.get_text() # copy) and defer text to _wait_one_response to avoid
out_dict = { # O(n) per-step cost that compounds to O(n^2).
"text": output_text, out_dict = {
"output_ids": output_token_ids, "text": None,
"meta_info": meta_info, "output_ids": state.output_ids,
} "meta_info": meta_info,
}
elif state.finished: elif state.finished:
out_dict = { out_dict = {
"text": state.get_text(), "text": state.get_text(),
@@ -1739,12 +1761,20 @@ class TokenizerManager(TokenizerCommunicatorMixin, TokenizerManagerScoreMixin):
output_token_ids = delta_output_ids output_token_ids = delta_output_ids
_slice_streaming_output_meta_info(meta_info, output_offset) _slice_streaming_output_meta_info(meta_info, output_offset)
state.last_output_offset = len(state.output_ids) state.last_output_offset = len(state.output_ids)
out_dict = {
"output_ids": output_token_ids,
"meta_info": meta_info,
}
elif state.finished:
out_dict = {
"output_ids": state.output_ids.copy(),
"meta_info": meta_info,
}
else: else:
output_token_ids = state.output_ids.copy() out_dict = {
out_dict = { "output_ids": state.output_ids,
"output_ids": output_token_ids, "meta_info": meta_info,
"meta_info": meta_info, }
}
elif state.finished: elif state.finished:
out_dict = { out_dict = {
"output_ids": state.output_ids.copy(), "output_ids": state.output_ids.copy(),