summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-08 22:06:35 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-08 22:06:35 +0200
commit7e48ca5b9143c2b5f1a6eb81f35e90e67e628404 (patch)
tree7cdb3cd65e715806f50ecf598b0b8e0fcc5e8faa
parentb5de54cfc298c696fcd090bb7d12e8a6b9a8bfe1 (diff)
Add evaluator config
-rw-r--r--llm/cmd/eval/logprobs.go3
-rw-r--r--llm/cmd/eval/ppl.go3
-rw-r--r--llm/evaluator.go11
3 files changed, 14 insertions, 3 deletions
diff --git a/llm/cmd/eval/logprobs.go b/llm/cmd/eval/logprobs.go
index 7228ea6..baac805 100644
--- a/llm/cmd/eval/logprobs.go
+++ b/llm/cmd/eval/logprobs.go
@@ -38,6 +38,9 @@ func logprobs() {
}
return r
+ }, llm.EvaluatorConfig{
+ BatchSize: 1,
+ NumWorkers: 4,
})
var insertStmt *sql.Stmt
diff --git a/llm/cmd/eval/ppl.go b/llm/cmd/eval/ppl.go
index 8838f86..330e4ab 100644
--- a/llm/cmd/eval/ppl.go
+++ b/llm/cmd/eval/ppl.go
@@ -36,6 +36,9 @@ func perplexity() {
v: total,
n: n,
}
+ }, llm.EvaluatorConfig{
+ BatchSize: 1,
+ NumWorkers: 4,
})
total := float64(0)
diff --git a/llm/evaluator.go b/llm/evaluator.go
index 8340ad4..8aa4a32 100644
--- a/llm/evaluator.go
+++ b/llm/evaluator.go
@@ -8,6 +8,11 @@ import (
"go.jknobloc.com/x/dataset"
)
+type EvaluatorConfig struct {
+ BatchSize int
+ NumWorkers int
+}
+
type Evaluator[R any] struct {
models []Causal
tokenizer Tokenizer
@@ -24,12 +29,12 @@ type Evaluator[R any] struct {
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) *Evaluator[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: 1, // TODO arg
- numWorkers: 4, // TODO arg
+ batchSize: cfg.BatchSize,
+ numWorkers: cfg.NumWorkers,
jobs: make(chan batch, 1024),
results: make(chan R, 1024),
callback: callback,