summaryrefslogtreecommitdiff
path: root/llm/cmd
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-03-09 18:59:37 +0100
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-03-09 19:22:07 +0100
commitbc88dbac521b0de3db441c37f8e946cc2b0d3669 (patch)
treea947f5dfb7dd1b46fc39c176928d3649a3924647 /llm/cmd
parentfd6fe450a35d8d79947eaebf15361cc04cf29959 (diff)
Add ppl command
Diffstat (limited to 'llm/cmd')
-rw-r--r--llm/cmd/ppl/ppl.go69
1 files changed, 69 insertions, 0 deletions
diff --git a/llm/cmd/ppl/ppl.go b/llm/cmd/ppl/ppl.go
new file mode 100644
index 0000000..622644a
--- /dev/null
+++ b/llm/cmd/ppl/ppl.go
@@ -0,0 +1,69 @@
+package main
+
+import (
+ "fmt"
+ "log"
+
+ "github.com/jonasknobloch/mbpe"
+ "github.com/jonasknobloch/x/dataset"
+ "github.com/jonasknobloch/x/gpt2"
+ "github.com/jonasknobloch/x/llm"
+)
+
+func main() {
+ d := data()
+ m := model()
+ t := tokenizer()
+
+ e := llm.NewEvaluator()
+
+ e.SetTokenizer(t)
+ e.AddModel(m)
+
+ ppl, err := e.Perplexity(d)
+
+ if err != nil {
+ log.Fatal(err)
+ }
+
+ fmt.Println(ppl)
+}
+
+func data() *dataset.Reader {
+ var miniPile *dataset.Reader
+
+ if r, err := dataset.NewReader("dataset/cmd/dataset/tmp/minipile/validation"); err != nil {
+ log.Fatal(err)
+ } else {
+ miniPile = r
+ }
+
+ return miniPile
+}
+
+func model() *gpt2.Model {
+ m := gpt2.NewModel("gpt2/models/base/model.onnx", "0")
+
+ if err := m.Init(); err != nil {
+ log.Fatal(err)
+ }
+
+ return m
+}
+
+func tokenizer() *mbpe.Tokenizer {
+ m := mbpe.NewMBPE()
+
+ if err := m.Load("gpt2/models/base/vocab.json", "gpt2/models/base/merges.txt"); err != nil {
+ log.Fatal(err)
+ }
+
+ t := mbpe.NewTokenizer(m)
+
+ byteLevel := mbpe.NewByteLevel(false)
+
+ t.SetPreTokenizer(byteLevel)
+ t.SetDecoder(byteLevel)
+
+ return t
+}