diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-07 00:01:23 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-07 00:01:23 +0200 |
| commit | b81fa93e098ca689cb6f73627968fe0b4d49decb (patch) | |
| tree | bb852b92c2761a0924a5cefcc3bce6d10b480822 | |
| parent | b9faf92148c7e6432ff0411783f444b7779fa36a (diff) | |
Run evaluator with score
| -rw-r--r-- | llm/cmd/eval/logprobs.go | 46 | ||||
| -rw-r--r-- | llm/cmd/eval/main.go | 2 | ||||
| -rw-r--r-- | llm/cmd/eval/ppl.go | 14 | ||||
| -rw-r--r-- | llm/evaluator.go | 4 | ||||
| -rw-r--r-- | llm/run.go | 6 | ||||
| -rw-r--r-- | llm/stat.go | 34 |
6 files changed, 19 insertions, 87 deletions
diff --git a/llm/cmd/eval/logprobs.go b/llm/cmd/eval/logprobs.go index bbc71aa..7228ea6 100644 --- a/llm/cmd/eval/logprobs.go +++ b/llm/cmd/eval/logprobs.go @@ -4,8 +4,6 @@ import ( "context" "database/sql" "log" - "math" - "slices" _ "github.com/duckdb/duckdb-go/v2" @@ -27,16 +25,14 @@ func logprobs() { m := model() t := tokenizer() - e := llm.NewEvaluator(m, t, func(j llm.Job, logits [][]float32, tokens []int) logProbs { - probs := selectLogProbs(logits, tokens) - + e := llm.NewEvaluator(m, t, func(j llm.Job, l []float32, tokens []int) logProbs { r := make(logProbs, len(tokens)) for i, token := range tokens { r[i] = logProb{ document: j.Document, token: token, - value: probs[i], + value: l[i], offset: i, } } @@ -113,41 +109,3 @@ func insert(stmt *sql.Stmt, prob logProb) error { return nil } - -func selectLogProbs(logits [][]float32, tokens []int) []float32 { - if len(logits) != len(tokens) { - panic("length mismatch") - } - - r := make([]float32, len(tokens)) - - for i, token := range tokens { - logprobs := logSoftmax(logits[i]) - - r[i] = logprobs[token] - } - - return r -} - -func logSoftmax(logits []float32) []float32 { - m := slices.Max(logits) - - s := float32(0.0) - r := make([]float32, len(logits)) - - for i, v := range logits { - e := float32(math.Exp(float64(v - m))) - - r[i] = v - s += e - } - - lse := float32(math.Log(float64(s))) + m - - for i := range r { - r[i] -= lse - } - - return r -} diff --git a/llm/cmd/eval/main.go b/llm/cmd/eval/main.go index 43752d9..b654e49 100644 --- a/llm/cmd/eval/main.go +++ b/llm/cmd/eval/main.go @@ -34,7 +34,7 @@ func data() *dataset.ParquetReader { } func model() *gpt2.Model { - m := gpt2.NewModel("gpt2/models/base/model.onnx", "0", gpt2.NewDefaultConfig(), true, true, false) + m := gpt2.NewModel("gpt2/models/base/model_eval.onnx", "0", gpt2.NewDefaultConfig(), false, false, true) if err := m.Init(); err != nil { log.Fatal(err) diff --git a/llm/cmd/eval/ppl.go b/llm/cmd/eval/ppl.go index 261114e..8838f86 100644 --- a/llm/cmd/eval/ppl.go +++ b/llm/cmd/eval/ppl.go @@ -21,11 +21,19 @@ func perplexity() { m := model() t := tokenizer() - e := llm.NewEvaluator(m, t, func(job llm.Job, logits [][]float32, tokens []int) pplResult { - p, n := llm.NegLogLikelihood(logits, tokens) + e := llm.NewEvaluator(m, t, func(job llm.Job, logProbs []float32, tokens []int) pplResult { + total := float64(0) + + n := 0 + + for _, p := range logProbs { + total -= float64(p) + + n++ + } return pplResult{ - v: p, + v: total, n: n, } }) diff --git a/llm/evaluator.go b/llm/evaluator.go index 17a7759..8340ad4 100644 --- a/llm/evaluator.go +++ b/llm/evaluator.go @@ -21,10 +21,10 @@ type Evaluator[R any] struct { jobs chan batch results chan R - callback func(job Job, logits [][]float32, tokens []int) R + callback func(job Job, logProbs []float32, tokens []int) R } -func NewEvaluator[R any](model Causal, tokenizer Tokenizer, callback func(job Job, logits [][]float32, tokens []int) R) *Evaluator[R] { +func NewEvaluator[R any](model Causal, tokenizer Tokenizer, callback func(job Job, logProbs []float32, tokens []int) R) *Evaluator[R] { return &Evaluator[R]{ models: []Causal{model}, // TODO multiple devices tokenizer: tokenizer, @@ -125,13 +125,13 @@ func (e *Evaluator[R]) execute(j *batch, device int) { m := e.models[device] - logits := make([][]float32, 0, len(job.Tokens)) + logProbs := make([]float32, 0, len(job.Tokens)-1) - if _, err := m.Generate(job.Tokens, 0, &logits); err != nil { + if err := m.Score(job.Tokens, &logProbs); err != nil { panic(err) // TODO handle } - l := logits[job.Seen-1 : len(logits)-1] + l := logProbs[job.Seen-1:] t := toInt(job.Tokens[job.Seen:]) r := e.callback(job, l, t) diff --git a/llm/stat.go b/llm/stat.go deleted file mode 100644 index 228e9f2..0000000 --- a/llm/stat.go +++ /dev/null @@ -1,34 +0,0 @@ -package llm - -import ( - "math" - "slices" -) - -func NegLogLikelihood(logits [][]float32, targets []int) (float64, int) { - if len(logits) != len(targets) { - panic("mismatched input lengths") - } - - total := float64(0) - - for i, target := range targets { - maxLogit := float64(slices.Max(logits[i])) - - sumExp := float64(0) - - for _, v := range logits[i] { - sumExp += math.Exp(float64(v) - maxLogit) - } - - logSumExp := maxLogit + math.Log(sumExp) - - targetLogit := float64(logits[i][target]) - - logProb := targetLogit - logSumExp - - total -= logProb - } - - return total, len(targets) -} |
