File size: 6,400 Bytes
6550ac5 | 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 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 | from typing import List, Union, Optional, Dict, Any
from transformers import AutoTokenizer
import os
os.environ["TOKENIZERS_PARALLELISM"] = "false"
class Tokenizer:
"""
轻量级封装,统一接口:
encode -> input_ids, attention_mask
decode -> 字符串
其余常用属性直接暴露。
"""
def __init__(self, model_name: str, trust_remote_code: bool = True):
self.tokenizer = AutoTokenizer.from_pretrained(
model_name,
trust_remote_code=trust_remote_code
)
# self.tokenizer.chat_template = """{% for message in messages %}{% if message['role'] == 'system' %}{% if message['content'] == '' %}<system> {{ '你是一个人工智能助手' }} </system>{% else %}<system> {{ message['content'] }} </system>{% endif %}{% elif message['role'] == 'user' %}<user> {{ message['content'] }} </user>{% elif message['role'] == 'assistant' %}<assistant>{{ message['content'] }}</assistant>{% endif %}{% endfor %}"""
self.tokenizer.chat_template = """{% for message in messages %}{% if message['role'] == 'system' %}<system> {{ message['content'] }} </system>{% elif message['role'] == 'user' %}<user> {{ message['content'] }} </user>{% elif message['role'] == 'assistant' %}<assistant>{{ message['content'] }}</assistant>{% endif %}{% endfor %}"""
# 如果 pad_token 不存在,统一用 eos_token 代替
if self.tokenizer.pad_token is None:
self.tokenizer.pad_token = self.tokenizer.eos_token
def encode(
self,
text: Union[str, List[str]],
max_length: Optional[int] = None,
padding: Optional[str] = "do_not_pad",
truncation: bool = False,
return_tensors: Optional[str] = None,
add_special_tokens: bool = False,
) -> Dict[str, Any]:
encoded = self.tokenizer(
text,
max_length=max_length,
padding=padding,
truncation=truncation,
return_tensors=return_tensors,
add_special_tokens=add_special_tokens
)
return encoded
def encode_chat(
self,
text: List[Dict[str, Any]],
) -> Dict[str, Any]:
encoded = self.tokenizer.apply_chat_template(
text,
tokenize=False
)
return encoded
def decode(
self,
token_ids: Union[List[int], List[List[int]]],
skip_special_tokens: bool = False,
clean_up_tokenization_spaces: bool = False
) -> Union[str, List[str]]:
if isinstance(token_ids[0], int): # 单条
return self.tokenizer.decode(
token_ids,
skip_special_tokens=skip_special_tokens,
clean_up_tokenization_spaces=clean_up_tokenization_spaces
)
# batch
return self.tokenizer.batch_decode(
token_ids,
skip_special_tokens=skip_special_tokens,
clean_up_tokenization_spaces=clean_up_tokenization_spaces
)
@property
def vocab_size(self) -> int:
return len(self.tokenizer)
@property
def pad_token_id(self) -> int:
return self.tokenizer.pad_token_id
@property
def eos_token_id(self) -> int:
return self.tokenizer.eos_token_id
@property
def bos_token_id(self) -> int:
return self.tokenizer.bos_token_id
@property
def unk_token_id(self) -> int:
return self.tokenizer.unk_token_id
# 兼容旧接口
def get_vocab_size(self) -> int:
return self.vocab_size
def get_eot_token(self) -> int:
"""End-Of-Text token,与 GPT 系列保持一致命名。"""
return self.eos_token_id
if __name__ == '__main__':
tokenizer = Tokenizer("./llama2_tokenizer_32k")
print(tokenizer.get_eot_token())
encode = tokenizer.encode_chat([{"role": "system", "content": ""}, {"role": "user", "content": "请告诉我在中国古代的“四大发明”是什么?"}, {"role": "assistant", "content": "中国古代的“四大发明”是指造纸术、印刷术、火药和指南针。这四项发明对世界文明的发展产生了深远的影响:\n\n1. **造纸术**:据史书记载,东汉时期的蔡伦改进了造纸工艺,使得纸张的生产更加便捷、成本更低,质量也更加优良。这一发明极大地促进了文化的传播和保存。\n\n2. **印刷术**:中国古代的印刷术最早可以追溯到唐代的雕版印刷,到了宋代发展出了活字印刷技术。印刷术的发明极大地促进了知识的传播,降低了书籍的成本,对教育和文化的发展起到了重要的推动作用。\n\n3. **火药**:火药最初是在唐代被发现的,最初可能用于医疗或炼金术。到了宋代,人们开始将火药用于军事目的,发明了各种火器。火药的发明改变了战争的面貌,对世界军事技术的发展产生了重大影响。\n\n4. **指南针**:指南针最初在中国被用于风水测量,后来逐渐发展成为航海导航的重要工具。这一发明极大地促进了海上航行技术的发展,对于新航路的开辟和世界地理大发现起到了关键作用。\n\n这四项发明不仅在中国历史上占有重要地位,而且对全世界的科技进步和文明发展都产生了深远的影响。"}])
print(type(encode))
print("orin length", len(encode))
encode = tokenizer.encode(str(encode), max_length=None, truncation=True, padding="do_not_pad")
for token in encode['input_ids']:
print(tokenizer.decode([token]))
decode = tokenizer.decode([encode['input_ids']])
print("token lens: ", len(encode['input_ids']))
print("encode: ", encode)
print("att_mask: ", encode['attention_mask'])
print("decode: ", decode)
print("vocab_size", tokenizer.get_vocab_size())
print("eot", tokenizer.get_eot_token(), tokenizer.decode([tokenizer.get_eot_token()]))
print("eos", tokenizer.eos_token_id, tokenizer.decode([tokenizer.eos_token_id]))
print("unk", tokenizer.unk_token_id, tokenizer.decode([tokenizer.unk_token_id]))
print("pad", tokenizer.pad_token_id, tokenizer.decode([tokenizer.pad_token_id]))
print("bos", tokenizer.bos_token_id, tokenizer.decode([tokenizer.bos_token_id]))
for i in range(10):
print(i, tokenizer.decode([i]))
|