[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
@@ -186,6 +186,103 @@ def test_streaming_tools_channel_before_think_close(chunk_size: int) -> None:
assert TOOLS_OPEN in content
_SKIPPED_THINK_REPLY = "The TTL is 20 minutes."
@pytest.mark.parametrize(
"text",
[
f"{_SKIPPED_THINK_REPLY}{RESPONSE_CLOSE}{MESSAGE_CLOSE}",
f"{THINK_OPEN}{_SKIPPED_THINK_REPLY}{RESPONSE_CLOSE}{MESSAGE_CLOSE}",
f"{_SKIPPED_THINK_REPLY}{RESPONSE_CLOSE}",
],
)
def test_non_stream_skipped_think_channel_is_content(text: str) -> None:
detector = KimiK3Detector(force_reasoning=True)
result = detector.detect_and_parse(text)
assert result.reasoning_text == ""
assert result.normal_text == _SKIPPED_THINK_REPLY
def test_non_stream_skipped_think_channel_does_not_affect_real_reasoning() -> None:
detector = KimiK3Detector(force_reasoning=True)
result = detector.detect_and_parse(
f"deep thought{THINK_CLOSE}{RESPONSE_OPEN}{_SKIPPED_THINK_REPLY}"
f"{RESPONSE_CLOSE}{MESSAGE_CLOSE}"
)
assert result.reasoning_text == "deep thought"
assert result.normal_text == _SKIPPED_THINK_REPLY
@pytest.mark.parametrize("suffix", ["", THINK_CLOSE])
def test_non_stream_skipped_think_before_tools_keeps_reply_as_content(
suffix: str,
) -> None:
"""A later think-close cannot reclassify a response closed before tools."""
detector = KimiK3Detector(force_reasoning=True)
text = f"{_SKIPPED_THINK_REPLY}{RESPONSE_CLOSE}{_TOOLS_CHANNEL}{suffix}"
result = detector.detect_and_parse(text)
assert result.reasoning_text == ""
assert result.normal_text == f"{_SKIPPED_THINK_REPLY}{_TOOLS_CHANNEL}"
@pytest.mark.parametrize("stream_reasoning", [False, True])
@pytest.mark.parametrize("chunk_size", [1, None])
def test_streaming_skipped_think_channel_is_chunk_independent(
stream_reasoning: bool, chunk_size: int | None
) -> None:
"""A skipped-think reply stays content across streaming boundaries."""
detector = KimiK3Detector(force_reasoning=True, stream_reasoning=stream_reasoning)
text = f"{_SKIPPED_THINK_REPLY}{RESPONSE_CLOSE}{MESSAGE_CLOSE}"
chunks = [text] if chunk_size is None else _chunks(text, chunk_size)
reasoning, content = _stream(detector, chunks)
assert content == _SKIPPED_THINK_REPLY
assert "<|" not in reasoning and "<|" not in content
if stream_reasoning:
assert _SKIPPED_THINK_REPLY.startswith(reasoning)
# Anything after the switch is plain content, not reasoning.
tail = detector.parse_streaming_increment("more")
assert tail.normal_text == "more" and not tail.reasoning_text
else:
assert reasoning == ""
@pytest.mark.parametrize("stream_reasoning", [False, True])
@pytest.mark.parametrize("chunk_size", [1, None])
def test_streaming_skipped_think_before_tools_is_chunk_independent(
stream_reasoning: bool, chunk_size: int | None
) -> None:
"""Response/tool routing is invariant to chunking and reasoning buffering."""
detector = KimiK3Detector(force_reasoning=True, stream_reasoning=stream_reasoning)
text = f"{_SKIPPED_THINK_REPLY}{RESPONSE_CLOSE}{_TOOLS_CHANNEL}"
chunks = [text] if chunk_size is None else _chunks(text, chunk_size)
reasoning, content = _stream(detector, chunks)
assert content == f"{_SKIPPED_THINK_REPLY}{_TOOLS_CHANNEL}"
assert RESPONSE_CLOSE not in reasoning
if stream_reasoning:
assert _SKIPPED_THINK_REPLY.startswith(reasoning)
else:
assert reasoning == ""
@pytest.mark.parametrize("stream_reasoning", [False, True])
@pytest.mark.parametrize("split_after_response_close", [False, True])
def test_streaming_skipped_think_ignores_delayed_think_close(
stream_reasoning: bool, split_after_response_close: bool
) -> None:
"""A delayed think-close is consumed without changing response routing."""
detector = KimiK3Detector(force_reasoning=True, stream_reasoning=stream_reasoning)
response = f"{_SKIPPED_THINK_REPLY}{RESPONSE_CLOSE}"
chunks = (
[response, THINK_CLOSE]
if split_after_response_close
else [response + THINK_CLOSE]
)
reasoning, content = _stream(detector, chunks)
assert reasoning == ""
assert content == _SKIPPED_THINK_REPLY
def test_reasoning_parser_registration() -> None:
assert isinstance(ReasoningParser("kimi_k3").detector, KimiK3Detector)
@@ -200,7 +297,7 @@ def _stream_with_finish(detector: KimiK3Detector, chunks: list[str]) -> tuple[st
("text", "reasoning", "content"),
[
(
f"bare answer{RESPONSE_CLOSE}{MESSAGE_CLOSE}",
f"bare answer{MESSAGE_CLOSE}",
"",
"bare answer",
),
@@ -223,12 +320,10 @@ def test_fnc_non_stream_skipped_think_vs_truncated_reasoning(
@pytest.mark.parametrize("chunk_size", [1, 5, 13])
def test_fnc_streaming_skipped_think_answer(chunk_size: int) -> None:
def test_fnc_streaming_message_close_recovers_at_finish(chunk_size: int) -> None:
detector = KimiK3Detector(force_reasoning=True, force_nonempty_content=True)
text = f"bare answer{RESPONSE_CLOSE}{MESSAGE_CLOSE}"
text = f"bare answer{MESSAGE_CLOSE}"
reasoning, content = _stream_with_finish(detector, _chunks(text, chunk_size))
# Streamed as reasoning in real time; finish() re-emits the cleaned
# payload as content once the channel close proves skipped-think.
assert reasoning == text
assert content == "bare answer"
@@ -289,7 +384,7 @@ def test_stream_reasoning_off_truncation_flushes_reasoning(
assert content == ""
def test_fnc_stream_reasoning_off_skipped_think_reemits_content() -> None:
def test_fnc_stream_reasoning_off_skipped_think_recovers_content() -> None:
detector = KimiK3Detector(
force_reasoning=True, stream_reasoning=False, force_nonempty_content=True
)