1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
|
package llm
import (
"errors"
"sync"
"sync/atomic"
"go.jknobloc.com/x/dataset"
)
type EvaluatorConfig struct {
BatchSize int
NumWorkers int
}
type Evaluator[R any] struct {
models []Causal
tokenizer Tokenizer
batchSize int
numWorkers int
scheduled atomic.Int64
completed atomic.Int64
jobs chan batch
results chan R
callback func(job Job, logProbs []float32, tokens []int) R
}
func NewEvaluator[R any](model Causal, tokenizer Tokenizer, callback func(job Job, logProbs []float32, tokens []int) R, cfg EvaluatorConfig) *Evaluator[R] {
return &Evaluator[R]{
models: []Causal{model}, // TODO multiple devices
tokenizer: tokenizer,
batchSize: cfg.BatchSize,
numWorkers: cfg.NumWorkers,
jobs: make(chan batch, 1024),
results: make(chan R, 1024),
callback: callback,
}
}
func (e *Evaluator[R]) Results() chan R {
return e.results
}
func (e *Evaluator[R]) RunAndCollect(title string, data dataset.Reader, window, stride int, callback func(R) error) error {
var wg sync.WaitGroup
var collectErr error
wg.Add(1)
go func() {
defer wg.Done()
for r := range e.results {
if err := callback(r); err != nil && collectErr == nil {
collectErr = err
}
}
}()
runErr := e.Run(title, data, window, stride)
wg.Wait()
return errors.Join(runErr, collectErr)
}
|