diff options
| -rw-r--r-- | llm/run.go | 39 | ||||
| -rw-r--r-- | llm/tokenbuffer.go | 55 | ||||
| -rw-r--r-- | llm/tokenbuffer_test.go | 30 |
3 files changed, 79 insertions, 45 deletions
@@ -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 { |
