diff options
Diffstat (limited to 'llm')
| -rw-r--r-- | llm/tokenbuffer.go | 16 | ||||
| -rw-r--r-- | llm/tokenbuffer_test.go | 44 |
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) } } |
