diff options
Diffstat (limited to 'tokenizer')
| -rw-r--r-- | tokenizer/bpe/bytelevel.go | 44 | ||||
| -rw-r--r-- | tokenizer/bpe/cmd/tokenize/main.go | 45 | ||||
| -rw-r--r-- | tokenizer/bpe/split/fsa.go | 230 | ||||
| -rw-r--r-- | tokenizer/bpe/utility.go | 2 |
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) |
