Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
105 changes: 86 additions & 19 deletions scripts/chat_grpc.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,7 +50,7 @@
import argparse
import sys
import time
from typing import List, Optional
from typing import Iterable, List, Optional


_HELP = """
Expand All @@ -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),
"<end_of_turn>",
"<turn|>",
"<|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"
Expand Down Expand Up @@ -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
)
Expand Down Expand Up @@ -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


Expand Down Expand Up @@ -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)

Expand Down
49 changes: 48 additions & 1 deletion tests/scripts/test_chat_grpc.py
Original file line number Diff line number Diff line change
@@ -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:
Expand Down Expand Up @@ -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 = "<eos>"
eot_token = "<turn|>"
unk_token_id = 3

def convert_tokens_to_ids(self, token):
return {
"<eos>": 1,
"<turn|>": 106,
"<end_of_turn>": 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),
Expand Down Expand Up @@ -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
Loading