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