summaryrefslogtreecommitdiff
path: root/llm/cmd/eval/main.go
blob: bd734b3362b0fde1d4d9eff68abd0a4f56c4ceba (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
package main

import (
	"log"

	"go.jknobloc.com/x/dataset"
	"go.jknobloc.com/x/gpt2"
	"go.jknobloc.com/x/shelf"
	"go.jknobloc.com/x/tokenizer/bpe"
)

func main() {
	if err := gpt2.InitializeEnvironment(); err != nil {
		log.Fatal(err)
	}

	perplexity()

	if err := gpt2.DestroyEnvironment(); err != nil {
		log.Fatal(err)
	}
}

func data() *dataset.ParquetReader {
	var miniPile *dataset.ParquetReader

	if r, err := dataset.NewParquetReader(shelf.Abs("data/minipile/validation")); err != nil {
		log.Fatal(err)
	} else {
		miniPile = r
	}

	return miniPile
}

func model() *gpt2.Model {
	opts := gpt2.Options{
		WithCache:    false,
		WithLogits:   false,
		WithLogProbs: true,
	}

	m := gpt2.NewModel(shelf.Abs("models/gpt2/model_eval.onnx"), gpt2.DefaultConfig(), opts)

	if err := m.Init(); err != nil {
		log.Fatal(err)
	}

	return m
}

func tokenizer() *bpe.Tokenizer {
	var tok *bpe.Tokenizer

	v := shelf.Abs("models/gpt2/vocab.json")
	m := shelf.Abs("models/gpt2/merges.txt")

	if t, err := bpe.NewTokenizerFromFiles(v, m, bpe.DefaultConfig()); err != nil {
		log.Fatal(err)
	} else {
		tok = t
	}

	return tok
}