summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-03-25 21:23:41 +0100
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-03-25 21:44:36 +0100
commita53c085de9f3e27dadf1f3f3025708d5c7542f92 (patch)
tree4f23d85ef225bf123a8b57f03967c77d51013e70
parent2327622c3671d887adf1240c5e90418e1ced8671 (diff)
Refactor evaluator initialization
-rw-r--r--llm/cmd/eval/logprobs.go8
-rw-r--r--llm/cmd/eval/ppl.go8
-rw-r--r--llm/evaluator.go18
-rw-r--r--llm/perplexity_test.go2
4 files changed, 9 insertions, 27 deletions
diff --git a/llm/cmd/eval/logprobs.go b/llm/cmd/eval/logprobs.go
index eaa7459..7c04626 100644
--- a/llm/cmd/eval/logprobs.go
+++ b/llm/cmd/eval/logprobs.go
@@ -24,15 +24,11 @@ type logProbs []logProb
func logprobs() {
d := data()
+
m := model()
t := tokenizer()
- e := llm.NewEvaluator[logProbs]()
-
- e.SetTokenizer(t)
- e.AddModel(m)
-
- e.SetCallback(func(j llm.Job, logits [][]float32, tokens []int) logProbs {
+ e := llm.NewEvaluator(m, t, func(j llm.Job, logits [][]float32, tokens []int) logProbs {
probs := selectLogProbs(logits, tokens)
r := make(logProbs, len(tokens))
diff --git a/llm/cmd/eval/ppl.go b/llm/cmd/eval/ppl.go
index fdbcff3..0a028d2 100644
--- a/llm/cmd/eval/ppl.go
+++ b/llm/cmd/eval/ppl.go
@@ -18,15 +18,11 @@ type pplResult struct {
func perplexity() {
d := data()
+
m := model()
t := tokenizer()
- e := llm.NewEvaluator[pplResult]()
-
- e.SetTokenizer(t)
- e.AddModel(m)
-
- e.SetCallback(func(job llm.Job, logits [][]float32, tokens []int) pplResult {
+ e := llm.NewEvaluator(m, t, func(job llm.Job, logits [][]float32, tokens []int) pplResult {
p, n := llm.NegLogLikelihood(logits, tokens)
return pplResult{
diff --git a/llm/evaluator.go b/llm/evaluator.go
index 24fe2d4..0eb5418 100644
--- a/llm/evaluator.go
+++ b/llm/evaluator.go
@@ -20,28 +20,18 @@ type Evaluator[R any] struct {
callback func(job Job, logits [][]float32, tokens []int) R
}
-func NewEvaluator[R any]() *Evaluator[R] {
+func NewEvaluator[R any](model Causal, tokenizer Tokenizer, callback func(job Job, logits [][]float32, tokens []int) R) *Evaluator[R] {
return &Evaluator[R]{
- models: make([]Causal, 0),
+ 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]) 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
}
diff --git a/llm/perplexity_test.go b/llm/perplexity_test.go
index 81f10af..fa974a7 100644
--- a/llm/perplexity_test.go
+++ b/llm/perplexity_test.go
@@ -25,7 +25,7 @@ func TestEvaluator_estimateJobs(t *testing.T) {
{tokens: 1, window: 1, stride: 1024, expected: 1},
}
- e := NewEvaluator[any]()
+ e := NewEvaluator[any](nil, nil, nil)
for _, tt := range tests {
t.Run(