[Kimi-K3] Recover the reply when the model skips the think channel (#37743)

Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com>
This commit is contained in:
zijiexia
2026-09-07 14:08:42 +00:00
committed by GitHub
co-authored by Xinyuan Tong
parent e4008de757
commit b5c9b68f03
2 changed files with 188 additions and 17 deletions
+87 -11
View File
@@ -541,6 +541,15 @@ class KimiK3Detector(BaseReasoningFormatDetector):
Post-reasoning content is unwrapped from the XTML ``response`` /
``message`` markers; a ``tools`` channel is passed through raw for the
kimi_k3 tool-call detector.
The model does not always honour the pre-filled think channel: on very long
prompts (~1M tokens) it sometimes emits a zero-length think section and
writes the reply directly, closing with
``<|close|>response<|sep|><|close|>message<|sep|>`` and never producing
``<|close|>think<|sep|>`` or ``<|open|>response<|sep|>``. A bare
``<|close|>response<|sep|>`` therefore proves the preceding text was the
response channel, and it is reported as content rather than reasoning
(see :meth:`_skipped_think_channel`).
"""
def __init__(
@@ -581,6 +590,8 @@ class KimiK3Detector(BaseReasoningFormatDetector):
self._reasoning_done = False
self._tools_passthrough = False
self._stream_text = ""
self._streamed_reasoning: list[str] = []
self._discard_delayed_think_close = False
def _clean_content(self, text: str) -> str:
tools_idx = text.find(TOOLS_OPEN)
@@ -596,6 +607,22 @@ class KimiK3Detector(BaseReasoningFormatDetector):
]
return min(found) if found else -1
@staticmethod
def _skipped_think_channel(
text: str,
start: int = 0,
think_close_idx: int = -1,
next_channel_idx: int = -1,
) -> bool:
response_close_idx = text.find(RESPONSE_CLOSE, start)
return response_close_idx != -1 and all(
boundary_idx == -1 or response_close_idx < boundary_idx
for boundary_idx in (think_close_idx, next_channel_idx)
)
def _clean_skipped_think_content(self, text: str) -> str:
return self._clean_content(text.replace(self.think_end_token, ""))
def detect_and_parse(self, text: str) -> StreamingParseResult:
in_reasoning = self._in_reasoning or self.think_start_token in text
if not in_reasoning and self.think_end_token not in text:
@@ -607,13 +634,17 @@ class KimiK3Detector(BaseReasoningFormatDetector):
start = open_idx + len(self.think_start_token) if open_idx != -1 else 0
close_idx = text.find(self.think_end_token, start)
tools_idx = text.find(self.tool_start_token, start)
channel_idx = self._next_channel_idx(text, start)
if self._skipped_think_channel(text, start, close_idx, channel_idx):
return StreamingParseResult(
normal_text=self._clean_skipped_think_content(text[start:])
)
if close_idx != -1 and tools_idx != -1 and tools_idx < close_idx:
return StreamingParseResult(
reasoning_text=strip_partial_marker_suffix(text[start:tools_idx]),
normal_text=self._clean_content(text[tools_idx:]),
)
if close_idx == -1:
channel_idx = self._next_channel_idx(text, start)
if channel_idx != -1:
return StreamingParseResult(
reasoning_text=strip_partial_marker_suffix(text[start:channel_idx]),
@@ -665,22 +696,38 @@ class KimiK3Detector(BaseReasoningFormatDetector):
close_idx = buf.find(self.think_end_token)
tools_idx = buf.find(self.tool_start_token)
channel_idx = self._next_channel_idx(buf)
if self._skipped_think_channel(
buf,
think_close_idx=close_idx,
next_channel_idx=channel_idx,
):
replay = "".join(self._streamed_reasoning)
self._streamed_reasoning.clear()
self._in_reasoning = False
self._reasoning_done = True
self._discard_delayed_think_close = True
return StreamingParseResult(
normal_text=(replay + self._drain_content()) or None
)
if close_idx != -1 and not (tools_idx != -1 and tools_idx < close_idx):
reasoning_text = buf[:close_idx]
self._buffer = buf[close_idx + len(self.think_end_token) :]
self._in_reasoning = False
self._reasoning_done = True
self._streamed_reasoning.clear()
return StreamingParseResult(
reasoning_text=reasoning_text or None,
normal_text=self._drain_content() or None,
)
channel_idx = self._next_channel_idx(buf)
if channel_idx != -1:
reasoning_text = strip_partial_marker_suffix(buf[:channel_idx])
self._buffer = buf[channel_idx:]
self._in_reasoning = False
self._reasoning_done = True
self._streamed_reasoning.clear()
self._tools_passthrough = buf.startswith(
self.tool_start_token, channel_idx
)
@@ -691,18 +738,26 @@ class KimiK3Detector(BaseReasoningFormatDetector):
if not self.stream_reasoning:
return StreamingParseResult()
markers = [self.think_end_token, self.tool_start_token, RESPONSE_OPEN]
markers = [
self.think_end_token,
self.tool_start_token,
RESPONSE_OPEN,
RESPONSE_CLOSE,
MESSAGE_CLOSE,
]
if not self.stripped_think_start:
markers.append(self.think_start_token)
holdback = partial_suffix_len(buf, markers)
emit = buf[: len(buf) - holdback] if holdback else buf
emit = strip_partial_marker_suffix(emit)
self._buffer = buf[len(emit) :]
self._streamed_reasoning.append(emit)
return StreamingParseResult(reasoning_text=emit)
return StreamingParseResult(normal_text=self._drain_content())
def finish(self) -> StreamingParseResult:
self._streamed_reasoning.clear()
if not self._force_nonempty_content:
return super().finish()
text, self._stream_text = self._stream_text, ""
@@ -724,25 +779,46 @@ class KimiK3Detector(BaseReasoningFormatDetector):
if not buf:
return ""
if self._tools_passthrough:
self._buffer = ""
return buf
holdback = (
partial_suffix_len(buf, [self.think_end_token])
if self._discard_delayed_think_close
else 0
)
emit = buf[: len(buf) - holdback] if holdback else buf
self._buffer = buf[len(emit) :]
if self._discard_delayed_think_close:
emit = emit.replace(self.think_end_token, "")
return emit
tools_idx = buf.find(TOOLS_OPEN)
if tools_idx != -1:
head = buf[:tools_idx]
holdback = (
partial_suffix_len(buf, [self.think_end_token])
if self._discard_delayed_think_close
else 0
)
emit = buf[: len(buf) - holdback] if holdback else buf
self._buffer = buf[len(emit) :]
head = emit[:tools_idx]
tail = emit[tools_idx:]
for marker in (RESPONSE_OPEN, RESPONSE_CLOSE, MESSAGE_CLOSE):
head = head.replace(marker, "")
if self._discard_delayed_think_close:
head = head.replace(self.think_end_token, "")
tail = tail.replace(self.think_end_token, "")
self._tools_passthrough = True
self._buffer = ""
return head + buf[tools_idx:]
return head + tail
holdback = partial_suffix_len(
buf, [RESPONSE_OPEN, RESPONSE_CLOSE, MESSAGE_CLOSE, TOOLS_OPEN]
)
markers = [RESPONSE_OPEN, RESPONSE_CLOSE, MESSAGE_CLOSE, TOOLS_OPEN]
if self._discard_delayed_think_close:
markers.append(self.think_end_token)
holdback = partial_suffix_len(buf, markers)
emit = buf[: len(buf) - holdback] if holdback else buf
self._buffer = buf[len(emit) :]
for marker in (RESPONSE_OPEN, RESPONSE_CLOSE, MESSAGE_CLOSE):
emit = emit.replace(marker, "")
if self._discard_delayed_think_close:
emit = emit.replace(self.think_end_token, "")
return emit