Team Ai
Modelpublic

Nevermined/test_haystack

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
handler.py60 linesDownload Raw Back to root
1import os2from haystack.utils import fetch_archive_from_http, clean_wiki_text, convert_files_to_docs3from haystack.schema import Answer4from haystack.document_stores import InMemoryDocumentStore5from haystack.pipelines import ExtractiveQAPipeline6from haystack.nodes import FARMReader, TfidfRetriever7import logging8import json9 10os.environ['TOKENIZERS_PARALLELISM'] ="false"11 12#Haystack Components13def start_haystack():14    document_store = InMemoryDocumentStore()15    load_and_write_data(document_store)16    retriever = TfidfRetriever(document_store=document_store)17    reader = FARMReader(model_name_or_path="deepset/roberta-base-squad2-distilled", use_gpu=True)18    pipeline = ExtractiveQAPipeline(reader, retriever)19    return pipeline20 21def load_and_write_data(document_store):22    23    # Get the absolute path of the script24    script_path = os.path.realpath(__file__)25    # Get the script directory26    script_dir = os.path.dirname(script_path)27    doc_dir = script_dir + "/dao_data"28    print("Loading data ...")29 30    docs = convert_files_to_docs(dir_path=doc_dir, clean_func=clean_wiki_text, split_paragraphs=True)31    document_store.write_documents(docs)32 33 34class EndpointHandler():35    def __init__(self, path=""):36        # load the optimized model     37        self.pipeline = start_haystack()38 39 40    def __call__(self, data):41        """42        Args:43            data (:obj:):44                includes the input data and the parameters for the inference.45        Return:46            A :obj:`list`:. The object returned should be a list of one list like [[{"label": 0.9939950108528137}]] containing :47                - "label": A string representing what the label/class is. There can be multiple labels.48                - "score": A score between 0 and 1 describing how confident the model is for this label/class.49        """50        inputs = data.pop("inputs", None)51        question = inputs.pop("question", None)52        if question is not None:53            prediction = self.pipeline.run(query=question, params={"Retriever": {"top_k": 10}, "Reader": {"top_k": 5}})54        else:55            return {}56        57        # postprocess the prediction58        response = { "answer": prediction['answers'][0].answer}59        return json.dumps(response)60