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