diff options
Diffstat (limited to 'research/knobloch/cmd/rankshift/main.go')
| -rw-r--r-- | research/knobloch/cmd/rankshift/main.go | 486 |
1 files changed, 486 insertions, 0 deletions
diff --git a/research/knobloch/cmd/rankshift/main.go b/research/knobloch/cmd/rankshift/main.go new file mode 100644 index 0000000..b81e7ea --- /dev/null +++ b/research/knobloch/cmd/rankshift/main.go @@ -0,0 +1,486 @@ +// Command rankshift reports how far the alignment objective moves sub-words in +// the merge order, split by whether the sub-word is a known morpheme. +// +// A variant's vocabulary can only be compared to the baseline's on the sub-words +// they have in common: an unshared sub-word has no baseline position, so there is +// nothing to displace. Over the shared set, each string holds a position in both +// vocabularies, and Spearman's rho between those two position vectors says how +// much the objective has reordered them. +// +// Rho alone is symmetric and says nothing about direction, so the signed +// displacement is reported alongside it. Position is merge order, so a lower id +// is an earlier, more important merge: a negative displacement means the sub-word +// gained rank. The expectation is that morphemes gain and non-morphemes lose, +// with the inverted objective reversing it. +// +// -intersect all (the default) compares every tokenizer on the sub-words the +// whole sweep has in common, matching what the *_shared plots do, so rows are +// comparable to each other. -intersect baseline uses each variant's own +// intersection with the baseline, which keeps far more data per row but makes +// the rows incomparable, since the set shrinks with alpha and does so by +// exactly the displacement being measured. +// +// Reordering only describes sub-words that survive. The added and dropped rows +// cover the rest: how many sub-words the variant swaps out of the baseline's +// vocabulary, and what share of each side is morphemic. +// +// Reads vocabularies only, so it runs in seconds and never touches the corpus. +package main + +import ( + "encoding/csv" + "flag" + "fmt" + "log" + "os" + "slices" + "strconv" + + "gonum.org/v1/gonum/stat" + + "go.jknobloc.com/x/research/knobloch" + "go.jknobloc.com/x/shelf" + "go.jknobloc.com/x/tokenizer/bpe" +) + +var alphas = []int{0, 10, 20, 30, 40, 50, 60, 70, 80, 90, 100} + +func name(alpha int, inverted bool) string { + if inverted { + return fmt.Sprintf("mi%03d", alpha) + } + + return fmt.Sprintf("m%03d", alpha) +} + +// names lists every tokenizer in the sweep, baseline first. +func names(baseline string) []string { + var all []string + + for _, inv := range []bool{false, true} { + for _, a := range alphas { + if n := name(a, inv); n != baseline { + all = append(all, n) + } + } + } + + return append([]string{baseline}, all...) +} + +// positions maps each sub-word string to its index in the merge order. +func positions(dir string) (map[string]int, error) { + tok, err := bpe.NewTokenizerFromFiles( + shelf.Abs(shelf.Item(dir)+"/vocab.json"), + shelf.Abs(shelf.Item(dir)+"/merges.txt"), + bpe.Config{Recover: false}) + + if err != nil { + return nil, err + } + + vocab := bpe.Vocab(tok) + p := make(map[string]int, len(vocab)) + + for id, token := range vocab { + p[token] = id + } + + return p, nil +} + +// spearman is the rank correlation between two position vectors. Positions are +// distinct within a vocabulary, so ranking them is just sorting -- no tie +// handling is needed, unlike the frequency ranks in stats.go. +func spearman(a, b []int) float64 { + if len(a) < 2 { + return 0 + } + + return stat.Correlation(ranks(a), ranks(b), nil) +} + +func ranks(v []int) []float64 { + order := make([]int, len(v)) + + for i := range order { + order[i] = i + } + + slices.SortFunc(order, func(x, y int) int { return v[x] - v[y] }) + + r := make([]float64, len(v)) + + for rank, i := range order { + r[i] = float64(rank + 1) + } + + return r +} + +func median(v []int) float64 { + if len(v) == 0 { + return 0 + } + + s := slices.Clone(v) + slices.Sort(s) + + if n := len(s); n%2 == 1 { + return float64(s[n/2]) + } else { + return float64(s[n/2-1]+s[n/2]) / 2 + } +} + +func mean(v []int) float64 { + if len(v) == 0 { + return 0 + } + + sum := 0 + + for _, x := range v { + sum += x + } + + return float64(sum) / float64(len(v)) +} + +func percent(n, total int) string { + if total == 0 { + return "" + } + + return fmt.Sprintf("%.1f", 100*float64(n)/float64(total)) +} + +// group holds the compared sub-words of one class and their positions either +// side, plus how many of them are morphemes. +type group struct { + base, variant []int + shift []int + morphemes int +} + +func (g *group) add(basePos, variantPos int, isMorph bool) { + g.base = append(g.base, basePos) + g.variant = append(g.variant, variantPos) + g.shift = append(g.shift, variantPos-basePos) + + if isMorph { + g.morphemes++ + } +} + +func (g *group) row() []string { + // gained counts sub-words that moved to an earlier merge, i.e. a lower id + gained := 0 + + for _, s := range g.shift { + if s < 0 { + gained++ + } + } + + return []string{ + strconv.Itoa(len(g.shift)), + percent(g.morphemes, len(g.shift)), + fmt.Sprintf("%.4f", spearman(g.base, g.variant)), + fmt.Sprintf("%.1f", median(g.shift)), + fmt.Sprintf("%.1f", mean(g.shift)), + percent(gained, len(g.shift)), + } +} + +// churn is a set of sub-words present on only one side, so it has no +// displacement -- only a size and a morpheme share. +func churn(n, morphemes int) []string { + return []string{strconv.Itoa(n), percent(morphemes, n), "", "", "", ""} +} + +// cohort is a fixed set of sub-word strings, followed across the sweep. The +// sub-words a variant introduces have no baseline position, so they never +// appear in the displacement columns; tracking them from the tokenizer that +// introduced them is the only way to see where they go. +type cohort struct { + tokens []string + origin float64 // median position in the tokenizer that introduced them +} + +func newCohort(tokens []string, at map[string]int) cohort { + return cohort{tokens: tokens, origin: cohortMedian(tokens, at)} +} + +func cohortMedian(tokens []string, at map[string]int) float64 { + present := make([]int, 0, len(tokens)) + + for _, t := range tokens { + if id, ok := at[t]; ok { + present = append(present, id) + } + } + + return median(present) +} + +func (c cohort) row(at map[string]int) []string { + present := 0 + + for _, t := range c.tokens { + if _, ok := at[t]; ok { + present++ + } + } + + med := cohortMedian(c.tokens, at) + + return []string{ + strconv.Itoa(len(c.tokens)), + strconv.Itoa(present), + percent(present, len(c.tokens)), + fmt.Sprintf("%.1f", med), + fmt.Sprintf("%+.1f", med-c.origin), + } +} + +func main() { + dirs := flag.String("dirs", "results/knobloch/minipile_19_ctrl/%s_minipile", + "tokenizer directory, %s is the alignment name") + baseline := flag.String("baseline", "m000", "alignment to compare against") + mode := flag.String("intersect", "all", + "compare on sub-words shared by `all` tokenizers, or only with the `baseline`") + out := flag.String("out", "rank_shift.csv", "output CSV") + track := flag.String("cohort", "", + "follow the sub-words this alignment introduces across the sweep, split by morpheme") + trackOut := flag.String("cohort-out", "rank_cohort.csv", "output CSV for -cohort") + + flag.Parse() + + if *mode != "all" && *mode != "baseline" { + log.Fatalf("unknown -intersect %q", *mode) + } + + sweep := names(*baseline) + + morph, err := knobloch.Morphemes(shelf.Abs(knobloch.SegmentsPath)) + + if err != nil { + log.Fatal(err) + } + + base, err := positions(fmt.Sprintf(*dirs, *baseline)) + + if err != nil { + log.Fatal(err) + } + + log.Printf("baseline %s: %d sub-words, %d known morphemes", *baseline, len(base), len(morph)) + + // with -intersect all the comparison set is fixed before any variant is + // scored, so every row is computed on the same sub-words. Built in its own + // pass to avoid holding the whole sweep's vocabularies in memory at once. + var common map[string]struct{} + + if *mode == "all" { + common = make(map[string]struct{}, len(base)) + + for token := range base { + common[token] = struct{}{} + } + + for _, n := range sweep[1:] { + p, err := positions(fmt.Sprintf(*dirs, n)) + + if err != nil { + log.Fatal(err) + } + + for token := range common { + if _, ok := p[token]; !ok { + delete(common, token) + } + } + } + + log.Printf("intersection over the whole sweep: %d sub-words", len(common)) + } + + // the cohort is fixed before the sweep is scored: the sub-words the chosen + // alignment adds to the baseline, split by whether they are morphemes + var cohorts map[string]cohort + + if *track != "" { + ref, err := positions(fmt.Sprintf(*dirs, *track)) + + if err != nil { + log.Fatal(err) + } + + var morphTokens, otherTokens []string + + for token := range ref { + if _, inBase := base[token]; inBase { + continue + } + + if _, isMorph := morph[token]; isMorph { + morphTokens = append(morphTokens, token) + } else { + otherTokens = append(otherTokens, token) + } + } + + cohorts = map[string]cohort{ + "morpheme": newCohort(morphTokens, ref), + "other": newCohort(otherTokens, ref), + } + + log.Printf("cohort from %s: %d added sub-words, %d morphemes (%s%%)", + *track, len(morphTokens)+len(otherTokens), len(morphTokens), + percent(len(morphTokens), len(morphTokens)+len(otherTokens))) + } + + f, err := os.Create(*out) + + if err != nil { + log.Fatal(err) + } + + defer f.Close() + + w := csv.NewWriter(f) + defer w.Flush() + + header := []string{"tokenizer", "alignment", "inverted", "class", + "n", "morpheme_pct", "spearman", "median_shift", "mean_shift", "gained_pct"} + + if err := w.Write(header); err != nil { + log.Fatal(err) + } + + var cw *csv.Writer + + if cohorts != nil { + cf, err := os.Create(*trackOut) + + if err != nil { + log.Fatal(err) + } + + defer cf.Close() + + cw = csv.NewWriter(cf) + defer cw.Flush() + + if err := cw.Write([]string{"tokenizer", "alignment", "inverted", "class", + "cohort", "present", "present_pct", "median_id", "shift_from_origin"}); err != nil { + log.Fatal(err) + } + } + + for _, n := range sweep[1:] { + variant, err := positions(fmt.Sprintf(*dirs, n)) + + if err != nil { + log.Fatal(err) + } + + var all, morpheme, other group + var dropped, droppedMorph, added, addedMorph int + + for token, basePos := range base { + _, isMorph := morph[token] + variantPos, ok := variant[token] + + if !ok { + dropped++ + + if isMorph { + droppedMorph++ + } + + continue + } + + if common != nil { + if _, keep := common[token]; !keep { + continue + } + } + + all.add(basePos, variantPos, isMorph) + + if isMorph { + morpheme.add(basePos, variantPos, isMorph) + } else { + other.add(basePos, variantPos, isMorph) + } + } + + for token := range variant { + if _, ok := base[token]; !ok { + added++ + + if _, isMorph := morph[token]; isMorph { + addedMorph++ + } + } + } + + alpha, inverted := n[len(n)-3:], n[:2] == "mi" + + a, err := strconv.Atoi(alpha) + + if err != nil { + log.Fatal(err) + } + + rows := []struct { + label string + cells []string + }{ + {"all", all.row()}, + {"morpheme", morpheme.row()}, + {"other", other.row()}, + {"added", churn(added, addedMorph)}, + {"dropped", churn(dropped, droppedMorph)}, + } + + for _, r := range rows { + row := append([]string{ + n, + fmt.Sprintf("%.1f", float64(a)/100), + strconv.FormatBool(inverted), + r.label, + }, r.cells...) + + if err := w.Write(row); err != nil { + log.Fatal(err) + } + } + + if cw != nil { + for _, c := range []string{"morpheme", "other"} { + row := append([]string{ + n, + fmt.Sprintf("%.1f", float64(a)/100), + strconv.FormatBool(inverted), + c, + }, cohorts[c].row(variant)...) + + if err := cw.Write(row); err != nil { + log.Fatal(err) + } + } + } + + log.Printf("%-6s compared %6d rho %.4f morph %+8.1f other %+8.1f "+ + "added %6d (%s%% morph) dropped %6d (%s%% morph)", + n, len(all.shift), spearman(all.base, all.variant), + median(morpheme.shift), median(other.shift), + added, percent(addedMorph, added), dropped, percent(droppedMorph, dropped)) + } + + log.Printf("wrote %s", *out) +} |
