summaryrefslogtreecommitdiff
path: root/llm
diff options
context:
space:
mode:
Diffstat (limited to 'llm')
-rw-r--r--llm/tokenbuffer.go26
-rw-r--r--llm/tokenbuffer_test.go58
2 files changed, 75 insertions, 9 deletions
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)
+ }
+}