diff options
Diffstat (limited to 'research/knobloch/cmd/tablea')
| -rw-r--r-- | research/knobloch/cmd/tablea/main.go | 144 |
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) +} |
