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
14 changes: 14 additions & 0 deletions src/imf/export.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,13 @@
EOS_ID = 1


def decode_tokens(tokens: list[int]) -> str:
"""Inverse of encode_bytes (byte-3 offsets, EOS-terminated)."""
return bytes((t - BYTE_OFFSET) % 256 for t in tokens if t >= BYTE_OFFSET).decode(
"utf-8", "replace"
)


def encode_bytes(text: str) -> list[int]:
"""Canonical byte-level tokenization: byte ids + trailing EOS."""
return [b + BYTE_OFFSET for b in text.encode("utf-8")] + [EOS_ID]
Expand Down Expand Up @@ -418,6 +425,13 @@ def onnx_greedy_kv(encoder_sess, kv_sess, text: str, max_len: int = 256) -> list
break
else:
joined = ",".join(str(t) for t in generated)
# Decoded-text guard: loops with varying punctuation never repeat
# a verbatim token window — catch the phrase itself echoing.
if len(generated) % 8 == 0:
text = decode_tokens(generated)
suffix = text[-16:]
if len(suffix) == 16 and text.count(suffix) >= 3:
break
return generated


Expand Down
27 changes: 27 additions & 0 deletions tests/test_decode_guard.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,3 +74,30 @@ def run(self, _, feeds):

out = onnx_greedy_kv(_FakeEncoder(), _StopKV(), "x", max_len=256)
assert out == [20, 21, 22, 23, 24]


def test_varying_separator_loop_is_cut():
"""Phrase + rotating punctuation never repeats a verbatim token
window — the live int8 failure mode. The decoded-text guard cuts it."""

class _RotateKV(_FakeKV):
seps = ['"', " ", "\n", ":"]

def run(self, _, feeds):
import numpy as np

logits = np.full((1, 1, 260), -1e9)
if self.step == 0:
logits[0, -1, 100] = 1e9 # phrase token
else:
mod = self.step % 4
if mod == 0:
logits[0, -1, 100] = 1e9 # phrase again
else:
# rotating separator tokens 200..203
logits[0, -1, 200 + ((self.step // 4) % 4)] = 1e9
self.step += 1
return [logits, np.zeros((1,)), np.zeros((1,))]

out = onnx_greedy_kv(_FakeEncoder(), _RotateKV(), "x", max_len=4096)
assert len(out) < 300
Loading