| 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' %}<system> {{ message['content'] }} </system>{% elif message['role'] == 'user' %}<user> {{ message['content'] }} </user>{% elif message['role'] == 'assistant' %}<assistant>{{ message['content'] }}</assistant>{% endif %}{% endfor %}""" |
| |
| 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 |
| ) |
| |
| 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])) |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|