summaryrefslogtreecommitdiff
path: root/llm/evaluator.go
blob: 8aa4a320856606e26d7102e8735dfe3b913aa4df (plain)
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)
}