summaryrefslogtreecommitdiff
path: root/llmc/cmd/peek/main.go
blob: caacb318889bd856bf97a4c22316913817446246 (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
package main

import (
	"fmt"
	"log"

	"golang.org/x/exp/constraints"

	"go.jknobloc.com/x/llmc"
	"go.jknobloc.com/x/tokenizer/bpe"
)

func main() {
	var a llmc.DataFile[uint16]
	var b llmc.DataFile[uint16]

	_ = must(llmc.Deserialize("artifacts/data/llmc/edu_fineweb100B/edu_fineweb_val_000000.bin", &a))
	_ = must(llmc.Deserialize("artifacts/test/llmc/edu_fineweb100B/edu_fineweb_val_000000.bin", &b))

	t := must(bpe.NewTokenizerFromFiles("gpt2/models/base/vocab.json", "gpt2/models/base/merges.txt"))

	s := decode(&a, 1024, t)
	k := decode(&b, 1024, t)

	fmt.Println(s)

	fmt.Println()

	fmt.Println(k)
}

func must[T any](v T, err error) T {
	if err != nil {
		log.Fatal(err)
	}

	return v
}

func decode[T constraints.Integer](src *llmc.DataFile[T], n int, t *bpe.Tokenizer) string {
	ids := make([]int, n)

	if len(src.Tokens) < n {
		panic("not enough tokens")
	}

	for i := range n {
		ids[i] = int(src.Tokens[i])
	}

	return t.Decode(ids)
}