Team Ai
Apppublic

MachineLearningReply/q-and-a-tool

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
document_qa_engine.py144 linesDownload Raw Back to root
1from typing import List2 3from haystack.dataclasses import ChatMessage4from pypdf import PdfReader5from haystack.utils import Secret6from haystack import Pipeline, Document, component7 8from haystack.components.preprocessors import DocumentCleaner, DocumentSplitter9from haystack.components.writers import DocumentWriter10from haystack.components.embedders import SentenceTransformersDocumentEmbedder, SentenceTransformersTextEmbedder11from haystack.document_stores.in_memory import InMemoryDocumentStore12from haystack.components.retrievers.in_memory import InMemoryEmbeddingRetriever13from haystack.components.builders import DynamicChatPromptBuilder14from haystack.components.generators.chat import OpenAIChatGenerator, HuggingFaceTGIChatGenerator15from haystack.document_stores.types import DuplicatePolicy16 17SENTENCE_RETREIVER_MODEL = "sentence-transformers/all-MiniLM-L6-v2"18 19MAX_TOKENS = 50020 21template = """22As a professional HR recruiter given the following information, answer the question shortly and concisely in 1 or 2 sentences.23 24Context:25{% for document in documents %}26    {{ document.content }}27{% endfor %}28 29Question: {{question}}30Answer:31"""32 33 34@component35class UploadedFileConverter:36    """37    A component to convert uploaded PDF files to Documents38    """39 40    @component.output_types(documents=List[Document])41    def run(self, uploaded_file):42        pdf = PdfReader(uploaded_file)43        documents = []44        # uploaded file name without .pdf at the end and with _ and page number at the end45        name = uploaded_file.name.rstrip('.PDF') + '_'46        for page in pdf.pages:47            documents.append(48                Document(49                    content=page.extract_text(),50                    meta={'name': name + f"_{page.page_number}"}))51        return {"documents": documents}52 53 54def create_ingestion_pipeline(document_store):55    doc_embedder = SentenceTransformersDocumentEmbedder(model=SENTENCE_RETREIVER_MODEL)56    doc_embedder.warm_up()57 58    pipeline = Pipeline()59    pipeline.add_component("converter", UploadedFileConverter())60    pipeline.add_component("cleaner", DocumentCleaner())61    pipeline.add_component("splitter",62                           DocumentSplitter(split_by="passage", split_length=100, split_overlap=10))63    pipeline.add_component("embedder", doc_embedder)64    pipeline.add_component("writer",65                           DocumentWriter(document_store=document_store, policy=DuplicatePolicy.OVERWRITE))66 67    pipeline.connect("converter", "cleaner")68    pipeline.connect("cleaner", "splitter")69    pipeline.connect("splitter", "embedder")70    pipeline.connect("embedder", "writer")71    return pipeline72 73 74def create_inference_pipeline(document_store, model_name, api_key):75    if model_name == "local LLM":76        generator = OpenAIChatGenerator(api_key=Secret.from_token("<local LLM doesn't need an API key>"),77                                        model=model_name,78                                        api_base_url="http://localhost:1234/v1",79                                        generation_kwargs={"max_tokens": MAX_TOKENS},80                                        )81    elif "gpt" in model_name:82        generator = OpenAIChatGenerator(api_key=Secret.from_token(api_key), model=model_name,83                                        generation_kwargs={"max_tokens": MAX_TOKENS, "temperature": 0},84                                        streaming_callback=lambda chunk: print(chunk.content, end="", flush=True),85                                        86                                        )87    else:88        generator = HuggingFaceTGIChatGenerator(token=Secret.from_token(api_key), model=model_name,89                                                generation_kwargs={"max_new_tokens": MAX_TOKENS}90                                                )91    pipeline = Pipeline()92    pipeline.add_component("text_embedder",93                           SentenceTransformersTextEmbedder(model=SENTENCE_RETREIVER_MODEL))94    pipeline.add_component("retriever", InMemoryEmbeddingRetriever(document_store, top_k=3))95    pipeline.add_component("prompt_builder",96                           DynamicChatPromptBuilder(runtime_variables=["query", "documents"]))97    pipeline.add_component("llm", generator)98    pipeline.connect("text_embedder.embedding", "retriever.query_embedding")99    pipeline.connect("retriever.documents", "prompt_builder.documents")100    pipeline.connect("prompt_builder.prompt", "llm.messages")101 102    return pipeline103 104 105class DocumentQAEngine:106    def __init__(self,107                 model_name,108                 api_key=None109                 ):110        self.api_key = api_key111        self.model_name = model_name112        document_store = InMemoryDocumentStore()113        self.chunks = []114        self.inference_pipeline = create_inference_pipeline(document_store, model_name, api_key)115        self.pdf_ingestion_pipeline = create_ingestion_pipeline(document_store)116 117    def ingest_pdf(self, uploaded_file):118        self.pdf_ingestion_pipeline.run({"converter": {"uploaded_file": uploaded_file}})119 120    def inference(self, query, input_messages: List[dict]):121        system_message = ChatMessage.from_system(122            "You are a professional analyzer of git repos, having access to the repo content. In 1-3 sentences")123        messages = [system_message]124        for message in input_messages:125            if message["role"] == "user":126                messages.append(ChatMessage.from_system(message["content"]))127            else:128                messages.append(129                    ChatMessage.from_user(message["content"]))130        messages.append(ChatMessage.from_user("""131        Relevant information from the uploaded repo:132            {% for doc in documents %}133                {{ doc.content }}134            {% endfor %}135 136            \nQuestion: {{query}}137            \nAnswer:138        """))139        res = self.inference_pipeline.run(data={"text_embedder": {"text": query},140                                                "prompt_builder": {"prompt_source": messages,141                                                                   "query": query142                                                                   }})143        return res["llm"]["replies"][0].content144