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