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.go43
1 files changed, 36 insertions, 7 deletions
diff --git a/tokenizer/bpe/validate.go b/tokenizer/bpe/validate.go
index 9c7a17b..083a418 100644
--- a/tokenizer/bpe/validate.go
+++ b/tokenizer/bpe/validate.go
@@ -87,7 +87,7 @@ func ByteCoverage(t *Tokenizer) bool {
return covered
}
-func ReachableMerges(t *Tokenizer, merges [][2]string) []bool {
+func reachableTokens(t *Tokenizer) map[string]struct{} {
atoi := Atoi(t)
reachable := make(map[string]struct{})
@@ -98,17 +98,46 @@ func ReachableMerges(t *Tokenizer, merges [][2]string) []bool {
}
}
- mask := make([]bool, len(merges))
-
- for i, merge := range merges {
- _, a := reachable[merge[0]]
- _, b := reachable[merge[1]]
+ for _, merge := range Merges(t) {
+ if _, ok := reachable[merge[0]]; !ok {
+ continue
+ }
- if !a || !b {
+ if _, ok := reachable[merge[1]]; !ok {
continue
}
reachable[merge[0]+merge[1]] = struct{}{}
+ }
+
+ return reachable
+}
+
+func ReachableTokens(t *Tokenizer, vocab []string) []bool {
+ reachable := reachableTokens(t)
+
+ mask := make([]bool, len(vocab))
+
+ for i, token := range vocab {
+ if _, ok := reachable[token]; !ok {
+ continue
+ }
+
+ mask[i] = true
+ }
+
+ return mask
+}
+
+func ReachableMerges(t *Tokenizer, merges [][2]string) []bool {
+ reachable := reachableTokens(t)
+
+ mask := make([]bool, len(merges))
+
+ for i, merge := range merges {
+ if _, ok := reachable[merge[0]+merge[1]]; !ok {
+ continue
+ }
mask[i] = true
}