summaryrefslogtreecommitdiff
path: root/llm/cmd/eval/perplexity.go
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-24 11:14:43 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-24 11:14:43 +0200
commit408536516bdfedc93d93226dc08280ed224361fc (patch)
treefc66349f2c2995644bf3caf1563bfccbad7f0496 /llm/cmd/eval/perplexity.go
parentb6c18b986996e8856483aff30c731acd4d79696f (diff)
Cleanup eval command
* Remove logprob extraction
Diffstat (limited to 'llm/cmd/eval/perplexity.go')
-rw-r--r--llm/cmd/eval/perplexity.go62
1 files changed, 62 insertions, 0 deletions
diff --git a/llm/cmd/eval/perplexity.go b/llm/cmd/eval/perplexity.go
new file mode 100644
index 0000000..ba92c99
--- /dev/null
+++ b/llm/cmd/eval/perplexity.go
@@ -0,0 +1,62 @@
+package main
+
+import (
+ "fmt"
+ "log"
+ "math"
+
+ "go.jknobloc.com/x/llm"
+)
+
+type pplResult struct {
+ v float64
+ n int
+}
+
+func perplexity() {
+ d := data()
+
+ m := model()
+ t := tokenizer()
+
+ 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: total,
+ n: n,
+ }
+ }, llm.EvaluatorConfig{
+ BatchSize: 1,
+ NumWorkers: 4,
+ })
+
+ total := float64(0)
+ n := 0
+
+ if err := e.RunAndCollect("Perplexity", d, 1024, 512, func(r pplResult) error {
+ total += r.v
+ n += r.n
+
+ return nil
+ }); err != nil {
+ log.Fatal(err)
+ }
+
+ avg := total / float64(n)
+ ppl := math.Exp(avg)
+
+ fmt.Println(ppl)
+
+ if err := m.Destroy(); err != nil {
+ log.Fatal(err)
+ }
+}