[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:
@@ -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
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user