summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--llm/run.go39
-rw-r--r--llm/tokenbuffer.go55
-rw-r--r--llm/tokenbuffer_test.go30
3 files changed, 79 insertions, 45 deletions
diff --git a/llm/run.go b/llm/run.go
index 7d29e21..abb4c0d 100644
--- a/llm/run.go
+++ b/llm/run.go
@@ -56,35 +56,38 @@ func (e *Evaluator[R]) Run(title string, data dataset.Reader, window, stride int
doc := 0
pos := 0
- for n, d := range data.Texts() {
+ for w := range tb.Stream(data.Texts()) {
+ n := w.Document
+
if n != doc {
pos = 0
doc = n
}
- for w, s := range tb.Push(n, d) {
- if s == 0 {
- s = 1 // first token as context
- }
+ tokens := w.Tokens
+ seen := w.Seen
- b.AddJob(Job{
- Document: doc,
- Position: pos,
- Tokens: w,
- Seen: s,
- })
+ if seen == 0 {
+ seen = 1 // first token as context
+ }
- pos++
+ b.AddJob(Job{
+ Document: doc,
+ Position: pos,
+ Tokens: tokens,
+ Seen: seen,
+ })
- if b.Size() == e.batchSize {
- e.jobs <- *b
+ pos++
- b = newBatch(e.batchSize)
+ if b.Size() == e.batchSize {
+ e.jobs <- *b
- e.scheduled.Add(int64(e.batchSize))
+ b = newBatch(e.batchSize)
- pb.SetTotal(int(e.scheduled.Load()))
- }
+ e.scheduled.Add(int64(e.batchSize))
+
+ pb.SetTotal(int(e.scheduled.Load()))
}
}
diff --git a/llm/tokenbuffer.go b/llm/tokenbuffer.go
index 76f6746..d19de4a 100644
--- a/llm/tokenbuffer.go
+++ b/llm/tokenbuffer.go
@@ -12,6 +12,12 @@ type TokenBuffer struct {
includeTail bool
}
+type TokenWindow struct {
+ Document int
+ Tokens []int64
+ Seen int
+}
+
func NewTokenBuffer(tokenizer Tokenizer, window, stride int) *TokenBuffer {
if stride > window {
panic("stride exceeds window")
@@ -40,12 +46,12 @@ func (tb *TokenBuffer) Position() int {
return tb.position
}
-func (tb *TokenBuffer) Push(document int, text string) iter.Seq2[[]int64, int] {
- return func(yield func([]int64, int) bool) {
+func (tb *TokenBuffer) Push(document int, text string) iter.Seq[TokenWindow] {
+ return func(yield func(window TokenWindow) bool) {
if tb.document != -1 && document != tb.document {
- tail, seen := tb.Tail()
+ w, ok := tb.Tail()
- if len(tail) > 0 && !yield(tail, seen) {
+ if ok && !yield(w) {
return
}
}
@@ -61,11 +67,15 @@ func (tb *TokenBuffer) Push(document int, text string) iter.Seq2[[]int64, int] {
tb.buffer = append(tb.buffer, ids...)
for len(tb.buffer) >= tb.window {
- w := make([]int64, tb.window)
+ w := TokenWindow{
+ Document: tb.document,
+ Tokens: make([]int64, tb.window),
+ Seen: min(tb.Position(), tb.window-tb.stride),
+ }
- copy(w, tb.buffer[:tb.window])
+ copy(w.Tokens, tb.buffer[:tb.window])
- if !yield(w, min(tb.Position(), tb.window-tb.stride)) {
+ if !yield(w) {
return
}
@@ -76,23 +86,44 @@ func (tb *TokenBuffer) Push(document int, text string) iter.Seq2[[]int64, int] {
}
}
-func (tb *TokenBuffer) Tail() ([]int64, int) {
+func (tb *TokenBuffer) Tail() (TokenWindow, bool) {
seen := min(tb.Position(), min(len(tb.buffer), tb.window-tb.stride))
+ w := TokenWindow{
+ Document: tb.document,
+ Seen: seen,
+ }
+
tb.document = -1
tb.position = 0
if !tb.includeTail || len(tb.buffer) == 0 {
tb.buffer = tb.buffer[:0]
- return nil, seen
+ return w, false
}
- w := make([]int64, len(tb.buffer))
+ w.Tokens = make([]int64, len(tb.buffer))
- copy(w, tb.buffer)
+ copy(w.Tokens, tb.buffer)
tb.buffer = tb.buffer[:0]
- return w, seen
+ return w, true
+}
+
+func (tb *TokenBuffer) Stream(docs iter.Seq2[int, string]) iter.Seq[TokenWindow] {
+ return func(yield func(window TokenWindow) bool) {
+ for n, d := range docs {
+ for w := range tb.Push(n, d) {
+ if !yield(w) {
+ return
+ }
+ }
+
+ if w, ok := tb.Tail(); ok && !yield(w) {
+ return
+ }
+ }
+ }
}
diff --git a/llm/tokenbuffer_test.go b/llm/tokenbuffer_test.go
index 072aa4f..789caa6 100644
--- a/llm/tokenbuffer_test.go
+++ b/llm/tokenbuffer_test.go
@@ -56,7 +56,7 @@ func TestTokenBuffer_Push(t *testing.T) {
var got [][]int64
for w := range tb.Push(0, tt.text) {
- got = append(got, w)
+ got = append(got, w.Tokens)
}
if len(got) != len(tt.expected) {
@@ -79,7 +79,7 @@ func TestTokenBuffer_PushAccumulates(t *testing.T) {
var got [][]int64
for w := range tb.Push(0, "ab") {
- got = append(got, w)
+ got = append(got, w.Tokens)
}
if len(got) != 0 {
@@ -87,7 +87,7 @@ func TestTokenBuffer_PushAccumulates(t *testing.T) {
}
for w := range tb.Push(0, "cd") {
- got = append(got, w)
+ got = append(got, w.Tokens)
}
expected := [][]int64{{97, 98, 99, 100}}
@@ -111,11 +111,11 @@ func TestTokenBuffer_Tail(t *testing.T) {
var got [][]int64
for w := range tb.Push(0, "abcdef") {
- got = append(got, w)
+ got = append(got, w.Tokens)
}
- if tail, _ := tb.Tail(); tail != nil {
- got = append(got, tail)
+ if w, ok := tb.Tail(); ok {
+ got = append(got, w.Tokens)
}
expected := [][]int64{{97, 98, 99, 100, 101}, {102}}
@@ -154,9 +154,9 @@ func TestTokenBuffer_Position(t *testing.T) {
}
}
- _, seen := tb.Tail()
+ tail, _ := tb.Tail()
- if tailPosition := positions[len(positions)-1] + seen; tailPosition != 6 {
+ if tailPosition := positions[len(positions)-1] + tail.Seen; tailPosition != 6 {
t.Errorf("expected tail positions=6 but got %d", tailPosition)
}
}
@@ -169,7 +169,7 @@ func TestTokenBuffer_DocumentBoundaryYieldsTail(t *testing.T) {
var a [][]int64
for w := range tb.Push(0, "abcde") {
- a = append(a, w)
+ a = append(a, w.Tokens)
}
if len(a) != 0 {
@@ -178,8 +178,8 @@ func TestTokenBuffer_DocumentBoundaryYieldsTail(t *testing.T) {
var b [][]int64
- for w, _ := range tb.Push(2, "fghij") {
- b = append(b, w)
+ for w := range tb.Push(2, "fghij") {
+ b = append(b, w.Tokens)
}
if len(b) != 1 {
@@ -208,12 +208,12 @@ func TestTokenBuffer_TailPosition(t *testing.T) {
t.Errorf("expected last position=4 but got %d", lastPosition)
}
- tail, seen := tb.Tail()
+ tail, _ := tb.Tail()
- tailPosition := lastPosition + seen
+ tailPosition := lastPosition + tail.Seen
- if !slices.Equal(tail, toInt64(byteTokenizer{}.Tokenize("gh"))) {
- t.Errorf("unexpected tail: %v", tail)
+ if !slices.Equal(tail.Tokens, toInt64(byteTokenizer{}.Tokenize("gh"))) {
+ t.Errorf("unexpected tail: %v", tail.Tokens)
}
if tailPosition != 6 {