blob: 3a710ace4e30ce0c0eb79121d11902f59aa99170 (
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.ConfigDefault(), 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
}
|