From f4860bfa53004e6e18968f0650207d5317eebdeb Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Thu, 2 Apr 2026 22:25:54 +0200 Subject: Add token buffer --- llm/tokenbuffer.go | 86 +++++++++++++++++++++++++ llm/tokenbuffer_test.go | 164 ++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 250 insertions(+) create mode 100644 llm/tokenbuffer.go create mode 100644 llm/tokenbuffer_test.go diff --git a/llm/tokenbuffer.go b/llm/tokenbuffer.go new file mode 100644 index 0000000..4df0af3 --- /dev/null +++ b/llm/tokenbuffer.go @@ -0,0 +1,86 @@ +package llm + +import "iter" + +type TokenBuffer struct { + tokenizer Tokenizer + window int + stride int + buffer []int64 + document int + includeTail bool +} + +func NewTokenBuffer(tokenizer Tokenizer, window, stride int) *TokenBuffer { + if stride > window { + panic("stride exceeds window") + } + + return &TokenBuffer{ + tokenizer: tokenizer, + window: window, + stride: stride, + buffer: make([]int64, 0, 2*window), + document: -1, + includeTail: true, + } +} + +func (tb *TokenBuffer) IncludeTail() bool { + return tb.includeTail +} + +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) { + if tb.document != -1 && document != tb.document { + if tb.includeTail && len(tb.buffer) > 0 { + if !yield(tb.buffer) { + return + } + } + + tb.buffer = nil + } + + tb.document = document + + if text == "" { + return + } + + ids := toInt64(tb.tokenizer.Tokenize(text)) + + tb.buffer = append(tb.buffer, ids...) + + for len(tb.buffer) >= tb.window { + w := make([]int64, tb.window) + + copy(w, tb.buffer[:tb.window]) + + if !yield(w) { + return + } + + tb.buffer = append(tb.buffer[:0], tb.buffer[tb.stride:]...) + } + } +} + +func (tb *TokenBuffer) Tail() []int64 { + if !tb.includeTail || len(tb.buffer) == 0 { + return nil + } + + w := make([]int64, len(tb.buffer)) + + copy(w, tb.buffer) + + tb.buffer = nil + tb.document = -1 + + return w +} diff --git a/llm/tokenbuffer_test.go b/llm/tokenbuffer_test.go new file mode 100644 index 0000000..860c5ab --- /dev/null +++ b/llm/tokenbuffer_test.go @@ -0,0 +1,164 @@ +package llm + +import ( + "fmt" + "slices" + "testing" +) + +type byteTokenizer struct{} + +func (bt byteTokenizer) Tokenize(text string) []int { + r := make([]int, len(text)) + + for i := 0; i < len(text); i++ { + r[i] = int(text[i]) + } + + return r +} + +func TestTokenBuffer_Push(t *testing.T) { + type gold struct { + window int + stride int + text string + expected [][]int64 + } + + tests := []gold{ + { + window: 5, stride: 5, + text: "abcdefghij", + expected: [][]int64{{97, 98, 99, 100, 101}, {102, 103, 104, 105, 106}}, + }, + { + window: 4, stride: 2, + text: "abcdef", + expected: [][]int64{{97, 98, 99, 100}, {99, 100, 101, 102}}, + }, + { + window: 10, stride: 5, + text: "abcde", + expected: nil, + }, + } + + for _, tt := range tests { + t.Run( + fmt.Sprintf("window%d_stride%d", tt.window, tt.stride), + + func(t *testing.T) { + tb := NewTokenBuffer(byteTokenizer{}, tt.window, tt.stride) + + tb.SetIncludeTail(false) + + var got [][]int64 + + for w := range tb.Push(0, tt.text) { + got = append(got, w) + } + + if len(got) != len(tt.expected) { + t.Fatalf("expected %d windows but got %d", len(tt.expected), len(got)) + } + + for i := range got { + if !slices.Equal(got[i], tt.expected[i]) { + t.Errorf("window %d: expected %v but got %v", i, tt.expected[i], got[i]) + } + } + }, + ) + } +} + +func TestTokenBuffer_PushAccumulates(t *testing.T) { + tb := NewTokenBuffer(byteTokenizer{}, 4, 4) + + var got [][]int64 + + for w := range tb.Push(0, "ab") { + got = append(got, w) + } + + if len(got) != 0 { + t.Fatalf("expected 0 windows but got %d", len(got)) + } + + for w := range tb.Push(0, "cd") { + got = append(got, w) + } + + expected := [][]int64{{97, 98, 99, 100}} + + if len(got) != len(expected) { + t.Fatalf("expected %d windows but got %d", len(expected), len(got)) + } + + for i := range got { + if !slices.Equal(got[i], expected[i]) { + t.Errorf("window %d: expected %v but got %v", i, expected[i], got[i]) + } + } +} + +func TestTokenBuffer_Tail(t *testing.T) { + tb := NewTokenBuffer(byteTokenizer{}, 5, 5) + + tb.SetIncludeTail(true) + + var got [][]int64 + + for w := range tb.Push(0, "abcdef") { + got = append(got, w) + } + + if tail := tb.Tail(); tail != nil { + got = append(got, tail) + } + + expected := [][]int64{{97, 98, 99, 100, 101}, {102}} + + if len(got) != len(expected) { + t.Fatalf("expected %d windows but got %d", len(expected), len(got)) + } + + for i := range got { + if !slices.Equal(got[i], expected[i]) { + t.Errorf("window %d: expected %v but got %v", i, expected[i], got[i]) + } + } +} + +func TestTokenBuffer_DocumentBoundaryYieldsTail(t *testing.T) { + tb := NewTokenBuffer(byteTokenizer{}, 10, 10) + + tb.SetIncludeTail(true) + + var a [][]int64 + + for w := range tb.Push(0, "abcde") { + a = append(a, w) + } + + if len(a) != 0 { + t.Fatalf("expected 0 windows but got %d", len(a)) + } + + var b [][]int64 + + for w := range tb.Push(2, "fghij") { + b = append(b, w) + } + + if len(b) != 1 { + t.Fatalf("expected 1 window but got %d", len(b)) + } + + expected := toInt64(byteTokenizer{}.Tokenize("abcde")) + + if !slices.Equal(b[0], expected) { + t.Errorf("expected %v but got %v", expected, b[0]) + } +} -- cgit v1.2.3