-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtokenizer.py
More file actions
34 lines (24 loc) · 1.02 KB
/
Copy pathtokenizer.py
File metadata and controls
34 lines (24 loc) · 1.02 KB
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
from typing import List
import tiktoken
import os
os.environ["TRANSFORMERS_VERBOSITY"] = "error"
os.environ["TRANSFORMERS_NO_ADVISORY_WARNINGS"] = "1"
from transformers import AutoTokenizer
from contracts import Provider
class Tokenizer:
def __init__(self, model: str, provider: Provider):
self.provider = provider
if provider == Provider.OPENAI:
self.tiktoken_tokenizer = tiktoken.encoding_for_model(model)
else:
self.transformers_tokenizer = AutoTokenizer.from_pretrained(model)
def encode(self, text: str) -> List[int]:
if self.provider == Provider.OPENAI:
return self.tiktoken_tokenizer.encode(text)
else:
return self.transformers_tokenizer.encode(text, add_special_tokens=False)
def decode(self, tokens: List[int]) -> str:
if self.provider == Provider.OPENAI:
return self.tiktoken_tokenizer.decode(tokens)
else:
return self.transformers_tokenizer.decode(tokens, add_special_tokens=False)