From 253bec78ca09d29e98caf401fdb3e027144ff918 Mon Sep 17 00:00:00 2001 From: fluffy314 Date: Sun, 12 Jul 2026 18:13:30 +0800 Subject: [PATCH] fix(chat): honor model end-of-turn tokens Register Gemma's end-of-turn token alongside generic EOS so natural answers terminate at the model boundary, with a strict repetition guard for malformed streams. Co-authored-by: Cursor --- scripts/chat_grpc.py | 105 ++++++++++++++++++++++++++------ tests/scripts/test_chat_grpc.py | 49 ++++++++++++++- 2 files changed, 134 insertions(+), 20 deletions(-) diff --git a/scripts/chat_grpc.py b/scripts/chat_grpc.py index d31ce6c1..529bd506 100755 --- a/scripts/chat_grpc.py +++ b/scripts/chat_grpc.py @@ -50,7 +50,7 @@ import argparse import sys import time -from typing import List, Optional +from typing import Iterable, List, Optional _HELP = """ @@ -62,6 +62,55 @@ """.strip() +def _resolve_eos_token_ids(tokenizer) -> List[int]: + """Return every tokenizer token that naturally ends a model turn.""" + resolved = set() + + def _add(values) -> None: + if values is None: + return + if isinstance(values, int): + resolved.add(int(values)) + return + if isinstance(values, Iterable) and not isinstance(values, (str, bytes)): + resolved.update(int(value) for value in values) + + _add(getattr(tokenizer, "eos_token_ids", None)) + _add(getattr(tokenizer, "eos_token_id", None)) + _add(getattr(tokenizer, "eot_token_id", None)) + + unk = getattr(tokenizer, "unk_token_id", None) + markers = { + getattr(tokenizer, "eos_token", None), + getattr(tokenizer, "eot_token", None), + "", + "", + "<|im_end|>", + } + for marker in markers: + if not marker: + continue + try: + token_id = tokenizer.convert_tokens_to_ids(marker) + except (AttributeError, TypeError, ValueError): + continue + if isinstance(token_id, int) and token_id >= 0 and token_id != unk: + resolved.add(token_id) + return sorted(resolved) + + +def _is_degenerate_loop(text: str, unit: int = 16) -> bool: + """Detect only a strict three-way consecutive repeated suffix.""" + if len(text) < unit * 3: + return False + first = text[-unit:] + return ( + bool(first.strip()) + and first == text[-2 * unit:-unit] + and first == text[-3 * unit:-2 * unit] + ) + + def _print_banner(address: str, tokenizer_id: str) -> None: print( f"Kakeya v0.3 chat — {address} ({tokenizer_id})\n" @@ -118,22 +167,34 @@ def _generate_and_print( if remaining <= 0: stop_reason = "client_safety_limit" break - for token_id in session.generate(max_tokens=remaining): - n += 1 - accumulated.append(token_id) - # Decode incrementally — tokenizer.decode on the running - # buffer gives the right text including BPE merges that - # span multiple tokens. - text_so_far = tokenizer.decode( - accumulated, skip_special_tokens=True, - ) - if hasattr(_generate_and_print, "_last_text"): - last = _generate_and_print._last_text - else: - last = "" - new_text = text_so_far[len(last):] - print(new_text, end="", flush=True) - _generate_and_print._last_text = text_so_far + stream = session.generate(max_tokens=remaining) + try: + for token_id in stream: + n += 1 + accumulated.append(token_id) + # Decode incrementally — tokenizer.decode on the running + # buffer gives the right text including BPE merges that + # span multiple tokens. + text_so_far = tokenizer.decode( + accumulated, skip_special_tokens=True, + ) + if hasattr(_generate_and_print, "_last_text"): + last = _generate_and_print._last_text + else: + last = "" + new_text = text_so_far[len(last):] + print(new_text, end="", flush=True) + _generate_and_print._last_text = text_so_far + if _is_degenerate_loop(text_so_far): + stop_reason = "degenerate_loop" + break + finally: + if stop_reason == "degenerate_loop": + close = getattr(stream, "close", None) + if close is not None: + close() + if stop_reason == "degenerate_loop": + break server_elapsed += float( getattr(session, "last_total_duration_seconds", 0.0) or 0.0 ) @@ -173,6 +234,13 @@ def _generate_and_print( file=sys.stderr, flush=True, ) + elif stop_reason == "degenerate_loop": + print( + "[generation stopped after a strict consecutive repetition; " + "use /reset before continuing]", + file=sys.stderr, + flush=True, + ) return n @@ -222,8 +290,7 @@ def main() -> int: print(f"[chat] loading tokenizer {args.tokenizer_id} ...", file=sys.stderr, flush=True) tokenizer = AutoTokenizer.from_pretrained(args.tokenizer_id) - eos = tokenizer.eos_token_id - eos_ids: List[int] = [int(eos)] if eos is not None else [] + eos_ids = _resolve_eos_token_ids(tokenizer) _print_banner(args.address, args.tokenizer_id) diff --git a/tests/scripts/test_chat_grpc.py b/tests/scripts/test_chat_grpc.py index ec7001a5..0ca041f9 100644 --- a/tests/scripts/test_chat_grpc.py +++ b/tests/scripts/test_chat_grpc.py @@ -1,6 +1,10 @@ from __future__ import annotations -from scripts.chat_grpc import _generate_and_print +from scripts.chat_grpc import ( + _generate_and_print, + _is_degenerate_loop, + _resolve_eos_token_ids, +) class Tokenizer: @@ -29,6 +33,32 @@ def generate(self, *, max_tokens): self.last_total_duration_seconds = seconds +class GemmaTokenizer: + eos_token_ids = 1 + eos_token_id = 1 + eot_token_id = 106 + eos_token = "" + eot_token = "" + unk_token_id = 3 + + def convert_tokens_to_ids(self, token): + return { + "": 1, + "": 106, + "": 3, + "<|im_end|>": 3, + }[token] + + +def test_resolves_gemma_end_of_turn_as_natural_eos(): + assert _resolve_eos_token_ids(GemmaTokenizer()) == [1, 106] + + +def test_degenerate_loop_guard_requires_three_consecutive_blocks(): + assert not _is_degenerate_loop("normal text " * 2) + assert _is_degenerate_loop("0123456789abcdef" * 3) + + def test_continues_max_token_chunks_until_eos(capsys): session = Session([ ([11, 12], 1, 1.0), @@ -69,3 +99,20 @@ def test_no_progress_breaks_continuation_loop(capsys): ]) assert _generate_and_print(session, Tokenizer(), [9], max_tokens=2) == 0 assert "stop=no_progress" in capsys.readouterr().err + + +def test_stops_and_closes_stream_on_strict_degenerate_loop(capsys): + class RepeatingTokenizer: + def decode(self, token_ids, *, skip_special_tokens=True): + assert skip_special_tokens + return "x" * len(token_ids) + + session = Session([ + ([1] * 64, 1, 1.0), + ]) + assert _generate_and_print( + session, RepeatingTokenizer(), [9], max_tokens=64, + ) == 48 + output = capsys.readouterr() + assert "stop=degenerate_loop" in output.err + assert "use /reset" in output.err