summaryrefslogtreecommitdiff
path: root/llm/evaluator.go
blob: 24fe2d4eb265280c37970202d6ef9f1df6ffbf57 (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
package llm

import (
	"sync/atomic"
)

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, logits [][]float32, tokens []int) R
}

func NewEvaluator[R any]() *Evaluator[R] {
	return &Evaluator[R]{
		models:     make([]Causal, 0),
		batchSize:  1, // TODO arg
		numWorkers: 4, // TODO arg
		jobs:       make(chan batch, 1024),
		results:    make(chan R, 1024),
	}
}

func (e *Evaluator[R]) AddModel(model Causal) {
	e.models = append(e.models, model)
}

func (e *Evaluator[R]) SetTokenizer(tokenizer Tokenizer) {
	e.tokenizer = tokenizer
}

func (e *Evaluator[R]) SetCallback(callback func(job Job, logits [][]float32, tokens []int) R) {
	e.callback = callback
}

func (e *Evaluator[R]) Results() chan R {
	return e.results
}