summaryrefslogtreecommitdiff
path: root/research/knobloch/cmd/tableb/main.go
diff options
context:
space:
mode:
Diffstat (limited to 'research/knobloch/cmd/tableb/main.go')
-rw-r--r--research/knobloch/cmd/tableb/main.go350
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)
+}