summaryrefslogtreecommitdiff
path: root/llm/cmd/eval/ppl.go
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-03-17 20:16:39 +0100
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-03-17 20:16:39 +0100
commit57fb765bdc5819f8a367ad75c731a50da68b8d85 (patch)
tree730647fc08847624c138a83f3e8329b6f4755ea6 /llm/cmd/eval/ppl.go
parent035eef3308edd4ec581061440337aaa10a37f573 (diff)
Update perplexity evaluator
* Introduce results channel * Introduce jobs channel * Generalize evaluation via callback
Diffstat (limited to 'llm/cmd/eval/ppl.go')
-rw-r--r--llm/cmd/eval/ppl.go80
1 files changed, 80 insertions, 0 deletions
diff --git a/llm/cmd/eval/ppl.go b/llm/cmd/eval/ppl.go
new file mode 100644
index 0000000..fdbcff3
--- /dev/null
+++ b/llm/cmd/eval/ppl.go
@@ -0,0 +1,80 @@
+package main
+
+import (
+ "fmt"
+ "log"
+ "math"
+ "strings"
+ "sync"
+
+ "go.jknobloc.com/x/dataset"
+ "go.jknobloc.com/x/llm"
+)
+
+type pplResult struct {
+ v float64
+ n int
+}
+
+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 {
+ p, n := llm.NegLogLikelihood(logits, tokens)
+
+ return pplResult{
+ v: p,
+ n: n,
+ }
+ })
+
+ results := e.Results()
+
+ var wg sync.WaitGroup
+
+ total := float64(0)
+ n := 0
+
+ wg.Add(1)
+
+ go func() {
+ defer wg.Done()
+
+ for r := range results {
+ total += r.v
+ n += r.n
+ }
+ }()
+
+ if err := e.Run(d, 1024, 512); err != nil {
+ log.Fatal(err)
+ }
+
+ wg.Wait()
+
+ avg := total / float64(n)
+ ppl := math.Exp(avg)
+
+ fmt.Println(ppl)
+}
+
+func joined() dataset.Reader {
+ miniPile := data()
+
+ docs := make([]string, 0)
+
+ for d := range miniPile.Texts("text") {
+ docs = append(docs, d)
+ }
+
+ j := dataset.NewStringReader(strings.Join(docs, "\n\n"))
+
+ return j
+}