summaryrefslogtreecommitdiff
path: root/research/knobloch/lesci.go
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-09-11 18:39:02 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-09-11 18:39:02 +0200
commit9e1b8c4bde9b0263a4c4d2278e3c283ca0eb07f1 (patch)
treebf6b604a150a196490b92f89ca40a335aca83d31 /research/knobloch/lesci.go
parent75e581b3bc19a73d0c1f48b32f2e58121c9741c1 (diff)
Diffstat (limited to 'research/knobloch/lesci.go')
-rw-r--r--research/knobloch/lesci.go260
1 files changed, 260 insertions, 0 deletions
diff --git a/research/knobloch/lesci.go b/research/knobloch/lesci.go
new file mode 100644
index 0000000..39f182c
--- /dev/null
+++ b/research/knobloch/lesci.go
@@ -0,0 +1,260 @@
+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)
+}