summaryrefslogtreecommitdiff
path: root/llm
diff options
context:
space:
mode:
Diffstat (limited to 'llm')
-rw-r--r--llm/tokenbuffer.go16
-rw-r--r--llm/tokenbuffer_test.go44
2 files changed, 32 insertions, 28 deletions
diff --git a/llm/tokenbuffer.go b/llm/tokenbuffer.go
index 19d7db9..76f6746 100644
--- a/llm/tokenbuffer.go
+++ b/llm/tokenbuffer.go
@@ -8,7 +8,7 @@ type TokenBuffer struct {
stride int
buffer []int64
document int
- seen int
+ position int
includeTail bool
}
@@ -23,7 +23,7 @@ func NewTokenBuffer(tokenizer Tokenizer, window, stride int) *TokenBuffer {
stride: stride,
buffer: make([]int64, 0, 2*window),
document: -1,
- seen: 0,
+ position: 0,
includeTail: true,
}
}
@@ -36,8 +36,8 @@ func (tb *TokenBuffer) SetIncludeTail(includeTail bool) {
tb.includeTail = includeTail
}
-func (tb *TokenBuffer) Seen() int {
- return tb.seen
+func (tb *TokenBuffer) Position() int {
+ return tb.position
}
func (tb *TokenBuffer) Push(document int, text string) iter.Seq2[[]int64, int] {
@@ -65,11 +65,11 @@ func (tb *TokenBuffer) Push(document int, text string) iter.Seq2[[]int64, int] {
copy(w, tb.buffer[:tb.window])
- if !yield(w, tb.seen) {
+ if !yield(w, min(tb.Position(), tb.window-tb.stride)) {
return
}
- tb.seen += tb.stride
+ tb.position += tb.stride
tb.buffer = append(tb.buffer[:0], tb.buffer[tb.stride:]...)
}
@@ -77,10 +77,10 @@ func (tb *TokenBuffer) Push(document int, text string) iter.Seq2[[]int64, int] {
}
func (tb *TokenBuffer) Tail() ([]int64, int) {
- seen := tb.seen
+ seen := min(tb.Position(), min(len(tb.buffer), tb.window-tb.stride))
tb.document = -1
- tb.seen = 0
+ tb.position = 0
if !tb.includeTail || len(tb.buffer) == 0 {
tb.buffer = tb.buffer[:0]
diff --git a/llm/tokenbuffer_test.go b/llm/tokenbuffer_test.go
index b03f267..072aa4f 100644
--- a/llm/tokenbuffer_test.go
+++ b/llm/tokenbuffer_test.go
@@ -131,31 +131,33 @@ func TestTokenBuffer_Tail(t *testing.T) {
}
}
-func TestTokenBuffer_Seen(t *testing.T) {
+func TestTokenBuffer_Position(t *testing.T) {
tb := NewTokenBuffer(byteTokenizer{}, 4, 2)
tb.SetIncludeTail(false)
- var seen []int
+ var positions []int
- for _, s := range tb.Push(0, "abcdefgh") {
- seen = append(seen, s)
+ for range tb.Push(0, "abcdefgh") {
+ positions = append(positions, tb.Position())
}
expected := []int{0, 2, 4}
- if len(seen) != len(expected) {
- t.Fatalf("expected %d seen values but got %d", len(expected), len(seen))
+ if len(positions) != len(expected) {
+ t.Fatalf("expected %d positions but got %d", len(expected), len(positions))
}
- 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])
+ for i := range positions {
+ if positions[i] != expected[i] {
+ t.Errorf("window %d: expected positions=%d but got positions=%d", i, expected[i], positions[i])
}
}
- if _, s := tb.Tail(); s != 6 {
- t.Errorf("expected tail seen=6 but got %d", s)
+ _, seen := tb.Tail()
+
+ if tailPosition := positions[len(positions)-1] + seen; tailPosition != 6 {
+ t.Errorf("expected tail positions=6 but got %d", tailPosition)
}
}
@@ -191,28 +193,30 @@ func TestTokenBuffer_DocumentBoundaryYieldsTail(t *testing.T) {
}
}
-func TestTokenBuffer_TailSeen(t *testing.T) {
+func TestTokenBuffer_TailPosition(t *testing.T) {
tb := NewTokenBuffer(byteTokenizer{}, 4, 2)
tb.SetIncludeTail(true)
- var lastSeen int
+ var lastPosition int
- for _, s := range tb.Push(0, "abcdefgh") {
- lastSeen = s
+ for range tb.Push(0, "abcdefgh") {
+ lastPosition = tb.Position()
}
- if lastSeen != 4 {
- t.Errorf("expected last seen=4 but got %d", lastSeen)
+ if lastPosition != 4 {
+ t.Errorf("expected last position=4 but got %d", lastPosition)
}
- tail, tailSeen := tb.Tail()
+ tail, seen := tb.Tail()
+
+ tailPosition := lastPosition + seen
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)
+ if tailPosition != 6 {
+ t.Errorf("expected tail position=6 but got %d", tailPosition)
}
}