CoolFace
Apppublic

chendl/compositional_test

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
test_distributed_retriever.py339 linesDownload Raw Back to rag
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