openbmb/cpm-bee-5b
7241
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 