blob: 88be5000005fb071a5d5da3a48ef45bdebf76a40 (
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
|
package main
import (
"log"
"go.jknobloc.com/x/dataset"
"go.jknobloc.com/x/gpt2"
"go.jknobloc.com/x/tokenizer/bpe"
)
func main() {
if err := gpt2.InitializeEnvironment(); err != nil {
log.Fatal(err)
}
perplexity()
// logprobs()
if err := gpt2.DestroyEnvironment(); err != nil {
log.Fatal(err)
}
}
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 {
opts := gpt2.Options{
WithCache: false,
WithLogits: false,
WithLogProbs: true,
}
m := gpt2.NewModel("gpt2/models/base/model_eval.onnx", "0", gpt2.DefaultConfig(), opts)
if err := m.Init(); err != nil {
log.Fatal(err)
}
return m
}
func tokenizer() *bpe.Tokenizer {
var tok *bpe.Tokenizer
if t, err := bpe.NewTokenizerFromFiles("gpt2/models/base/vocab.json", "gpt2/models/base/merges.txt"); err != nil {
log.Fatal(err)
} else {
tok = t
}
return tok
}
|