summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-14 19:44:06 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-14 19:44:06 +0200
commitd9bed216c45aa1e0b817f4bccc7969bdac1adbff (patch)
tree2d865fecae1e59d16626893a9a8f1404fff89f81
parent0c97f22069d77cf3cac92d68173e38f83b258a9a (diff)
Fix reachable merges validator
* Previous implementation ignored byte replacements
-rw-r--r--research/lesci/extract.go28
-rw-r--r--tokenizer/bpe/validate.go33
2 files changed, 49 insertions, 12 deletions
diff --git a/research/lesci/extract.go b/research/lesci/extract.go
index fb24316..c7c85e8 100644
--- a/research/lesci/extract.go
+++ b/research/lesci/extract.go
@@ -52,22 +52,34 @@ func Rules(tokenizer llm.Tokenizer, merges [][2]string) (tensor.Dense[int64], []
rules := tensor.NewDense[int64]([]int{len(merges), 3}, nil)
+ atoi := make(map[string]int)
+
+ vocab := bpe.Vocab(tokenizer.(*bpe.Tokenizer))
+
+ for i, token := range vocab {
+ if _, ok := atoi[token]; ok {
+ continue
+ }
+
+ atoi[token] = i
+ }
+
for i, merge := range merges {
if !valid[i] {
continue
}
- a := tokenizer.Tokenize(merge[0])
- b := tokenizer.Tokenize(merge[1])
- c := tokenizer.Tokenize(merge[0] + merge[1])
+ a, x := atoi[merge[0]]
+ b, y := atoi[merge[1]]
+ c, z := atoi[merge[0]+merge[1]]
- if len(a) != 1 || len(b) != 1 || len(c) != 1 {
- panic("unexpected token IDs")
+ if !x || !y || !z {
+ panic("unknown token")
}
- rules.Set([]int{i, 0}, int64(a[0]))
- rules.Set([]int{i, 1}, int64(b[0]))
- rules.Set([]int{i, 2}, int64(c[0]))
+ rules.Set([]int{i, 0}, int64(a))
+ rules.Set([]int{i, 1}, int64(b))
+ rules.Set([]int{i, 2}, int64(c))
}
return rules, valid
diff --git a/tokenizer/bpe/validate.go b/tokenizer/bpe/validate.go
index 84910cf..ffa8e37 100644
--- a/tokenizer/bpe/validate.go
+++ b/tokenizer/bpe/validate.go
@@ -108,14 +108,39 @@ func ByteCoverage(t *Tokenizer) bool {
}
func ReachableMerges(t *Tokenizer, merges [][2]string) []bool {
+ atoi := make(map[string]int)
+
+ vocab := Vocab(t)
+
+ for i, token := range vocab {
+ if _, ok := atoi[token]; ok {
+ continue
+ }
+
+ atoi[token] = i
+ }
+
+ reachable := make(map[string]struct{})
+
+ for _, a := range InitialAlphabet() {
+ if _, ok := atoi[string(a)]; ok {
+ reachable[string(a)] = struct{}{}
+ }
+ }
+
mask := make([]bool, len(merges))
for i, merge := range merges {
- a := t.Tokenize(merge[0])
- b := t.Tokenize(merge[1])
- c := t.Tokenize(merge[0] + merge[1])
+ _, a := reachable[merge[0]]
+ _, b := reachable[merge[1]]
+
+ if !a || !b {
+ continue
+ }
+
+ reachable[merge[0]+merge[1]] = struct{}{}
- mask[i] = len(a) == 1 && len(b) == 1 && len(c) == 1
+ mask[i] = true
}
return mask