summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-03 15:22:15 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-03 15:45:29 +0200
commita61bc9ff890ea50d41a406cd1ffe554341551d2b (patch)
tree3159fa6cb444d88fae7ac7abee92420d551831c7
parentdc0c59b97243645a9fdc040460e51d290085374e (diff)
Use token buffer to generate jobs
-rw-r--r--llm/run.go102
-rw-r--r--llm/run_test.go45
2 files changed, 36 insertions, 111 deletions
diff --git a/llm/run.go b/llm/run.go
index 227a365..1843c94 100644
--- a/llm/run.go
+++ b/llm/run.go
@@ -52,93 +52,63 @@ func (e *Evaluator[R]) Run(title string, data dataset.Reader, window, stride int
return int(e.completed.Load())
})
- s := 0 // skipped
- n := 0 // total
-
- defer func() {
- fmt.Printf("skipped %d of %d texts\n", s, n)
- }()
+ n := 0
+ m := 0
defer pb.Close()
- for d := range data.Texts("text") {
- tokens := toInt64(e.tokenizer.Tokenize(d))
-
- if len(tokens) == 0 {
- n++
- s++
-
- continue
- }
+ tb := NewTokenBuffer(e.tokenizer, window, stride)
- e.schedule(n, tokens, window, stride, e.batchSize)
+ tb.SetIncludeTail(false)
- pb.SetTotal(pb.Total() + e.estimateJobs(tokens, window, stride))
+ b := newBatch(e.batchSize)
- n++
- }
-
- close(e.jobs)
-
- wg.Wait()
-
- close(e.results)
-
- return nil
-}
-
-func (e *Evaluator[R]) estimateJobs(tokens []int64, window, stride int) int {
- if len(tokens) < window {
- return 0
- }
+ for d := range data.Texts("text") {
+ for w, s := range tb.Push(n, d) {
+ if s == 0 {
+ s = 1 // first token as context
+ }
- windows := ((len(tokens) - window) / stride) + 1
- // jobs := (windows + batchSize - 1) / e.batchSize
+ b.AddJob(Job{
+ Document: n,
+ Position: m,
+ Tokens: w,
+ Seen: s,
+ })
- return windows
-}
+ m++
-func (e *Evaluator[R]) schedule(uid int, tokens []int64, contextSize, stride, batchSize int) {
- b := newBatch(batchSize)
+ if b.Size() == e.batchSize {
+ e.jobs <- *b
- seen := 1 // first token as context
- n := 0
+ b = newBatch(e.batchSize)
- // for i := 0; i+contextSize <= len(tokens); i += stride {
- for i := 0; i < len(tokens); i += stride {
- if b.Size() == batchSize {
- e.jobs <- *b
+ e.scheduled.Add(int64(e.batchSize))
- b = newBatch(batchSize)
+ pb.SetTotal(int(e.scheduled.Load()))
+ }
}
- j := min(i+contextSize, len(tokens))
-
- if j-i < contextSize {
- break // don't add jobs with partial windows
- }
+ n++
- // if j < seen {
- // break // don't add jobs with no new tokens
- // }
+ m = 0
+ }
- b.AddJob(Job{
- Document: uid,
- Position: n,
- Tokens: tokens[i:j],
- Seen: seen - i,
- })
+ if s := b.Size(); s > 0 {
+ e.jobs <- *b
- seen = i + contextSize
+ e.scheduled.Add(int64(s))
- n++
+ pb.SetTotal(int(e.scheduled.Load()))
}
- if b.Size() > 0 {
- e.jobs <- *b
- }
+ close(e.jobs)
- e.scheduled.Add(int64(n))
+ wg.Wait()
+
+ close(e.results)
+
+ return nil
}
func (e *Evaluator[R]) execute(j *batch, device int) {
diff --git a/llm/run_test.go b/llm/run_test.go
deleted file mode 100644
index fa974a7..0000000
--- a/llm/run_test.go
+++ /dev/null
@@ -1,45 +0,0 @@
-package llm
-
-import (
- "fmt"
- "testing"
-)
-
-func TestEvaluator_estimateJobs(t *testing.T) {
- type gold struct {
- tokens int
- window int
- stride int
- expected int
- }
-
- tests := []gold{
- {tokens: 20, window: 10, stride: 3, expected: 4},
- {tokens: 20, window: 10, stride: 4, expected: 3},
- {tokens: 20, window: 10, stride: 5, expected: 3},
-
- {tokens: 0, window: 1, stride: 1, expected: 0},
- {tokens: 1, window: 1, stride: 1, expected: 1},
-
- {tokens: 0, window: 1024, stride: 1, expected: 0},
- {tokens: 1, window: 1, stride: 1024, expected: 1},
- }
-
- e := NewEvaluator[any](nil, nil, nil)
-
- for _, tt := range tests {
- t.Run(
- fmt.Sprintf("tokens%d_window%d_stride%d", tt.tokens, tt.window, tt.stride),
-
- func(t *testing.T) {
- tokens := make([]int64, tt.tokens)
-
- got := e.estimateJobs(tokens, tt.window, tt.stride)
-
- if got != tt.expected {
- t.Errorf("expected %d but got (%d, %d, %d) = %d", tt.expected, tt.tokens, tt.window, tt.stride, got)
- }
- },
- )
- }
-}