summaryrefslogtreecommitdiff
path: root/tokenizer/bpe/bytelevel.go
blob: dfecb53ab0d39dcd72baf10c6a7b4aeb3661b4ae (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
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
}