summaryrefslogtreecommitdiff
path: root/research/knobloch/cmd/window/main.go
diff options
context:
space:
mode:
Diffstat (limited to 'research/knobloch/cmd/window/main.go')
-rw-r--r--research/knobloch/cmd/window/main.go254
1 files changed, 254 insertions, 0 deletions
diff --git a/research/knobloch/cmd/window/main.go b/research/knobloch/cmd/window/main.go
new file mode 100644
index 0000000..055e3cb
--- /dev/null
+++ b/research/knobloch/cmd/window/main.go
@@ -0,0 +1,254 @@
+// Command window reports how frequent the tokens in the RDD estimation window
+// are, for every alignment level.
+//
+// The setup mirrors lesci. The window and its constituent filter come from the
+// counterfactual tokenizer, which is what lesci builds its rules from. Corpus
+// frequencies come from whichever tokenizer actually produces the token: below
+// the cutoff that is the model's own tokenizer, above it the token does not
+// exist for the model and only the counterfactual produces it.
+//
+// Frequencies are counted over a dictionary, so pointing -dict at another split
+// reports the same statistic over that split.
+package main
+
+import (
+ "encoding/csv"
+ "flag"
+ "fmt"
+ "log"
+ "os"
+ "runtime"
+ "slices"
+ "strconv"
+ "sync"
+
+ "github.com/jonasknobloch/mbpe"
+ "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)
+}
+
+// summary of one group of window tokens
+type summary struct {
+ n int
+ zero int
+ median int
+ q25 int
+ q75 int
+ mean float64
+}
+
+func summarise(freqs []int) summary {
+ s := summary{n: len(freqs)}
+
+ if len(freqs) == 0 {
+ return s
+ }
+
+ slices.Sort(freqs)
+
+ var sum int64
+
+ for _, f := range freqs {
+ sum += int64(f)
+
+ if f == 0 {
+ s.zero++
+ }
+ }
+
+ s.median = freqs[len(freqs)/2]
+ s.q25 = freqs[len(freqs)/4]
+ s.q75 = freqs[3*len(freqs)/4]
+ s.mean = float64(sum) / float64(len(freqs))
+
+ return s
+}
+
+type row struct {
+ name string
+ alignment string
+
+ all summary
+ observed summary
+ oov summary
+}
+
+func encode(dir shelf.Item, items []mbpe.Chunk) ([]int, *bpe.Tokenizer) {
+ tok, err := bpe.NewTokenizerFromFiles(
+ shelf.Abs(dir+"/vocab.json"), shelf.Abs(dir+"/merges.txt"),
+ bpe.Config{Recover: false})
+
+ if err != nil {
+ log.Fatal(err)
+ }
+
+ // the dictionary is already pre-tokenised
+ bpe.MBPE(tok).SetPreTokenizer(&knobloch.NoPreTok{})
+
+ counts := make([]int, len(bpe.Vocab(tok)))
+
+ for _, v := range items {
+ for _, id := range tok.Encode(v.Src()) {
+ counts[id] += v.N()
+ }
+ }
+
+ return counts, tok
+}
+
+func main() {
+ dict := flag.String("dict", "results/knobloch/minipile/dict.txt", "shelf-relative dictionary")
+ model := flag.String("model", "tokenizers/minipile/tokenizer_gpt2_50256_%s_minipile", "model tokenizer directory, %s is the alignment name")
+ 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, i.e. the vocabulary size of the model under test")
+ window := flag.Int("window", 5000, "half-width of the estimation window in token ids")
+ out := flag.String("out", "window_frequencies.csv", "output CSV")
+ workers := flag.Int("workers", runtime.NumCPU(), "parallel encodes")
+
+ flag.Parse()
+
+ // the window selector reads these
+ knobloch.LesciCutoff = *cutoff
+ knobloch.LesciWindow = *window
+
+ d := mbpe.NewDict()
+
+ if err := d.Load(shelf.Abs(shelf.Item(*dict))); err != nil {
+ log.Fatal(err)
+ }
+
+ items := d.Items()
+
+ log.Printf("dictionary: %d pre-token types", len(items))
+ log.Printf("cutoff %d, window ids [%d, %d)", *cutoff, *cutoff-*window, *cutoff+*window)
+ log.Printf("observed side from %s, counterfactual side from %s", *model, *ctrl)
+
+ type job struct {
+ name string
+ alignment string
+ }
+
+ var jobs []job
+
+ for _, inv := range []bool{false, true} {
+ for _, a := range alphas {
+ jobs = append(jobs, job{name(a, inv), fmt.Sprintf("%.1f", float64(a)/100)})
+ }
+ }
+
+ rows := make([]row, len(jobs))
+
+ var wg sync.WaitGroup
+
+ queue := make(chan int)
+
+ for i := 0; i < *workers; i++ {
+ wg.Add(1)
+
+ go func() {
+ defer wg.Done()
+
+ for idx := range queue {
+ j := jobs[idx]
+
+ observed, _ := encode(shelf.Item(fmt.Sprintf(*model, j.name)), items)
+ counterfactual, ctfTok := encode(shelf.Item(fmt.Sprintf(*ctrl, j.name)), items)
+
+ if len(counterfactual) <= *cutoff+*window {
+ log.Fatalf("%s: counterfactual vocabulary %d does not reach the top of the window %d",
+ j.name, len(counterfactual), *cutoff+*window)
+ }
+
+ w := knobloch.WindowFrequencies(observed, counterfactual, ctfTok)
+
+ var all, obs, oov []int
+
+ for _, f := range w {
+ all = append(all, f.Freq)
+
+ if f.OOV {
+ oov = append(oov, f.Freq)
+ } else {
+ obs = append(obs, f.Freq)
+ }
+ }
+
+ rows[idx] = row{
+ name: j.name,
+ alignment: j.alignment,
+ all: summarise(all),
+ observed: summarise(obs),
+ oov: summarise(oov),
+ }
+
+ log.Printf("%s: window %d tokens, median %d (observed %d, oov %d)",
+ j.name, rows[idx].all.n, rows[idx].all.median,
+ rows[idx].observed.median, rows[idx].oov.median)
+ }
+ }()
+ }
+
+ for i := range jobs {
+ queue <- i
+ }
+
+ close(queue)
+ wg.Wait()
+
+ file, err := os.Create(*out)
+
+ if err != nil {
+ log.Fatal(err)
+ }
+
+ defer file.Close()
+
+ w := csv.NewWriter(file)
+
+ header := []string{"tokenizer", "alignment", "cutoff", "window"}
+
+ for _, g := range []string{"", "observed_", "oov_"} {
+ header = append(header,
+ g+"n_tokens", g+"n_zero", g+"median_freq", g+"q25_freq", g+"q75_freq", g+"mean_freq")
+ }
+
+ if err := w.Write(header); err != nil {
+ log.Fatal(err)
+ }
+
+ for _, r := range rows {
+ rec := []string{
+ r.name, r.alignment, strconv.Itoa(*cutoff), strconv.Itoa(*window),
+ }
+
+ for _, s := range []summary{r.all, r.observed, r.oov} {
+ rec = append(rec,
+ strconv.Itoa(s.n), strconv.Itoa(s.zero), strconv.Itoa(s.median),
+ strconv.Itoa(s.q25), strconv.Itoa(s.q75),
+ strconv.FormatFloat(s.mean, 'f', 2, 64))
+ }
+
+ if err := w.Write(rec); err != nil {
+ log.Fatal(err)
+ }
+ }
+
+ w.Flush()
+
+ if err := w.Error(); err != nil {
+ log.Fatal(err)
+ }
+
+ log.Printf("wrote %s", *out)
+}