package knobloch import ( "fmt" "image/color" "go.jknobloc.com/x/tokenizer/bpe" "gonum.org/v1/plot" "gonum.org/v1/plot/plotter" "gonum.org/v1/plot/vg" ) // The vocabulary cutoff the lesci experiment splits on, and the window either // side of it that the figure covers. var ( LesciCutoff = 32768 LesciWindow = 5000 ) // lesciMask selects the vocab ids inside the window around the cutoff, dropping // every token that is itself a constituent of another merge rule in that window. // // This is lesci.Window followed by lesci.Filter (research/lesci/lesci.go), // reimplemented rather than imported: lesci is a separate module and its version // works on a tensor.Dense[int64] merge table we would have to build anyway. The // behaviour is mirrored exactly so these figures compare against that chapter. // // That includes lesci.Filter only dropping constituents below the cutoff, which // makes the filter asymmetric: above the cutoff nothing is dropped. lesci.OutOfVocab // is not applied, since keeping only the tokens above the cutoff would empty // half of this figure. func lesciMask(n int, t *bpe.Tokenizer) []bool { atoi := bpe.Atoi(t) lo := int64(LesciCutoff - LesciWindow) hi := int64(LesciCutoff + LesciWindow) inWindow := func(id int64) bool { return id >= lo && id < hi } // every token used to build another token whose result lands in the window constituent := make(map[int64]struct{}) for _, merge := range bpe.Merges(t) { c, ok := atoi[merge[0]+merge[1]] if !ok || !inWindow(c) { continue } if a, ok := atoi[merge[0]]; ok { constituent[a] = struct{}{} } if b, ok := atoi[merge[1]]; ok { constituent[b] = struct{}{} } } mask := make([]bool, n) for id := range mask { if !inWindow(int64(id)) { continue } // as in lesci.Filter: only in-vocab tokens are dropped if _, ok := constituent[int64(id)]; ok && int64(id) < int64(LesciCutoff) { continue } mask[id] = true } return mask } // WindowTokens returns the token strings that make up the RDD estimation window, // split by side of the cutoff. Membership depends only on the counterfactual // vocabulary and its merges, so this needs no corpus pass. func WindowTokens(ctf *bpe.Tokenizer) (observed, oov []string) { itoa := bpe.Itoa(ctf) for id, keep := range lesciMask(len(bpe.Vocab(ctf)), ctf) { if !keep { continue } if id >= LesciCutoff { oov = append(oov, itoa[int64(id)]) } else { observed = append(observed, itoa[int64(id)]) } } return observed, oov } // WindowFrequency is one token of the RDD estimation window, with the corpus // frequency the tokenization-bias estimator would read for it. type WindowFrequency struct { ID int Freq int // below the cutoff the token exists in the model's own vocabulary and its // frequency is read from that tokenizer; at or above it the token is // out-of-vocabulary and only the counterfactual tokenizer produces it OOV bool } // WindowFrequencies returns the tokens of the RDD estimation window with their // corpus frequencies, mirroring how lesci reads the two sides of the cutoff. // // The window and its constituent filter are taken from the counterfactual // tokenizer, as in lesci, which builds its rules from bpe.Merges of the // counterfactual. Frequencies come from whichever tokenizer actually produces // the token: the model's own tokenizer below the cutoff, the counterfactual // above it, where the token does not exist for the model at all. // // observed is indexed by the model tokenizer's vocab ids and counterfactual by // the counterfactual's. The two must agree on ids below the cutoff, which holds // when the smaller vocabulary is an id-preserving prefix of the larger. func WindowFrequencies(observed, counterfactual []int, ctf *bpe.Tokenizer) []WindowFrequency { mask := lesciMask(len(counterfactual), ctf) out := make([]WindowFrequency, 0, 2*LesciWindow) for id, keep := range mask { if !keep { continue } w := WindowFrequency{ID: id, OOV: id >= LesciCutoff} if w.OOV { w.Freq = counterfactual[id] } else if id < len(observed) { w.Freq = observed[id] } out = append(out, w) } return out } // plotMorphemeFractionLesci is the fraction overview restricted to the lesci // window. The axis is the token id rather than the frequency rank, since the // cutoff is a position in merge order and has no meaning on a rank axis. func plotMorphemeFractionLesci(m []int, shared, unshared []bool, t *bpe.Tokenizer, out string) error { morph, err := isMorpheme(t) if err != nil { return err } lesci := lesciMask(len(m), t) var series []fractionSeries for _, c := range []struct { keep []bool color color.NRGBA label string }{ {nil, color.NRGBA{R: 90, G: 90, B: 90, A: 255}, "All"}, {shared, opaque(colorOther), "Shared"}, {unshared, opaque(colorMorpheme), "Unshared"}, } { segments, overall, shown := idFractionCurve(m, and(lesci, c.keep), morph, LesciCutoff-LesciWindow, LesciCutoff+LesciWindow) if shown == 0 { continue } series = append(series, fractionSeries{ segments: segments, overall: overall, color: c.color, label: fmt.Sprintf("%s (%.1f%%, n=%d)", c.label, overall, shown), }) } if len(series) == 0 { return fmt.Errorf("no tokens selected") } return renderMorphemeFractionLesci(series, out) } // and intersects two masks; a nil second mask leaves the first untouched. func and(a, b []bool) []bool { if b == nil { return a } r := make([]bool, len(a)) for i := range a { r[i] = a[i] && b[i] } return r } func renderMorphemeFractionLesci(series []fractionSeries, out string) error { p := plot.New() p.X.Label.Text = "Token ID" p.Y.Label.Text = "Morphemes in window (%)" p.X.Min = float64(LesciCutoff - LesciWindow) p.X.Max = float64(LesciCutoff + LesciWindow) p.Y.Min = 0 p.Y.Max = 100 p.Add(plotter.NewGrid()) cutoff, err := plotter.NewLine(plotter.XYs{ {X: float64(LesciCutoff), Y: 0}, {X: float64(LesciCutoff), Y: 100}, }) if err != nil { return err } cutoff.Color = color.NRGBA{R: 131, G: 131, B: 131, A: 255} cutoff.Width = vg.Points(1) cutoff.Dashes = []vg.Length{vg.Points(4), vg.Points(3)} p.Add(cutoff) p.Legend.Add(fmt.Sprintf("Cutoff (%d)", LesciCutoff), cutoff) for _, s := range series { for i, pts := range s.segments { line, err := plotter.NewLine(pts) if err != nil { return err } line.Color = s.color line.Width = vg.Points(1.5) p.Add(line) if i == 0 { p.Legend.Add(s.label, line) } } } p.Legend.Top = true p.Legend.Padding = vg.Points(4) return p.Save(12*vg.Inch, 6*vg.Inch, out) }