summaryrefslogtreecommitdiff
path: root/research/knobloch/cmd/tablea/main.go
diff options
context:
space:
mode:
Diffstat (limited to 'research/knobloch/cmd/tablea/main.go')
-rw-r--r--research/knobloch/cmd/tablea/main.go144
1 files changed, 144 insertions, 0 deletions
diff --git a/research/knobloch/cmd/tablea/main.go b/research/knobloch/cmd/tablea/main.go
new file mode 100644
index 0000000..37fa877
--- /dev/null
+++ b/research/knobloch/cmd/tablea/main.go
@@ -0,0 +1,144 @@
+// Command tablea writes table A for a chosen set of tokenizers rather than the
+// whole grid.
+//
+// Table A is a per-tokenizer statistic, so unlike tables B and C it needs no
+// intersection and therefore no encode of the tokenizers that are not shown.
+// Restricting the set cuts the work proportionally: three alignment levels at
+// five vocabulary sizes is fifteen encodes instead of the full hundred and ten.
+//
+// The alignment column carries an " inv" suffix for the inverted family, which
+// the full-grid command does not add.
+package main
+
+import (
+ "flag"
+ "fmt"
+ "log"
+ "runtime"
+ "strconv"
+ "strings"
+ "sync"
+
+ "github.com/jonasknobloch/mbpe"
+ "go.jknobloc.com/x/research/knobloch"
+ "go.jknobloc.com/x/shelf"
+)
+
+var vocabSizes = []int{8192, 16384, 32768, 50256, 100512}
+
+// spec derives the alignment level from the tokenizer name, e.g. m050 or mi100.
+func spec(size int, name string) (knobloch.TableSpec, error) {
+ inverted := strings.HasPrefix(name, "mi")
+
+ digits := strings.TrimPrefix(strings.TrimPrefix(name, "mi"), "m")
+
+ alpha, err := strconv.Atoi(digits)
+
+ if err != nil {
+ return knobloch.TableSpec{}, fmt.Errorf("%s: cannot read an alignment level from the name", name)
+ }
+
+ alignment := fmt.Sprintf("%.1f", float64(alpha)/100)
+
+ if inverted {
+ alignment += " inv"
+ }
+
+ return knobloch.TableSpec{
+ Name: name,
+ Alignment: alignment,
+ Inverted: inverted,
+ VocabSize: size,
+ Dir: shelf.Item(fmt.Sprintf("tokenizers/minipile/tokenizer_gpt2_%d_%s_minipile", size, name)),
+ }, nil
+}
+
+func main() {
+ dict := flag.String("dict", "results/knobloch/minipile/dict.txt", "shelf-relative dictionary")
+ names := flag.String("tokenizers", "m000,m050,m100", "comma separated tokenizer names")
+ sizes := flag.String("sizes", "", "comma separated vocabulary sizes, empty for all")
+ out := flag.String("out", "table_a_alphas.csv", "output CSV")
+ workers := flag.Int("workers", runtime.NumCPU(), "parallel encodes")
+
+ flag.Parse()
+
+ selected := vocabSizes
+
+ if *sizes != "" {
+ selected = nil
+
+ for _, s := range strings.Split(*sizes, ",") {
+ v, err := strconv.Atoi(strings.TrimSpace(s))
+
+ if err != nil {
+ log.Fatalf("bad vocabulary size %q", s)
+ }
+
+ selected = append(selected, v)
+ }
+ }
+
+ var specs []knobloch.TableSpec
+
+ for _, size := range selected {
+ for _, name := range strings.Split(*names, ",") {
+ s, err := spec(size, strings.TrimSpace(name))
+
+ if err != nil {
+ log.Fatal(err)
+ }
+
+ specs = append(specs, s)
+ }
+ }
+
+ 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))
+ log.Printf("%d encodes: %s at %v", len(specs), *names, selected)
+
+ rows := make([]knobloch.TableARow, 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)
+ }
+
+ rows[idx] = knobloch.BuildTableA(e)
+
+ log.Printf("%s at %d: encoded", specs[idx].Name, specs[idx].VocabSize)
+ }
+ }()
+ }
+
+ for i := range specs {
+ queue <- i
+ }
+
+ close(queue)
+ wg.Wait()
+
+ if err := knobloch.WriteTableA(*out, rows); err != nil {
+ log.Fatal(err)
+ }
+
+ log.Printf("wrote %s", *out)
+}