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