chendl/compositional_test
1
1import json2import os3import shutil4import sys5import tempfile6import unittest7from unittest import TestCase8from unittest.mock import patch9 10import faiss11import numpy as np12from datasets import Dataset13 14from transformers import BartConfig, BartTokenizer, DPRConfig, DPRQuestionEncoderTokenizer, RagConfig15from transformers.file_utils import is_datasets_available, is_faiss_available, is_psutil_available, is_torch_available16from transformers.integrations import is_ray_available17from transformers.models.bert.tokenization_bert import VOCAB_FILES_NAMES as DPR_VOCAB_FILES_NAMES18from transformers.models.rag.retrieval_rag import CustomHFIndex, RagRetriever19from transformers.models.roberta.tokenization_roberta import VOCAB_FILES_NAMES as BART_VOCAB_FILES_NAMES20from transformers.testing_utils import require_ray21 22 23sys.path.append(os.path.join(os.getcwd())) # noqa: E402 # noqa: E402 # isort:skip24 25if is_torch_available():26 from distributed_pytorch_retriever import RagPyTorchDistributedRetriever # noqa: E402 # isort:skip27else:28 RagPyTorchDistributedRetriever = None29 30if is_ray_available():31 import ray # noqa: E402 # isort:skip32 from distributed_ray_retriever import RagRayDistributedRetriever, RayRetriever # noqa: E402 # isort:skip33else:34 ray = None35 RagRayDistributedRetriever = None36 RayRetriever = None37 38 39def require_distributed_retrieval(test_case):40 """41 Decorator marking a test that requires a set of dependencies necessary for pefrorm retrieval with42 :class:`~transformers.RagRetriever`.43 44 These tests are skipped when respective libraries are not installed.45 46 """47 if not (is_datasets_available() and is_faiss_available() and is_psutil_available()):48 test_case = unittest.skip("test requires Datasets, Faiss, psutil")(test_case)49 return test_case50 51 52@require_distributed_retrieval53class RagRetrieverTest(TestCase):54 def setUp(self):55 self.tmpdirname = tempfile.mkdtemp()56 self.retrieval_vector_size = 857 58 # DPR tok59 vocab_tokens = [60 "[UNK]",61 "[CLS]",62 "[SEP]",63 "[PAD]",64 "[MASK]",65 "want",66 "##want",67 "##ed",68 "wa",69 "un",70 "runn",71 "##ing",72 ",",73 "low",74 "lowest",75 ]76 dpr_tokenizer_path = os.path.join(self.tmpdirname, "dpr_tokenizer")77 os.makedirs(dpr_tokenizer_path, exist_ok=True)78 self.vocab_file = os.path.join(dpr_tokenizer_path, DPR_VOCAB_FILES_NAMES["vocab_file"])79 with open(self.vocab_file, "w", encoding="utf-8") as vocab_writer:80 vocab_writer.write("".join([x + "\n" for x in vocab_tokens]))81 82 # BART tok83 vocab = [84 "l",85 "o",86 "w",87 "e",88 "r",89 "s",90 "t",91 "i",92 "d",93 "n",94 "\u0120",95 "\u0120l",96 "\u0120n",97 "\u0120lo",98 "\u0120low",99 "er",100 "\u0120lowest",101 "\u0120newer",102 "\u0120wider",103 "<unk>",104 ]105 vocab_tokens = dict(zip(vocab, range(len(vocab))))106 merges = ["#version: 0.2", "\u0120 l", "\u0120l o", "\u0120lo w", "e r", ""]107 self.special_tokens_map = {"unk_token": "<unk>"}108 109 bart_tokenizer_path = os.path.join(self.tmpdirname, "bart_tokenizer")110 os.makedirs(bart_tokenizer_path, exist_ok=True)111 self.vocab_file = os.path.join(bart_tokenizer_path, BART_VOCAB_FILES_NAMES["vocab_file"])112 self.merges_file = os.path.join(bart_tokenizer_path, BART_VOCAB_FILES_NAMES["merges_file"])113 with open(self.vocab_file, "w", encoding="utf-8") as fp:114 fp.write(json.dumps(vocab_tokens) + "\n")115 with open(self.merges_file, "w", encoding="utf-8") as fp:116 fp.write("\n".join(merges))117 118 def get_dpr_tokenizer(self) -> DPRQuestionEncoderTokenizer:119 return DPRQuestionEncoderTokenizer.from_pretrained(os.path.join(self.tmpdirname, "dpr_tokenizer"))120 121 def get_bart_tokenizer(self) -> BartTokenizer:122 return BartTokenizer.from_pretrained(os.path.join(self.tmpdirname, "bart_tokenizer"))123 124 def tearDown(self):125 shutil.rmtree(self.tmpdirname)126 127 def get_dummy_dataset(self):128 dataset = Dataset.from_dict(129 {130 "id": ["0", "1"],131 "text": ["foo", "bar"],132 "title": ["Foo", "Bar"],133 "embeddings": [np.ones(self.retrieval_vector_size), 2 * np.ones(self.retrieval_vector_size)],134 }135 )136 dataset.add_faiss_index("embeddings", string_factory="Flat", metric_type=faiss.METRIC_INNER_PRODUCT)137 return dataset138 139 def get_dummy_pytorch_distributed_retriever(140 self, init_retrieval: bool, port=12345141 ) -> RagPyTorchDistributedRetriever:142 dataset = self.get_dummy_dataset()143 config = RagConfig(144 retrieval_vector_size=self.retrieval_vector_size,145 question_encoder=DPRConfig().to_dict(),146 generator=BartConfig().to_dict(),147 )148 with patch("transformers.models.rag.retrieval_rag.load_dataset") as mock_load_dataset:149 mock_load_dataset.return_value = dataset150 retriever = RagPyTorchDistributedRetriever(151 config,152 question_encoder_tokenizer=self.get_dpr_tokenizer(),153 generator_tokenizer=self.get_bart_tokenizer(),154 )155 if init_retrieval:156 retriever.init_retrieval(port)157 return retriever158 159 def get_dummy_ray_distributed_retriever(self, init_retrieval: bool) -> RagRayDistributedRetriever:160 # Have to run in local mode because sys.path modifications at top of161 # file are not propogated to remote workers.162 # https://stackoverflow.com/questions/54338013/parallel-import-a-python-file-from-sibling-folder163 ray.init(local_mode=True)164 config = RagConfig(165 retrieval_vector_size=self.retrieval_vector_size,166 question_encoder=DPRConfig().to_dict(),167 generator=BartConfig().to_dict(),168 )169 remote_cls = ray.remote(RayRetriever)170 workers = [remote_cls.remote() for _ in range(1)]171 with patch("transformers.models.rag.retrieval_rag.load_dataset") as mock_load_dataset:172 mock_load_dataset.return_value = self.get_dummy_dataset()173 retriever = RagRayDistributedRetriever(174 config,175 question_encoder_tokenizer=self.get_dpr_tokenizer(),176 generator_tokenizer=self.get_bart_tokenizer(),177 retrieval_workers=workers,178 )179 if init_retrieval:180 retriever.init_retrieval()181 return retriever182 183 def get_dummy_custom_hf_index_pytorch_retriever(self, init_retrieval: bool, from_disk: bool, port=12345):184 dataset = self.get_dummy_dataset()185 config = RagConfig(186 retrieval_vector_size=self.retrieval_vector_size,187 question_encoder=DPRConfig().to_dict(),188 generator=BartConfig().to_dict(),189 index_name="custom",190 )191 if from_disk:192 config.passages_path = os.path.join(self.tmpdirname, "dataset")193 config.index_path = os.path.join(self.tmpdirname, "index.faiss")194 dataset.get_index("embeddings").save(os.path.join(self.tmpdirname, "index.faiss"))195 dataset.drop_index("embeddings")196 dataset.save_to_disk(os.path.join(self.tmpdirname, "dataset"))197 del dataset198 retriever = RagPyTorchDistributedRetriever(199 config,200 question_encoder_tokenizer=self.get_dpr_tokenizer(),201 generator_tokenizer=self.get_bart_tokenizer(),202 )203 else:204 retriever = RagPyTorchDistributedRetriever(205 config,206 question_encoder_tokenizer=self.get_dpr_tokenizer(),207 generator_tokenizer=self.get_bart_tokenizer(),208 index=CustomHFIndex(config.retrieval_vector_size, dataset),209 )210 if init_retrieval:211 retriever.init_retrieval(port)212 return retriever213 214 def get_dummy_custom_hf_index_ray_retriever(self, init_retrieval: bool, from_disk: bool):215 # Have to run in local mode because sys.path modifications at top of216 # file are not propogated to remote workers.217 # https://stackoverflow.com/questions/54338013/parallel-import-a-python-file-from-sibling-folder218 ray.init(local_mode=True)219 dataset = self.get_dummy_dataset()220 config = RagConfig(221 retrieval_vector_size=self.retrieval_vector_size,222 question_encoder=DPRConfig().to_dict(),223 generator=BartConfig().to_dict(),224 index_name="custom",225 )226 remote_cls = ray.remote(RayRetriever)227 workers = [remote_cls.remote() for _ in range(1)]228 if from_disk:229 config.passages_path = os.path.join(self.tmpdirname, "dataset")230 config.index_path = os.path.join(self.tmpdirname, "index.faiss")231 dataset.get_index("embeddings").save(os.path.join(self.tmpdirname, "index.faiss"))232 dataset.drop_index("embeddings")233 dataset.save_to_disk(os.path.join(self.tmpdirname, "dataset"))234 del dataset235 retriever = RagRayDistributedRetriever(236 config,237 question_encoder_tokenizer=self.get_dpr_tokenizer(),238 generator_tokenizer=self.get_bart_tokenizer(),239 retrieval_workers=workers,240 index=CustomHFIndex.load_from_disk(241 vector_size=config.retrieval_vector_size,242 dataset_path=config.passages_path,243 index_path=config.index_path,244 ),245 )246 else:247 retriever = RagRayDistributedRetriever(248 config,249 question_encoder_tokenizer=self.get_dpr_tokenizer(),250 generator_tokenizer=self.get_bart_tokenizer(),251 retrieval_workers=workers,252 index=CustomHFIndex(config.retrieval_vector_size, dataset),253 )254 if init_retrieval:255 retriever.init_retrieval()256 return retriever257 258 def distributed_retriever_check(self, retriever: RagRetriever, hidden_states: np.array, n_docs: int) -> None:259 retrieved_doc_embeds, doc_ids, doc_dicts = retriever.retrieve(hidden_states, n_docs=n_docs)260 self.assertEqual(retrieved_doc_embeds.shape, (2, n_docs, self.retrieval_vector_size))261 self.assertEqual(len(doc_dicts), 2)262 self.assertEqual(sorted(doc_dicts[0]), ["embeddings", "id", "text", "title"])263 self.assertEqual(len(doc_dicts[0]["id"]), n_docs)264 self.assertEqual(doc_dicts[0]["id"][0], "1") # max inner product is reached with second doc265 self.assertEqual(doc_dicts[1]["id"][0], "0") # max inner product is reached with first doc266 self.assertListEqual(doc_ids.tolist(), [[1], [0]])267 268 def test_pytorch_distributed_retriever_retrieve(self):269 n_docs = 1270 hidden_states = np.array(271 [np.ones(self.retrieval_vector_size), -np.ones(self.retrieval_vector_size)], dtype=np.float32272 )273 274 self.distributed_retriever_check(275 self.get_dummy_pytorch_distributed_retriever(init_retrieval=True), hidden_states, n_docs276 )277 278 def test_custom_hf_index_pytorch_retriever_retrieve(self):279 n_docs = 1280 hidden_states = np.array(281 [np.ones(self.retrieval_vector_size), -np.ones(self.retrieval_vector_size)], dtype=np.float32282 )283 284 self.distributed_retriever_check(285 self.get_dummy_custom_hf_index_pytorch_retriever(init_retrieval=True, from_disk=False),286 hidden_states,287 n_docs,288 )289 290 def test_custom_pytorch_distributed_retriever_retrieve_from_disk(self):291 n_docs = 1292 hidden_states = np.array(293 [np.ones(self.retrieval_vector_size), -np.ones(self.retrieval_vector_size)], dtype=np.float32294 )295 296 self.distributed_retriever_check(297 self.get_dummy_custom_hf_index_pytorch_retriever(init_retrieval=True, from_disk=True),298 hidden_states,299 n_docs,300 )301 302 @require_ray303 def test_ray_distributed_retriever_retrieve(self):304 n_docs = 1305 hidden_states = np.array(306 [np.ones(self.retrieval_vector_size), -np.ones(self.retrieval_vector_size)], dtype=np.float32307 )308 309 self.distributed_retriever_check(310 self.get_dummy_ray_distributed_retriever(init_retrieval=True), hidden_states, n_docs311 )312 ray.shutdown()313 314 @require_ray315 def test_custom_hf_index_ray_retriever_retrieve(self):316 n_docs = 1317 hidden_states = np.array(318 [np.ones(self.retrieval_vector_size), -np.ones(self.retrieval_vector_size)], dtype=np.float32319 )320 with self.assertRaises(ValueError):321 self.distributed_retriever_check(322 self.get_dummy_custom_hf_index_ray_retriever(init_retrieval=True, from_disk=False),323 hidden_states,324 n_docs,325 )326 ray.shutdown()327 328 @require_ray329 def test_custom_ray_distributed_retriever_retrieve_from_disk(self):330 n_docs = 1331 hidden_states = np.array(332 [np.ones(self.retrieval_vector_size), -np.ones(self.retrieval_vector_size)], dtype=np.float32333 )334 335 self.distributed_retriever_check(336 self.get_dummy_custom_hf_index_ray_retriever(init_retrieval=True, from_disk=True), hidden_states, n_docs337 )338 ray.shutdown()339 