CoolFace
Modelpublic

ScalableMath/Lean-STaR-plus

sourceHugging Faceupdated 2y agoView on Hugging Face
2likes16downloads
tokenization_internlm2_fast.py215 linesDownload Raw Back to root
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