diff options
Diffstat (limited to 'tokenizer/bpe')
| -rw-r--r-- | tokenizer/bpe/validate.go | 43 |
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 } |
