diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-03 15:22:15 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-03 15:45:29 +0200 |
| commit | a61bc9ff890ea50d41a406cd1ffe554341551d2b (patch) | |
| tree | 3159fa6cb444d88fae7ac7abee92420d551831c7 /llm | |
| parent | dc0c59b97243645a9fdc040460e51d290085374e (diff) | |
Use token buffer to generate jobs
Diffstat (limited to 'llm')
| -rw-r--r-- | llm/run.go | 102 | ||||
| -rw-r--r-- | llm/run_test.go | 45 |
2 files changed, 36 insertions, 111 deletions
@@ -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) - } - }, - ) - } -} |
