summaryrefslogtreecommitdiff
path: root/llm/run.go
diff options
context:
space:
mode:
Diffstat (limited to 'llm/run.go')
-rw-r--r--llm/run.go102
1 files changed, 36 insertions, 66 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) {