summaryrefslogtreecommitdiff
path: root/tokenizer/bpe/tokenizer.go
blob: 7e2459a19a308909120ee418ccd8e7b34cf8e093 (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"

type Tokenizer struct {
	mbpe   *mbpe.Tokenizer
	config Config
}

func NewTokenizer(mbpe *mbpe.Tokenizer, cfg Config) *Tokenizer {
	return &Tokenizer{
		mbpe:   mbpe,
		config: cfg,
	}
}

func (t *Tokenizer) Encode(s string) []int {
	return t.mbpe.Tokenize(s)
}

func (t *Tokenizer) Decode(ids []int) string {
	d := t.mbpe.Decoder()

	m, ok := t.mbpe.Model().(*mbpe.MBPE)

	if !ok {
		panic("unsupported model")
	}

	return d.Decode(m.ToString(ids))
}

func (t *Tokenizer) Tokenize(s string) (ids []int) {
	if t.config.Recover {
		defer func() {
			if r := recover(); r != nil {
				// defer UnknownRunes(t, s)

				ids = []int{}
			}
		}()
	}

	return t.Encode(s)
}