summaryrefslogtreecommitdiff
path: root/llm/evaluator.go
blob: 0eb541887ee624a4050a21d62bf02835cf23e45a (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
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](model Causal, tokenizer Tokenizer, callback func(job Job, logits [][]float32, tokens []int) R) *Evaluator[R] {
	return &Evaluator[R]{
		models:     []Causal{model}, // TODO multiple devices
		tokenizer:  tokenizer,
		batchSize:  1, // TODO arg
		numWorkers: 4, // TODO arg
		jobs:       make(chan batch, 1024),
		results:    make(chan R, 1024),
		callback:   callback,
	}
}

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