CoolFace
Modelpublic

openbmb/cpm-bee-5b

sourceHugging Faceupdated 3y agoView on Hugging Face
7likes241downloads
test_tokenization_cpmbee.py188 linesDownload Raw Back to root
1# coding=utf-82# Copyright 2022 The HuggingFace Inc. team. All rights reserved.3#4# Licensed under the Apache License, Version 2.0 (the "License");5# you may not use this file except in compliance with the License.6# You may obtain a copy of the License at7#8#     http://www.apache.org/licenses/LICENSE-2.09#10# Unless required by applicable law or agreed to in writing, software11# distributed under the License is distributed on an "AS IS" BASIS,12# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.13# See the License for the specific language governing permissions and14# limitations under the License.15""" Testing suite for the PyTorch CpmBee tokenizer. """16 17import os18import unittest19 20from transformers.models.cpmbee.tokenization_cpmbee import VOCAB_FILES_NAMES, CpmBeeTokenizer21from transformers.tokenization_utils import AddedToken22 23from ...test_tokenization_common import TokenizerTesterMixin24 25 26class CPMBeeTokenizationTest(TokenizerTesterMixin, unittest.TestCase):27    tokenizer_class = CpmBeeTokenizer28    test_rust_tokenizer = False29 30    def setUp(self):31        super().setUp()32 33        vocab_tokens = [34            "<d>",35            "</d>",36            "<s>",37            "</s>",38            "</_>",39            "<unk>",40            "<pad>",41            "<mask>",42            "</n>",43            "我",44            "是",45            "C",46            "P",47            "M",48            "B",49            "e",50            "e",51        ]52        self.vocab_file = os.path.join(self.tmpdirname, VOCAB_FILES_NAMES["vocab_file"])53        vocab_tokens = list(set(vocab_tokens))54        with open(self.vocab_file, "w", encoding="utf-8") as vocab_writer:55            vocab_writer.write("".join([x + "\n" for x in vocab_tokens]))56 57    # override test_add_tokens_tokenizer because <...> is special token in CpmBeeTokenizer.58    def test_add_tokens_tokenizer(self):59        tokenizers = self.get_tokenizers(do_lower_case=False)60        for tokenizer in tokenizers:61            with self.subTest(f"{tokenizer.__class__.__name__}"):62                vocab_size = tokenizer.vocab_size63                all_size = len(tokenizer)64 65                self.assertNotEqual(vocab_size, 0)66 67                # We usually have added tokens from the start in tests because our vocab fixtures are68                # smaller than the original vocabs - let's not assert this69                # self.assertEqual(vocab_size, all_size)70 71                new_toks = ["aaaaa bbbbbb", "cccccccccdddddddd"]72                added_toks = tokenizer.add_tokens(new_toks)73                vocab_size_2 = tokenizer.vocab_size74                all_size_2 = len(tokenizer)75 76                self.assertNotEqual(vocab_size_2, 0)77                self.assertEqual(vocab_size, vocab_size_2)78                self.assertEqual(added_toks, len(new_toks))79                self.assertEqual(all_size_2, all_size + len(new_toks))80 81                tokens = tokenizer.encode("aaaaa bbbbbb low cccccccccdddddddd l", add_special_tokens=False)82 83                self.assertGreaterEqual(len(tokens), 4)84                self.assertGreater(tokens[0], tokenizer.vocab_size - 1)85                self.assertGreater(tokens[-2], tokenizer.vocab_size - 1)86 87                new_toks_2 = {"eos_token": ">>>>|||<||<<|<<", "pad_token": "<<<<<|||;;;||;"}88                added_toks_2 = tokenizer.add_special_tokens(new_toks_2)89                vocab_size_3 = tokenizer.vocab_size90                all_size_3 = len(tokenizer)91 92                self.assertNotEqual(vocab_size_3, 0)93                self.assertEqual(vocab_size, vocab_size_3)94                self.assertEqual(added_toks_2, len(new_toks_2))95                self.assertEqual(all_size_3, all_size_2 + len(new_toks_2))96 97                tokens = tokenizer.encode(98                    ">>>>|||<||<<|<< aaaaabbbbbb low cccccccccdddddddd <<<<<|||;;;||; l", add_special_tokens=False99                )100 101                self.assertGreaterEqual(len(tokens), 6)102                self.assertGreater(tokens[0], tokenizer.vocab_size - 1)103                self.assertGreater(tokens[0], tokens[1])104                self.assertGreater(tokens[-2], tokenizer.vocab_size - 1)105                self.assertGreater(tokens[-2], tokens[-3])106                self.assertEqual(tokens[0], tokenizer.eos_token_id)107                self.assertEqual(tokens[-2], tokenizer.pad_token_id)108 109    def test_added_tokens_do_lower_case(self):110        tokenizers = self.get_tokenizers(do_lower_case=True)111        for tokenizer in tokenizers:112            with self.subTest(f"{tokenizer.__class__.__name__}"):113                if not hasattr(tokenizer, "do_lower_case") or not tokenizer.do_lower_case:114                    continue115 116                special_token = tokenizer.all_special_tokens[0]117 118                text = special_token + " aaaaa bbbbbb low cccccccccdddddddd l " + special_token119                text2 = special_token + " AAAAA BBBBBB low CCCCCCCCCDDDDDDDD l " + special_token120 121                toks_before_adding = tokenizer.tokenize(text)  # toks before adding new_toks122 123                new_toks = ["aaaaa bbbbbb", "cccccccccdddddddd", "AAAAA BBBBBB", "CCCCCCCCCDDDDDDDD"]124                added = tokenizer.add_tokens([AddedToken(tok, lstrip=True, rstrip=True) for tok in new_toks])125 126                toks_after_adding = tokenizer.tokenize(text)127                toks_after_adding2 = tokenizer.tokenize(text2)128 129                # Rust tokenizers dont't lowercase added tokens at the time calling `tokenizer.add_tokens`,130                # while python tokenizers do, so new_toks 0 and 2 would be treated as the same, so do new_toks 1 and 3.131                self.assertIn(added, [2, 4])132 133                self.assertListEqual(toks_after_adding, toks_after_adding2)134                self.assertTrue(135                    len(toks_before_adding) > len(toks_after_adding),  # toks_before_adding should be longer136                )137 138                # Check that none of the special tokens are lowercased139                sequence_with_special_tokens = "A " + " yEs ".join(tokenizer.all_special_tokens) + " B"140                # Convert the tokenized list to str as some special tokens are tokenized like normal tokens141                # which have a prefix spacee e.g. the mask token of Albert, and cannot match the original142                # special tokens exactly.143                tokenized_sequence = "".join(tokenizer.tokenize(sequence_with_special_tokens))144 145                for special_token in tokenizer.all_special_tokens:146                    self.assertTrue(special_token in tokenized_sequence)147 148        tokenizers = self.get_tokenizers(do_lower_case=True)149        for tokenizer in tokenizers:150            with self.subTest(f"{tokenizer.__class__.__name__}"):151                if hasattr(tokenizer, "do_lower_case") and tokenizer.do_lower_case:152                    continue153 154                special_token = tokenizer.all_special_tokens[0]155 156                text = special_token + " aaaaa bbbbbb low cccccccccdddddddd l " + special_token157                text2 = special_token + " AAAAA BBBBBB low CCCCCCCCCDDDDDDDD l " + special_token158 159                toks_before_adding = tokenizer.tokenize(text)  # toks before adding new_toks160 161                new_toks = ["aaaaa bbbbbb", "cccccccccdddddddd", "AAAAA BBBBBB", "CCCCCCCCCDDDDDDDD"]162                added = tokenizer.add_tokens([AddedToken(tok, lstrip=True, rstrip=True) for tok in new_toks])163                self.assertIn(added, [2, 4])164 165                toks_after_adding = tokenizer.tokenize(text)166                toks_after_adding2 = tokenizer.tokenize(text2)167 168                self.assertEqual(len(toks_after_adding), len(toks_after_adding2))  # Length should still be the same169                self.assertNotEqual(170                    toks_after_adding[1], toks_after_adding2[1]171                )  # But at least the first non-special tokens should differ172                self.assertTrue(173                    len(toks_before_adding) > len(toks_after_adding),  # toks_before_adding should be longer174                )175 176    def test_pre_tokenization(self):177        tokenizer = CpmBeeTokenizer.from_pretrained("openbmb/cpm-bee-10b")178        texts = {"input": "你好,", "<ans>": ""}179        tokens = tokenizer(texts)180        tokens = tokens["input_ids"][0]181 182        input_tokens = [6, 8, 7, 6, 65678, 7, 6, 10273, 246, 7, 6, 9, 7]183        self.assertListEqual(tokens, input_tokens)184 185        normalized_text = "<s><root></s><s>input</s><s>你好,</s><s><ans></s>"186        reconstructed_text = tokenizer.decode(tokens)187        self.assertEqual(reconstructed_text, normalized_text)188