summaryrefslogtreecommitdiff
path: root/tokenizer/bpe/split/fsa.go
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-03-24 15:37:50 +0100
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-03-24 15:42:10 +0100
commitf796560cb4de0b6de118c00be54c58206bad5b60 (patch)
treeda12930d66ec6511abb723b3ef811a310c07bdb0 /tokenizer/bpe/split/fsa.go
parenta63589d6584548b286825f4fb23202c35382032c (diff)
Expose finite state automaton
Diffstat (limited to 'tokenizer/bpe/split/fsa.go')
-rw-r--r--tokenizer/bpe/split/fsa.go230
1 files changed, 230 insertions, 0 deletions
diff --git a/tokenizer/bpe/split/fsa.go b/tokenizer/bpe/split/fsa.go
new file mode 100644
index 0000000..8c95c89
--- /dev/null
+++ b/tokenizer/bpe/split/fsa.go
@@ -0,0 +1,230 @@
+package split
+
+import (
+ "strings"
+ "unicode"
+ "unicode/utf8"
+)
+
+const (
+ RuneUnicode32 = iota
+ RuneWhitespaceNotUnicode32
+ RuneLetter
+ RuneNumber
+ RuneOther
+ StateInitial
+ StateU32
+ StateWhitespaceNotUnicode32
+ StateLetter
+ StateNumber
+ StateOther
+ StateWhitespaceLookAhead
+)
+
+type FSA struct {
+ state int
+ input []rune
+ static []string
+}
+
+func NewFSA() *FSA {
+ return &FSA{
+ state: StateInitial,
+ input: make([]rune, 0),
+ static: []string{"'s", "'t", "'re", "'m", "'ll", "'d"},
+ }
+}
+
+func (f *FSA) Reset() {
+ f.state = StateInitial
+ f.input = make([]rune, 0)
+}
+
+func (f *FSA) Read(next rune) bool {
+ var r int
+
+ if next == 32 {
+ r = RuneUnicode32
+ } else if unicode.IsSpace(next) {
+ r = RuneWhitespaceNotUnicode32
+ } else if unicode.IsLetter(next) {
+ r = RuneLetter
+ } else if unicode.IsNumber(next) {
+ r = RuneNumber
+ } else {
+ r = RuneOther
+ }
+
+ if f.state == StateInitial {
+ switch r {
+ case RuneUnicode32:
+ f.input = append(f.input, next)
+ f.state = StateU32
+
+ break
+ case RuneWhitespaceNotUnicode32:
+ f.input = append(f.input, next)
+ f.state = StateWhitespaceNotUnicode32
+
+ break
+ case RuneLetter:
+ f.input = append(f.input, next)
+ f.state = StateLetter
+
+ break
+ case RuneNumber:
+ f.input = append(f.input, next)
+ f.state = StateNumber
+
+ break
+ default:
+ f.input = append(f.input, next)
+ f.state = StateOther
+ }
+ } else if f.state == StateU32 {
+ switch r {
+ case RuneUnicode32:
+ f.input = append(f.input, next)
+ f.state = StateWhitespaceLookAhead
+
+ break
+ case RuneWhitespaceNotUnicode32:
+ f.input = append(f.input, next)
+ f.state = StateWhitespaceLookAhead
+
+ break
+ case RuneLetter:
+ f.input = append(f.input, next)
+ f.state = StateLetter
+
+ break
+ case RuneNumber:
+ f.input = append(f.input, next)
+ f.state = StateNumber
+
+ break
+ default:
+ f.input = append(f.input, next)
+ f.state = StateOther
+ }
+ } else if f.state == StateWhitespaceNotUnicode32 {
+ switch r {
+ case RuneUnicode32:
+ f.input = append(f.input, next)
+ f.state = StateWhitespaceLookAhead
+
+ break
+ case RuneWhitespaceNotUnicode32:
+ f.input = append(f.input, next)
+ f.state = StateWhitespaceLookAhead
+
+ break
+ default:
+ return false
+ }
+ } else if f.state == StateNumber {
+ switch r {
+ case RuneNumber:
+ f.input = append(f.input, next)
+ f.state = StateNumber
+
+ break
+ default:
+ return false
+ }
+ } else if f.state == StateLetter {
+ switch r {
+ case RuneLetter:
+ f.input = append(f.input, next)
+ f.state = StateLetter
+
+ break
+ default:
+ return false
+ }
+ } else if f.state == StateOther {
+ switch r {
+ case RuneOther:
+ f.input = append(f.input, next)
+ f.state = StateOther
+
+ break
+ default:
+ return false
+ }
+ } else if f.state == StateWhitespaceLookAhead {
+ switch r {
+ case RuneUnicode32:
+ f.input = append(f.input, next)
+ f.state = StateWhitespaceLookAhead
+
+ break
+ case RuneWhitespaceNotUnicode32:
+ f.input = append(f.input, next)
+ f.state = StateWhitespaceLookAhead
+
+ break
+ default:
+ return false
+ }
+ } else {
+ panic("invalid state")
+ }
+
+ return true
+}
+
+func (f *FSA) FindAll(s string) []string {
+ var findAll func(runes []rune, matches []string) []string
+
+ findAll = func(runes []rune, matches []string) []string {
+ s = string(runes)
+
+ for _, v := range f.static {
+ if strings.HasPrefix(s, v) {
+ matches = append(matches, v)
+ runes = runes[utf8.RuneCountInString(v):]
+
+ if len(runes) == 0 {
+ return matches
+ }
+
+ return findAll(runes, matches)
+ }
+ }
+
+ for i, r := range runes {
+ ok := f.Read(r)
+
+ if !ok {
+ if f.state == StateInitial {
+ return matches
+ }
+
+ if f.state == StateWhitespaceLookAhead {
+ matches = append(matches, string(f.input[:len(f.input)-1]))
+
+ f.Reset()
+
+ return findAll(runes[i-1:], matches)
+ }
+
+ matches = append(matches, string(f.input))
+
+ f.Reset()
+
+ return findAll(runes[i:], matches)
+ }
+ }
+
+ if len(f.input) > 0 {
+ matches = append(matches, string(f.input))
+ }
+
+ return matches
+ }
+
+ defer f.Reset()
+
+ return findAll([]rune(s), make([]string, 0))
+}