summaryrefslogtreecommitdiff
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
parenta63589d6584548b286825f4fb23202c35382032c (diff)
Expose finite state automaton
-rw-r--r--tokenizer/bpe/bytelevel.go44
-rw-r--r--tokenizer/bpe/cmd/tokenize/main.go45
-rw-r--r--tokenizer/bpe/split/fsa.go230
-rw-r--r--tokenizer/bpe/utility.go2
4 files changed, 320 insertions, 1 deletions
diff --git a/tokenizer/bpe/bytelevel.go b/tokenizer/bpe/bytelevel.go
new file mode 100644
index 0000000..4059d0d
--- /dev/null
+++ b/tokenizer/bpe/bytelevel.go
@@ -0,0 +1,44 @@
+package bpe
+
+import (
+ "github.com/jonasknobloch/mbpe"
+ "go.jknobloc.com/x/tokenizer/bpe/split"
+)
+
+type ByteLevel struct {
+ *mbpe.ByteLevel
+ addPrefixSpace bool
+ fsa *split.FSA
+}
+
+func NewByteLevel(addPrefixSpace bool) *ByteLevel {
+ return &ByteLevel{
+ ByteLevel: mbpe.NewByteLevel(addPrefixSpace),
+ addPrefixSpace: addPrefixSpace,
+ fsa: split.NewFSA(),
+ }
+}
+
+func (p *ByteLevel) PreTokenize(phrase string) []string {
+ if phrase == "" {
+ return []string{}
+ }
+
+ if p.addPrefixSpace && phrase[0] != ' ' {
+ phrase = " " + phrase
+ }
+
+ compounds := p.fsa.FindAll(phrase)
+
+ for i, compound := range compounds {
+ r := ""
+
+ for _, b := range []byte(compound) {
+ r += mbpe.BytesChar[b]
+ }
+
+ compounds[i] = r
+ }
+
+ return compounds
+}
diff --git a/tokenizer/bpe/cmd/tokenize/main.go b/tokenizer/bpe/cmd/tokenize/main.go
index aaf10e7..1b0b868 100644
--- a/tokenizer/bpe/cmd/tokenize/main.go
+++ b/tokenizer/bpe/cmd/tokenize/main.go
@@ -1,7 +1,11 @@
package main
import (
+ "flag"
"log"
+ "os"
+ "runtime"
+ "runtime/pprof"
"sync/atomic"
"time"
@@ -11,7 +15,30 @@ import (
"go.jknobloc.com/x/tui"
)
+var cpuprofile = flag.String("cpuprofile", "", "write cpu profile to `file`")
+var memprofile = flag.String("memprofile", "", "write memory profile to `file`")
+
func main() {
+ flag.Parse()
+
+ if *cpuprofile != "" {
+ var file *os.File
+
+ if f, err := os.Create(*cpuprofile); err != nil {
+ log.Fatal("could not create CPU profile: ", err)
+ } else {
+ file = f
+
+ defer file.Close()
+ }
+
+ if err := pprof.StartCPUProfile(file); err != nil {
+ log.Fatal("could not start CPU profile: ", err)
+ }
+
+ defer pprof.StopCPUProfile()
+ }
+
reader := data()
t := tokenizer()
@@ -41,6 +68,24 @@ func main() {
n++
}
+
+ if *memprofile != "" {
+ var file *os.File
+
+ if f, err := os.Create(*memprofile); err != nil {
+ log.Fatal("could not create memory profile: ", err)
+ } else {
+ file = f
+
+ defer file.Close()
+ }
+
+ runtime.GC()
+
+ if err := pprof.Lookup("allocs").WriteTo(file, 0); err != nil {
+ log.Fatal("could not write memory profile: ", err)
+ }
+ }
}
func data() *dataset.ParquetReader {
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))
+}
diff --git a/tokenizer/bpe/utility.go b/tokenizer/bpe/utility.go
index 8ad0ca2..f47c073 100644
--- a/tokenizer/bpe/utility.go
+++ b/tokenizer/bpe/utility.go
@@ -11,7 +11,7 @@ func NewTokenizerFromFiles(vocab, merges string) (*Tokenizer, error) {
tokenizer := mbpe.NewTokenizer(model)
- pre := mbpe.NewByteLevel(false)
+ pre := NewByteLevel(false)
tokenizer.SetPreTokenizer(pre)
tokenizer.SetDecoder(pre)