// 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) }