ScalableMath/Lean-STaR-plus
216
1# coding=utf-82# Copyright (c) The InternLM team and The HuggingFace Inc. team. All rights reserved.3#4# This code is based on transformers/src/transformers/models/llama/tokenization_llama_fast.py5#6# Licensed under the Apache License, Version 2.0 (the "License");7# you may not use this file except in compliance with the License.8# You may obtain a copy of the License at9#10# http://www.apache.org/licenses/LICENSE-2.011#12# Unless required by applicable law or agreed to in writing, software13# distributed under the License is distributed on an "AS IS" BASIS,14# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.15# See the License for the specific language governing permissions and16# limitations under the License.17 18"""Tokenization Fast class for InternLM."""19import os20from shutil import copyfile21from typing import Any, Dict, Optional, Tuple22 23from tokenizers import processors, decoders, Tokenizer, normalizers24from tokenizers.models import BPE25 26from transformers.tokenization_utils_fast import PreTrainedTokenizerFast27from transformers.utils import logging28 29from transformers.convert_slow_tokenizer import (30 SLOW_TO_FAST_CONVERTERS,31 SpmConverter,32 SentencePieceExtractor,33)34 35from .tokenization_internlm2 import InternLM2Tokenizer36 37logger = logging.get_logger(__name__)38 39VOCAB_FILES_NAMES = {"vocab_file": "./tokenizer.model"}40 41# Modified from transformers.convert_slow_tokenizer.LlamaConverter42class InternLM2Converter(SpmConverter):43 handle_byte_fallback = True44 45 def vocab(self, proto):46 vocab = [47 ("<unk>", 0.0),48 ("<s>", 0.0),49 ("</s>", 0.0),50 ]51 vocab += [(piece.piece, piece.score) for piece in proto.pieces[3:]]52 return vocab53 54 def unk_id(self, proto):55 unk_id = 056 return unk_id57 58 def decoder(self, replacement, add_prefix_space):59 decoders_sequence = [60 decoders.Replace("▁", " "),61 decoders.ByteFallback(),62 decoders.Fuse(),63 ]64 if self.proto.normalizer_spec.add_dummy_prefix:65 decoders_sequence.append(decoders.Strip(content=" ", left=1))66 return decoders.Sequence(decoders_sequence)67 68 def tokenizer(self, proto):69 model_type = proto.trainer_spec.model_type70 vocab_scores = self.vocab(proto)71 # special tokens72 added_tokens = self.original_tokenizer.added_tokens_decoder73 for i in range(len(vocab_scores)):74 piece, score = vocab_scores[i]75 if i in added_tokens:76 vocab_scores[i] = (added_tokens[i].content, score)77 if model_type == 1:78 raise RuntimeError("InternLM2 is supposed to be a BPE model!")79 80 elif model_type == 2:81 _, merges = SentencePieceExtractor(self.original_tokenizer.vocab_file).extract(vocab_scores)82 bpe_vocab = {word: i for i, (word, _score) in enumerate(vocab_scores)}83 tokenizer = Tokenizer(84 BPE(bpe_vocab, merges, unk_token=proto.trainer_spec.unk_piece, fuse_unk=True, byte_fallback=True)85 )86 tokenizer.add_special_tokens(87 [ added_token for index, added_token in added_tokens.items()]88 )89 else:90 raise Exception(91 "You're trying to run a `Unigram` model but you're file was trained with a different algorithm"92 )93 94 return tokenizer95 96 def normalizer(self, proto):97 normalizers_list = []98 if proto.normalizer_spec.add_dummy_prefix:99 normalizers_list.append(normalizers.Prepend(prepend="▁"))100 normalizers_list.append(normalizers.Replace(pattern=" ", content="▁"))101 return normalizers.Sequence(normalizers_list)102 103 def pre_tokenizer(self, replacement, add_prefix_space):104 return None105 106SLOW_TO_FAST_CONVERTERS["InternLM2Tokenizer"] = InternLM2Converter107 108 109# Modified from transformers.model.llama.tokenization_llama_fast.LlamaTokenizerFast -> InternLM2TokenizerFast110class InternLM2TokenizerFast(PreTrainedTokenizerFast):111 vocab_files_names = VOCAB_FILES_NAMES112 slow_tokenizer_class = InternLM2Tokenizer113 padding_side = "left"114 model_input_names = ["input_ids", "attention_mask"]115 _auto_class = "AutoTokenizer"116 117 def __init__(118 self,119 vocab_file,120 unk_token="<unk>",121 bos_token="<s>",122 eos_token="</s>",123 pad_token="</s>",124 sp_model_kwargs: Optional[Dict[str, Any]] = None,125 add_bos_token=True,126 add_eos_token=False,127 decode_with_prefix_space=False,128 clean_up_tokenization_spaces=False,129 **kwargs,130 ):131 super().__init__(132 vocab_file=vocab_file,133 unk_token=unk_token,134 bos_token=bos_token,135 eos_token=eos_token,136 pad_token=pad_token,137 sp_model_kwargs=sp_model_kwargs,138 add_bos_token=add_bos_token,139 add_eos_token=add_eos_token,140 decode_with_prefix_space=decode_with_prefix_space,141 clean_up_tokenization_spaces=clean_up_tokenization_spaces,142 **kwargs,143 )144 self._add_bos_token = add_bos_token145 self._add_eos_token = add_eos_token146 self.update_post_processor()147 self.vocab_file = vocab_file148 149 @property150 def can_save_slow_tokenizer(self) -> bool:151 return os.path.isfile(self.vocab_file) if self.vocab_file else False152 153 def update_post_processor(self):154 """155 Updates the underlying post processor with the current `bos_token` and `eos_token`.156 """157 bos = self.bos_token158 bos_token_id = self.bos_token_id159 if bos is None and self.add_bos_token:160 raise ValueError("add_bos_token = True but bos_token = None")161 162 eos = self.eos_token163 eos_token_id = self.eos_token_id164 if eos is None and self.add_eos_token:165 raise ValueError("add_eos_token = True but eos_token = None")166 167 single = f"{(bos+':0 ') if self.add_bos_token else ''}$A:0{(' '+eos+':0') if self.add_eos_token else ''}"168 pair = f"{single}{(' '+bos+':1') if self.add_bos_token else ''} $B:1{(' '+eos+':1') if self.add_eos_token else ''}"169 170 special_tokens = []171 if self.add_bos_token:172 special_tokens.append((bos, bos_token_id))173 if self.add_eos_token:174 special_tokens.append((eos, eos_token_id))175 self._tokenizer.post_processor = processors.TemplateProcessing(176 single=single, pair=pair, special_tokens=special_tokens177 )178 179 @property180 def add_eos_token(self):181 return self._add_eos_token182 183 @property184 def add_bos_token(self):185 return self._add_bos_token186 187 @add_eos_token.setter188 def add_eos_token(self, value):189 self._add_eos_token = value190 self.update_post_processor()191 192 @add_bos_token.setter193 def add_bos_token(self, value):194 self._add_bos_token = value195 self.update_post_processor()196 197 def save_vocabulary(self, save_directory: str, filename_prefix: Optional[str] = None) -> Tuple[str]:198 if not self.can_save_slow_tokenizer:199 raise ValueError(200 "Your fast tokenizer does not have the necessary information to save the vocabulary for a slow "201 "tokenizer."202 )203 204 if not os.path.isdir(save_directory):205 logger.error(f"Vocabulary path ({save_directory}) should be a directory")206 return207 out_vocab_file = os.path.join(208 save_directory, (filename_prefix + "-" if filename_prefix else "") + VOCAB_FILES_NAMES["vocab_file"]209 )210 211 if os.path.abspath(self.vocab_file) != os.path.abspath(out_vocab_file):212 copyfile(self.vocab_file, out_vocab_file)213 214 return (out_vocab_file,)215 