From ef2a19b58792dbd274a03f4696719652e56e4360 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Thu, 2 Apr 2026 23:02:09 +0200 Subject: Track already seen tokens --- llm/tokenbuffer.go | 26 ++++++++++++++++------ llm/tokenbuffer_test.go | 58 +++++++++++++++++++++++++++++++++++++++++++++++-- 2 files changed, 75 insertions(+), 9 deletions(-) (limited to 'llm') diff --git a/llm/tokenbuffer.go b/llm/tokenbuffer.go index 4df0af3..f83858d 100644 --- a/llm/tokenbuffer.go +++ b/llm/tokenbuffer.go @@ -8,6 +8,7 @@ type TokenBuffer struct { stride int buffer []int64 document int + seen int includeTail bool } @@ -22,6 +23,7 @@ func NewTokenBuffer(tokenizer Tokenizer, window, stride int) *TokenBuffer { stride: stride, buffer: make([]int64, 0, 2*window), document: -1, + seen: 0, includeTail: true, } } @@ -34,16 +36,21 @@ func (tb *TokenBuffer) SetIncludeTail(includeTail bool) { tb.includeTail = includeTail } -func (tb *TokenBuffer) Push(document int, text string) iter.Seq[[]int64] { - return func(yield func([]int64) bool) { +func (tb *TokenBuffer) Seen() int { + return tb.seen +} + +func (tb *TokenBuffer) Push(document int, text string) iter.Seq2[[]int64, int] { + return func(yield func([]int64, int) bool) { if tb.document != -1 && document != tb.document { if tb.includeTail && len(tb.buffer) > 0 { - if !yield(tb.buffer) { + if !yield(tb.buffer, tb.seen) { return } } tb.buffer = nil + tb.seen = 0 } tb.document = document @@ -61,26 +68,31 @@ func (tb *TokenBuffer) Push(document int, text string) iter.Seq[[]int64] { copy(w, tb.buffer[:tb.window]) - if !yield(w) { + if !yield(w, tb.seen) { return } + tb.seen += tb.stride + tb.buffer = append(tb.buffer[:0], tb.buffer[tb.stride:]...) } } } -func (tb *TokenBuffer) Tail() []int64 { +func (tb *TokenBuffer) Tail() ([]int64, int) { if !tb.includeTail || len(tb.buffer) == 0 { - return nil + return nil, tb.seen } w := make([]int64, len(tb.buffer)) copy(w, tb.buffer) + seen := tb.seen + tb.buffer = nil tb.document = -1 + tb.seen = 0 - return w + return w, seen } diff --git a/llm/tokenbuffer_test.go b/llm/tokenbuffer_test.go index 860c5ab..b03f267 100644 --- a/llm/tokenbuffer_test.go +++ b/llm/tokenbuffer_test.go @@ -114,7 +114,7 @@ func TestTokenBuffer_Tail(t *testing.T) { got = append(got, w) } - if tail := tb.Tail(); tail != nil { + if tail, _ := tb.Tail(); tail != nil { got = append(got, tail) } @@ -131,6 +131,34 @@ func TestTokenBuffer_Tail(t *testing.T) { } } +func TestTokenBuffer_Seen(t *testing.T) { + tb := NewTokenBuffer(byteTokenizer{}, 4, 2) + + tb.SetIncludeTail(false) + + var seen []int + + for _, s := range tb.Push(0, "abcdefgh") { + seen = append(seen, s) + } + + expected := []int{0, 2, 4} + + if len(seen) != len(expected) { + t.Fatalf("expected %d seen values but got %d", len(expected), len(seen)) + } + + for i := range seen { + if seen[i] != expected[i] { + t.Errorf("window %d: expected seen=%d but got seen=%d", i, expected[i], seen[i]) + } + } + + if _, s := tb.Tail(); s != 6 { + t.Errorf("expected tail seen=6 but got %d", s) + } +} + func TestTokenBuffer_DocumentBoundaryYieldsTail(t *testing.T) { tb := NewTokenBuffer(byteTokenizer{}, 10, 10) @@ -148,7 +176,7 @@ func TestTokenBuffer_DocumentBoundaryYieldsTail(t *testing.T) { var b [][]int64 - for w := range tb.Push(2, "fghij") { + for w, _ := range tb.Push(2, "fghij") { b = append(b, w) } @@ -162,3 +190,29 @@ func TestTokenBuffer_DocumentBoundaryYieldsTail(t *testing.T) { t.Errorf("expected %v but got %v", expected, b[0]) } } + +func TestTokenBuffer_TailSeen(t *testing.T) { + tb := NewTokenBuffer(byteTokenizer{}, 4, 2) + + tb.SetIncludeTail(true) + + var lastSeen int + + for _, s := range tb.Push(0, "abcdefgh") { + lastSeen = s + } + + if lastSeen != 4 { + t.Errorf("expected last seen=4 but got %d", lastSeen) + } + + tail, tailSeen := tb.Tail() + + if !slices.Equal(tail, toInt64(byteTokenizer{}.Tokenize("gh"))) { + t.Errorf("unexpected tail: %v", tail) + } + + if tailSeen != 6 { + t.Errorf("expected tail seen=6 but got %d", tailSeen) + } +} -- cgit v1.3.1