summaryrefslogtreecommitdiff
path: root/research/knobloch/cmd/tables/main.go
blob: 27c3f0c027930b5c0c7f0a32edf576bc2dcca3d6 (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
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
// Command tables writes the tokenizer characterisation CSVs described in
// FREQUENCIES.md.
//
// Coverage is the full grid: every alignment level in both directions, at every
// vocabulary size we have tokenizers for. Statistics are over the MiniPile train
// dictionary.
//
// Encoding the dictionary dominates the runtime and the encodes are independent,
// so they run on a worker pool. One vocabulary size is processed at a time,
// since table C compares tokenizers within a size and the comparison inputs for
// a whole size have to be resident together.
package main

import (
	"flag"
	"fmt"
	"log"
	"runtime"
	"sync"

	"github.com/jonasknobloch/mbpe"
	"go.jknobloc.com/x/research/knobloch"
	"go.jknobloc.com/x/shelf"
)

var (
	vocabSizes = []int{8192, 16384, 32768, 50256, 100512}
	alphas     = []int{0, 10, 20, 30, 40, 50, 60, 70, 80, 90, 100}
)

// name is the tokenizer directory suffix, e.g. m050 or mi050.
func name(alpha int, inverted bool) string {
	if inverted {
		return fmt.Sprintf("mi%03d", alpha)
	}

	return fmt.Sprintf("m%03d", alpha)
}

func spec(size, alpha int, inverted bool) knobloch.TableSpec {
	n := name(alpha, inverted)

	return knobloch.TableSpec{
		Name:      n,
		Alignment: fmt.Sprintf("%.1f", float64(alpha)/100),
		Inverted:  inverted,
		VocabSize: size,
		Dir:       shelf.Item(fmt.Sprintf("tokenizers/minipile/tokenizer_gpt2_%d_%s_minipile", size, n)),
	}
}

func main() {
	dict := flag.String("dict", "results/knobloch/minipile/dict.txt", "shelf-relative dictionary")
	prefix := flag.String("prefix", "table", "output file prefix")
	workers := flag.Int("workers", runtime.NumCPU(), "parallel encodes")
	only := flag.Int("only", 0, "restrict to a single vocabulary size, for smoke runs")

	flag.Parse()

	sizes := vocabSizes

	if *only != 0 {
		sizes = []int{*only}
	}

	d := mbpe.NewDict()

	if err := d.Load(shelf.Abs(shelf.Item(*dict))); err != nil {
		log.Fatal(err)
	}

	items := d.Items()

	log.Printf("dictionary: %d pre-token types", len(items))

	var (
		rowsA []knobloch.TableARow
		rowsB []knobloch.TableBRow

		aligned    []knobloch.TableCRow
		inverted   []knobloch.TableCRow
		invVsAlign []knobloch.TableCRow
	)

	for _, size := range sizes {
		// every tokenizer at this size, aligned then inverted
		var specs []knobloch.TableSpec

		for _, inv := range []bool{false, true} {
			for _, a := range alphas {
				specs = append(specs, spec(size, a, inv))
			}
		}

		dirs := make([]shelf.Item, 0, len(specs))

		for _, s := range specs {
			dirs = append(dirs, s.Dir)
		}

		// table B measures sharing against the intersection of every vocabulary
		// at this size, so the baseline is just another tokenizer
		shared, err := knobloch.VocabIntersection(dirs)

		if err != nil {
			log.Fatal(err)
		}

		log.Printf("vocab %d: %d tokenizers, intersection %d tokens", size, len(specs), len(shared))

		encoded := make([]*knobloch.Encoded, len(specs))

		var wg sync.WaitGroup

		queue := make(chan int)

		for i := 0; i < *workers; i++ {
			wg.Add(1)

			go func() {
				defer wg.Done()

				for idx := range queue {
					e, err := knobloch.EncodeDict(specs[idx], items)

					if err != nil {
						log.Fatal(err)
					}

					encoded[idx] = e
				}
			}()
		}

		for i := range specs {
			queue <- i
		}

		close(queue)
		wg.Wait()

		log.Printf("vocab %d: encoded", size)

		// index by (alpha, direction) for the pairings below
		at := func(alpha int, inv bool) *knobloch.Encoded {
			for i, s := range specs {
				if s.Inverted == inv && s.Name == name(alpha, inv) {
					return encoded[i]
				}
			}

			log.Fatalf("missing encode for %d/%s", size, name(alpha, inv))

			return nil
		}

		base := at(0, false)

		for _, e := range encoded {
			rowsA = append(rowsA, knobloch.BuildTableA(e))

			b, err := knobloch.BuildTableB(e, shared)

			if err != nil {
				log.Fatal(err)
			}

			rowsB = append(rowsB, b)
		}

		for _, a := range alphas {
			al := at(a, false)
			inv := at(a, true)

			// aligned vs the alpha=0 baseline
			r, err := knobloch.BuildTableC(al, base, items)

			if err != nil {
				log.Fatal(err)
			}

			aligned = append(aligned, r)

			// inverted vs the alpha=0 baseline
			r, err = knobloch.BuildTableC(inv, base, items)

			if err != nil {
				log.Fatal(err)
			}

			inverted = append(inverted, r)

			// inverted vs aligned at matching alpha
			r, err = knobloch.BuildTableC(inv, al, items)

			if err != nil {
				log.Fatal(err)
			}

			invVsAlign = append(invVsAlign, r)
		}
	}

	for _, out := range []struct {
		name string
		err  error
	}{
		{*prefix + "_a.csv", knobloch.WriteTableA(*prefix+"_a.csv", rowsA)},
		{*prefix + "_b.csv", knobloch.WriteTableB(*prefix+"_b.csv", rowsB)},
		{*prefix + "_c_aligned.csv", knobloch.WriteTableC(*prefix+"_c_aligned.csv", aligned)},
		{*prefix + "_c_inverted.csv", knobloch.WriteTableC(*prefix+"_c_inverted.csv", inverted)},
		{*prefix + "_c_inverted_vs_aligned.csv", knobloch.WriteTableC(*prefix+"_c_inverted_vs_aligned.csv", invVsAlign)},
	} {
		if out.err != nil {
			log.Fatal(out.err)
		}

		log.Printf("wrote %s", out.name)
	}
}