summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-02 22:25:54 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-02 22:25:54 +0200
commitf4860bfa53004e6e18968f0650207d5317eebdeb (patch)
tree53a61f0d75d78c64a712b8a31e4bfeef4df7e289
parentfa4e25b4992c25cdec41bdccb8af909b93fc9c39 (diff)
Add token buffer
-rw-r--r--llm/tokenbuffer.go86
-rw-r--r--llm/tokenbuffer_test.go164
2 files changed, 250 insertions, 0 deletions
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])
+ }
+}