diff options
Diffstat (limited to 'research/knobloch/cmd/window/main.go')
| -rw-r--r-- | research/knobloch/cmd/window/main.go | 254 |
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) +} |
