diff options
Diffstat (limited to 'research/knobloch/cmd/tableb')
| -rw-r--r-- | research/knobloch/cmd/tableb/main.go | 350 |
1 files changed, 350 insertions, 0 deletions
diff --git a/research/knobloch/cmd/tableb/main.go b/research/knobloch/cmd/tableb/main.go new file mode 100644 index 0000000..eca5ec9 --- /dev/null +++ b/research/knobloch/cmd/tableb/main.go @@ -0,0 +1,350 @@ +// Command tableb writes table B for a chosen set of tokenizers at one +// vocabulary size, rather than the whole grid. +// +// Everything comes from the dictionary, the segmentation and the tokenizers +// themselves. Nothing here reads a trained model, so the table stays a property +// of the tokenizer and the corpus. +// +// Two schemas: +// +// current what WriteTableB emits today: morph_count, the two morph shares, +// pct_rank_morph / pct_rank_other, shared_morph / shared_other +// legacy the columns of the published table_b.csv: types_used, the shares +// denominated by used types, absolute median frequencies, and the two +// shared_* columns measured against the baseline rather than an +// intersection +// +// The two disagree about more than names. The current schema divides the type +// share by the whole vocabulary and reports the median's percentile rank, which +// is comparable across vocabulary sizes; the legacy one divides by used types +// and reports the median itself, which is not. At a single size either is fine. +// +// Unlike table A, table B needs a vocabulary intersection for the current +// schema's shared_* columns. Building that only reads vocab.json files and costs +// nothing, so it still covers all 22 tokenizers at the size by default even when +// far fewer rows are written — otherwise "shared" would silently mean a +// different thing here than in the full-grid table. Pass -intersect listed to +// restrict it. The legacy schema ignores it and uses -baseline instead. +// +// The alignment column carries an " inv" suffix for the inverted family, which +// the full-grid command does not add. +package main + +import ( + "encoding/csv" + "flag" + "fmt" + "log" + "os" + "runtime" + "slices" + "strconv" + "strings" + "sync" + + "github.com/jonasknobloch/mbpe" + "go.jknobloc.com/x/research/knobloch" + "go.jknobloc.com/x/shelf" +) + +var alphaIDs = []string{"000", "010", "020", "030", "040", "050", "060", "070", "080", "090", "100"} + +func dir(size int, name string) shelf.Item { + return shelf.Item(fmt.Sprintf("tokenizers/minipile/tokenizer_gpt2_%d_%s_minipile", size, name)) +} + +// 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: dir(size, name), + }, nil +} + +// median matches the package's own: the upper-middle element of the sorted +// slice, not the average of the middle two. +func median(v []int) int { + if len(v) == 0 { + return 0 + } + + slices.Sort(v) + + return v[len(v)/2] +} + +func share(part, whole int) float64 { + if whole == 0 { + return 0 + } + + return 100 * float64(part) / float64(whole) +} + +// classify splits a vocabulary into the tokens the segmentation calls morphemes +// and the rest. +func classify(e *knobloch.Encoded, morph map[string]int) (morphemes, other map[string]struct{}) { + morphemes = make(map[string]struct{}) + other = make(map[string]struct{}) + + for id := range e.Counts { + token := e.Itoa[int64(id)] + + if _, ok := morph[token]; ok { + morphemes[token] = struct{}{} + } else { + other[token] = struct{}{} + } + } + + return morphemes, other +} + +func retained(base map[string]struct{}, e *knobloch.Encoded) float64 { + kept := 0 + + for id := range e.Counts { + if _, ok := base[e.Itoa[int64(id)]]; ok { + kept++ + } + } + + return share(kept, len(base)) +} + +func legacyRow(e *knobloch.Encoded, morph map[string]int, baseMorph, baseOther map[string]struct{}) []string { + var morphFreq, otherFreq []int + var morphTokens, otherTokens int64 + + typesUsed, typesMorph := 0, 0 + + for id, f := range e.Counts { + token := e.Itoa[int64(id)] + + if f > 0 { + typesUsed++ + } + + if _, ok := morph[token]; ok { + typesMorph++ + morphTokens += int64(f) + morphFreq = append(morphFreq, f) + } else { + otherTokens += int64(f) + otherFreq = append(otherFreq, f) + } + } + + return []string{ + e.Spec.Name, strconv.Itoa(e.Spec.VocabSize), e.Spec.Alignment, + strconv.FormatFloat(e.Fertility(), 'f', 4, 64), + strconv.Itoa(typesUsed), strconv.Itoa(typesMorph), + strconv.FormatFloat(share(typesMorph, typesUsed), 'f', 2, 64), + strconv.FormatFloat(share(int(morphTokens), int(morphTokens+otherTokens)), 'f', 2, 64), + strconv.Itoa(median(morphFreq)), strconv.Itoa(median(otherFreq)), + strconv.FormatFloat(retained(baseMorph, e), 'f', 2, 64), + strconv.FormatFloat(retained(baseOther, e), 'f', 2, 64), + } +} + +func main() { + dict := flag.String("dict", "results/knobloch/minipile/dict.txt", "shelf-relative dictionary") + size := flag.Int("size", 50256, "vocabulary size") + names := flag.String("tokenizers", "", "comma separated tokenizer names, empty for the 11 aligned") + schema := flag.String("schema", "legacy", "column set: legacy or current") + baseline := flag.String("baseline", "m000", "legacy schema: tokenizer the shared_* columns measure retention against") + intersect := flag.String("intersect", "all", "current schema: vocabularies behind shared_*, all (22) or listed") + out := flag.String("out", "table_b_aligned.csv", "output CSV") + workers := flag.Int("workers", runtime.NumCPU(), "parallel encodes") + + flag.Parse() + + var selected []string + + if *names == "" { + for _, a := range alphaIDs { + selected = append(selected, "m"+a) + } + } else { + for _, n := range strings.Split(*names, ",") { + selected = append(selected, strings.TrimSpace(n)) + } + } + + // the legacy shared_* columns need the baseline encoded even if it is not a + // row, so add it and drop it again before writing + encode := slices.Clone(selected) + + if *schema == "legacy" && !slices.Contains(encode, *baseline) { + encode = append(encode, *baseline) + } + + 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("segmentation: %s", shelf.Abs(knobloch.SegmentsPath)) + log.Printf("%d encodes at vocab %d: %s", len(encode), *size, strings.Join(encode, ",")) + + encoded := make([]*knobloch.Encoded, len(encode)) + + 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 { + s, err := spec(*size, encode[idx]) + + if err != nil { + log.Fatal(err) + } + + e, err := knobloch.EncodeDict(s, items) + + if err != nil { + log.Fatal(err) + } + + encoded[idx] = e + + log.Printf("%s: encoded", encode[idx]) + } + }() + } + + for i := range encode { + queue <- i + } + + close(queue) + wg.Wait() + + at := func(name string) *knobloch.Encoded { + for i, n := range encode { + if n == name { + return encoded[i] + } + } + + log.Fatalf("missing encode for %s", name) + + return nil + } + + if *schema == "current" { + var basis []shelf.Item + + if *intersect == "listed" { + for _, n := range selected { + basis = append(basis, dir(*size, n)) + } + } else { + for _, p := range []string{"m", "mi"} { + for _, a := range alphaIDs { + basis = append(basis, dir(*size, p+a)) + } + } + } + + shared, err := knobloch.VocabIntersection(basis) + + if err != nil { + log.Fatal(err) + } + + log.Printf("intersection over %d vocabularies: %d shared tokens", len(basis), len(shared)) + + rows := make([]knobloch.TableBRow, 0, len(selected)) + + for _, n := range selected { + row, err := knobloch.BuildTableB(at(n), shared) + + if err != nil { + log.Fatal(err) + } + + rows = append(rows, row) + } + + if err := knobloch.WriteTableB(*out, rows); err != nil { + log.Fatal(err) + } + + log.Printf("wrote %s", *out) + + return + } + + morph, err := knobloch.Morphemes(shelf.Abs(knobloch.SegmentsPath)) + + if err != nil { + log.Fatal(err) + } + + baseMorph, baseOther := classify(at(*baseline), morph) + + log.Printf("baseline %s: %d morphemes, %d other", *baseline, len(baseMorph), len(baseOther)) + + file, err := os.Create(*out) + + if err != nil { + log.Fatal(err) + } + + defer file.Close() + + w := csv.NewWriter(file) + + if err := w.Write([]string{ + "tokenizer", "vocab_size", "alignment", "fertility", + "types_used", "types_morph", "types_morph_share", "tokens_morph_share", + "median_freq_morph", "median_freq_other", + "types_shared_morph_share", "types_shared_other_share", + }); err != nil { + log.Fatal(err) + } + + for _, n := range selected { + if err := w.Write(legacyRow(at(n), morph, baseMorph, baseOther)); err != nil { + log.Fatal(err) + } + } + + w.Flush() + + if err := w.Error(); err != nil { + log.Fatal(err) + } + + log.Printf("wrote %s", *out) +} |
