summaryrefslogtreecommitdiff
path: root/llm/cmd/eval
diff options
context:
space:
mode:
Diffstat (limited to 'llm/cmd/eval')
-rw-r--r--llm/cmd/eval/main.go53
-rw-r--r--llm/cmd/eval/ppl.go80
2 files changed, 133 insertions, 0 deletions
diff --git a/llm/cmd/eval/main.go b/llm/cmd/eval/main.go
new file mode 100644
index 0000000..05b3f96
--- /dev/null
+++ b/llm/cmd/eval/main.go
@@ -0,0 +1,53 @@
+package main
+
+import (
+ "log"
+
+ "github.com/jonasknobloch/mbpe"
+
+ "go.jknobloc.com/x/dataset"
+ "go.jknobloc.com/x/gpt2"
+)
+
+func main() {
+ perplexity()
+}
+
+func data() *dataset.ParquetReader {
+ var miniPile *dataset.ParquetReader
+
+ if r, err := dataset.NewParquetReader("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
+}
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
+}