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