X-Orange's picture
Upload folder using huggingface_hub
6550ac5 verified
Raw
History Blame Contribute Delete
6.4 kB
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]))