diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-09-11 18:39:02 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-09-11 18:39:02 +0200 |
| commit | 9e1b8c4bde9b0263a4c4d2278e3c283ca0eb07f1 (patch) | |
| tree | bf6b604a150a196490b92f89ca40a335aca83d31 /research/knobloch/cmd/windowoverlap/main.go | |
| parent | 75e581b3bc19a73d0c1f48b32f2e58121c9741c1 (diff) | |
Diffstat (limited to 'research/knobloch/cmd/windowoverlap/main.go')
| -rw-r--r-- | research/knobloch/cmd/windowoverlap/main.go | 190 |
1 files changed, 190 insertions, 0 deletions
diff --git a/research/knobloch/cmd/windowoverlap/main.go b/research/knobloch/cmd/windowoverlap/main.go new file mode 100644 index 0000000..3c651d1 --- /dev/null +++ b/research/knobloch/cmd/windowoverlap/main.go @@ -0,0 +1,190 @@ +// Command windowoverlap reports how much the RDD estimation windows of different +// alignment levels contain the same tokens. +// +// The window is a slice of vocabulary-id space, so each alignment puts different +// tokens in it. If the windows overlap little, the bias estimator is reading a +// largely different token population for each tokenizer, which bears on whether +// their estimates are comparable. +// +// Membership depends only on the counterfactual vocabulary and its merges, so +// this runs in seconds and never touches the corpus. +package main + +import ( + "encoding/csv" + "flag" + "fmt" + "log" + "os" + "strconv" + + "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) +} + +func set(v []string) map[string]struct{} { + s := make(map[string]struct{}, len(v)) + + for _, x := range v { + s[x] = struct{}{} + } + + return s +} + +// share of a that is also in b +func overlap(a []string, b map[string]struct{}) float64 { + if len(a) == 0 { + return 0 + } + + n := 0 + + for _, x := range a { + if _, ok := b[x]; ok { + n++ + } + } + + return 100 * float64(n) / float64(len(a)) +} + +type entry struct { + name string + alignment string + observed []string + oov []string +} + +func main() { + ctrl := flag.String("ctrl", "results/knobloch/minipile_19_ctrl/%s_minipile", "counterfactual tokenizer directory, %s is the alignment name") + cutoff := flag.Int("cutoff", 50256, "RDD cutoff") + window := flag.Int("window", 5000, "half-width of the estimation window in token ids") + out := flag.String("out", "window_overlap.csv", "output CSV") + + flag.Parse() + + knobloch.LesciCutoff = *cutoff + knobloch.LesciWindow = *window + + log.Printf("cutoff %d, window ids [%d, %d)", *cutoff, *cutoff-*window, *cutoff+*window) + + var entries []entry + + for _, inv := range []bool{false, true} { + for _, a := range alphas { + n := name(a, inv) + + dir := shelf.Item(fmt.Sprintf(*ctrl, n)) + + tok, err := bpe.NewTokenizerFromFiles( + shelf.Abs(dir+"/vocab.json"), shelf.Abs(dir+"/merges.txt"), + bpe.Config{Recover: false}) + + if err != nil { + log.Fatal(err) + } + + obs, oov := knobloch.WindowTokens(tok) + + entries = append(entries, entry{n, fmt.Sprintf("%.1f", float64(a)/100), obs, oov}) + } + } + + // whole-window membership, ignoring which side of the cutoff a token landed + // on, plus the side-matched version; the gap between them is tokens that are + // in both windows but have crossed the cutoff + all := func(e entry) []string { return append(append([]string{}, e.observed...), e.oov...) } + + interAll := set(all(entries[0])) + interObs := set(entries[0].observed) + interOOV := set(entries[0].oov) + + for _, e := range entries[1:] { + cur := set(all(e)) + + for t := range interAll { + if _, ok := cur[t]; !ok { + delete(interAll, t) + } + } + + cur = set(e.observed) + + for t := range interObs { + if _, ok := cur[t]; !ok { + delete(interObs, t) + } + } + + cur = set(e.oov) + + for t := range interOOV { + if _, ok := cur[t]; !ok { + delete(interOOV, t) + } + } + } + + baseAll := set(all(entries[0])) + baseObs := set(entries[0].observed) + baseOOV := set(entries[0].oov) + + log.Printf("intersection across %d tokenizers: %d whole window, %d observed, %d oov", + len(entries), len(interAll), len(interObs), len(interOOV)) + + file, err := os.Create(*out) + + if err != nil { + log.Fatal(err) + } + + defer file.Close() + + w := csv.NewWriter(file) + + if err := w.Write([]string{ + "tokenizer", "alignment", "cutoff", "window", + "n_window", "n_observed", "n_oov", + "window_shared_base", "window_shared_all", + "observed_shared_base", "oov_shared_base", + "observed_shared_all", "oov_shared_all", + }); err != nil { + log.Fatal(err) + } + + for _, e := range entries { + if err := w.Write([]string{ + e.name, e.alignment, strconv.Itoa(*cutoff), strconv.Itoa(*window), + strconv.Itoa(len(e.observed) + len(e.oov)), + strconv.Itoa(len(e.observed)), strconv.Itoa(len(e.oov)), + strconv.FormatFloat(overlap(all(e), baseAll), 'f', 2, 64), + strconv.FormatFloat(overlap(all(e), interAll), 'f', 2, 64), + strconv.FormatFloat(overlap(e.observed, baseObs), 'f', 2, 64), + strconv.FormatFloat(overlap(e.oov, baseOOV), 'f', 2, 64), + strconv.FormatFloat(overlap(e.observed, interObs), 'f', 2, 64), + strconv.FormatFloat(overlap(e.oov, interOOV), 'f', 2, 64), + }); err != nil { + log.Fatal(err) + } + } + + w.Flush() + + if err := w.Error(); err != nil { + log.Fatal(err) + } + + log.Printf("wrote %s", *out) +} |
