summaryrefslogtreecommitdiff
path: root/llm/cmd
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-07 00:01:23 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-07 00:01:23 +0200
commitb81fa93e098ca689cb6f73627968fe0b4d49decb (patch)
treebb852b92c2761a0924a5cefcc3bce6d10b480822 /llm/cmd
parentb9faf92148c7e6432ff0411783f444b7779fa36a (diff)
Run evaluator with score
Diffstat (limited to 'llm/cmd')
-rw-r--r--llm/cmd/eval/logprobs.go46
-rw-r--r--llm/cmd/eval/main.go2
-rw-r--r--llm/cmd/eval/ppl.go14
3 files changed, 14 insertions, 48 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,
}
})