Skip to content

Commit e2b6747

Browse files
author
Ronald Tse
committed
fix(imf): decoded-text repetition guard for varying-punctuation loops
The token-window guard bounds periodic loops, but live int8 output loops with rotating punctuation (phrase + varied separator) never repeats a verbatim window. Cut when the recent 16 decoded chars echo 3+ times. decode_tokens mirrors the TS runtime (% 256).
1 parent ef8b445 commit e2b6747

2 files changed

Lines changed: 41 additions & 0 deletions

File tree

‎src/imf/export.py‎

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -38,6 +38,13 @@
3838
EOS_ID = 1
3939

4040

41+
def decode_tokens(tokens: list[int]) -> str:
42+
"""Inverse of encode_bytes (byte-3 offsets, EOS-terminated)."""
43+
return bytes((t - BYTE_OFFSET) % 256 for t in tokens if t >= BYTE_OFFSET).decode(
44+
"utf-8", "replace"
45+
)
46+
47+
4148
def encode_bytes(text: str) -> list[int]:
4249
"""Canonical byte-level tokenization: byte ids + trailing EOS."""
4350
return [b + BYTE_OFFSET for b in text.encode("utf-8")] + [EOS_ID]
@@ -418,6 +425,13 @@ def onnx_greedy_kv(encoder_sess, kv_sess, text: str, max_len: int = 256) -> list
418425
break
419426
else:
420427
joined = ",".join(str(t) for t in generated)
428+
# Decoded-text guard: loops with varying punctuation never repeat
429+
# a verbatim token window — catch the phrase itself echoing.
430+
if len(generated) % 8 == 0:
431+
text = decode_tokens(generated)
432+
suffix = text[-16:]
433+
if len(suffix) == 16 and text.count(suffix) >= 3:
434+
break
421435
return generated
422436

423437

‎tests/test_decode_guard.py‎

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -74,3 +74,30 @@ def run(self, _, feeds):
7474

7575
out = onnx_greedy_kv(_FakeEncoder(), _StopKV(), "x", max_len=256)
7676
assert out == [20, 21, 22, 23, 24]
77+
78+
79+
def test_varying_separator_loop_is_cut():
80+
"""Phrase + rotating punctuation never repeats a verbatim token
81+
window — the live int8 failure mode. The decoded-text guard cuts it."""
82+
83+
class _RotateKV(_FakeKV):
84+
seps = ['"', " ", "\n", ":"]
85+
86+
def run(self, _, feeds):
87+
import numpy as np
88+
89+
logits = np.full((1, 1, 260), -1e9)
90+
if self.step == 0:
91+
logits[0, -1, 100] = 1e9 # phrase token
92+
else:
93+
mod = self.step % 4
94+
if mod == 0:
95+
logits[0, -1, 100] = 1e9 # phrase again
96+
else:
97+
# rotating separator tokens 200..203
98+
logits[0, -1, 200 + ((self.step // 4) % 4)] = 1e9
99+
self.step += 1
100+
return [logits, np.zeros((1,)), np.zeros((1,))]
101+
102+
out = onnx_greedy_kv(_FakeEncoder(), _RotateKV(), "x", max_len=4096)
103+
assert len(out) < 300

0 commit comments

Comments
 (0)