MachineLearningReply/q-and-a-tool
0
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 