summaryrefslogtreecommitdiff
path: root/tokenizer/bpe/validate.go
diff options
context:
space:
mode:
Diffstat (limited to 'tokenizer/bpe/validate.go')
-rw-r--r--tokenizer/bpe/validate.go108
1 files changed, 108 insertions, 0 deletions
diff --git a/tokenizer/bpe/validate.go b/tokenizer/bpe/validate.go
new file mode 100644
index 0000000..f9bd265
--- /dev/null
+++ b/tokenizer/bpe/validate.go
@@ -0,0 +1,108 @@
+package bpe
+
+import (
+ "fmt"
+ "slices"
+
+ "github.com/jonasknobloch/mbpe"
+)
+
+func InitialAlphabet() []rune {
+ alphabet := make([]rune, 256)
+
+ bc := mbpe.BytesChar
+
+ for i := 0; i < 256; i++ {
+ b := uint8(i)
+
+ runes := []rune(bc[b])
+
+ if len(runes) != 1 {
+ panic("unexpected replacement")
+ }
+
+ alphabet[i] = runes[0]
+ }
+
+ slices.Sort(alphabet)
+
+ return alphabet
+}
+
+func UnknownRunes(t *Tokenizer, s string) []rune {
+ atoi := make(map[string]int)
+
+ vocab := t.mbpe.Model().(*mbpe.MBPE).Vocab()
+
+ for i, token := range vocab {
+ if _, ok := atoi[token]; ok {
+ continue
+ }
+
+ atoi[token] = i
+ }
+
+ unknown := make(map[rune]struct{})
+
+ chunks := t.mbpe.PreTokenizer().PreTokenize(s)
+
+ for _, chunk := range chunks {
+ for _, r := range chunk {
+ i, ok := atoi[string(r)]
+
+ if !ok {
+ fmt.Printf("%d %s %v not in vocabulary\n", i, string(r), []byte(string(r)))
+
+ if _, ok := unknown[r]; ok {
+ continue
+ }
+
+ unknown[r] = struct{}{}
+ }
+ }
+ }
+
+ result := make([]rune, 0, len(unknown))
+
+ for r := range unknown {
+ result = append(result, r)
+ }
+
+ slices.Sort(result)
+
+ return result
+}
+
+func ByteCoverage(t *Tokenizer) bool {
+ atoi := make(map[string]int)
+
+ vocab := t.mbpe.Model().(*mbpe.MBPE).Vocab()
+
+ for i, token := range vocab {
+ if _, ok := atoi[token]; ok {
+ continue
+ }
+
+ atoi[token] = i
+ }
+
+ bc := mbpe.BytesChar
+
+ covered := true
+
+ for i := 0; i < 256; i++ {
+ c, ok := bc[byte(i)]
+
+ if !ok {
+ panic("not in replacement table")
+ }
+
+ if _, ok := atoi[c]; !ok {
+ fmt.Printf("%d %s %v not in vocabulary\n", i, c, []byte(c))
+
+ covered = false
+ }
+ }
+
+ return covered
+}