Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
base.py369 linesDownload Raw Back to retrieval_qa
1"""Chain for question-answering against a vector database."""2 3from __future__ import annotations4 5import inspect6from abc import abstractmethod7from typing import Any8 9from langchain_core._api import deprecated10from langchain_core.callbacks import (11    AsyncCallbackManagerForChainRun,12    CallbackManagerForChainRun,13    Callbacks,14)15from langchain_core.documents import Document16from langchain_core.language_models import BaseLanguageModel17from langchain_core.prompts import PromptTemplate18from langchain_core.retrievers import BaseRetriever19from langchain_core.vectorstores import VectorStore20from pydantic import ConfigDict, Field, model_validator21from typing_extensions import override22 23from langchain_classic.chains.base import Chain24from langchain_classic.chains.combine_documents.base import BaseCombineDocumentsChain25from langchain_classic.chains.combine_documents.stuff import StuffDocumentsChain26from langchain_classic.chains.llm import LLMChain27from langchain_classic.chains.question_answering import load_qa_chain28from langchain_classic.chains.question_answering.stuff_prompt import PROMPT_SELECTOR29 30 31@deprecated(32    since="0.2.13",33    removal="2.0.0",34    alternative="langchain.agents.create_agent",35    addendum=(36        "Build new RAG flows with `create_agent` and a retrieval tool. See "37        "https://docs.langchain.com/oss/python/langchain/rag"38    ),39)40class BaseRetrievalQA(Chain):41    """Base class for question-answering chains."""42 43    combine_documents_chain: BaseCombineDocumentsChain44    """Chain to use to combine the documents."""45    input_key: str = "query"46    output_key: str = "result"47    return_source_documents: bool = False48    """Return the source documents or not."""49 50    model_config = ConfigDict(51        populate_by_name=True,52        arbitrary_types_allowed=True,53        extra="forbid",54    )55 56    @property57    def input_keys(self) -> list[str]:58        """Input keys."""59        return [self.input_key]60 61    @property62    def output_keys(self) -> list[str]:63        """Output keys."""64        _output_keys = [self.output_key]65        if self.return_source_documents:66            _output_keys = [*_output_keys, "source_documents"]67        return _output_keys68 69    @classmethod70    def from_llm(71        cls,72        llm: BaseLanguageModel,73        prompt: PromptTemplate | None = None,74        callbacks: Callbacks = None,75        llm_chain_kwargs: dict | None = None,76        **kwargs: Any,77    ) -> BaseRetrievalQA:78        """Initialize from LLM."""79        _prompt = prompt or PROMPT_SELECTOR.get_prompt(llm)80        llm_chain = LLMChain(81            llm=llm,82            prompt=_prompt,83            callbacks=callbacks,84            **(llm_chain_kwargs or {}),85        )86        document_prompt = PromptTemplate(87            input_variables=["page_content"],88            template="Context:\n{page_content}",89        )90        combine_documents_chain = StuffDocumentsChain(91            llm_chain=llm_chain,92            document_variable_name="context",93            document_prompt=document_prompt,94            callbacks=callbacks,95        )96 97        return cls(98            combine_documents_chain=combine_documents_chain,99            callbacks=callbacks,100            **kwargs,101        )102 103    @classmethod104    def from_chain_type(105        cls,106        llm: BaseLanguageModel,107        chain_type: str = "stuff",108        chain_type_kwargs: dict | None = None,109        **kwargs: Any,110    ) -> BaseRetrievalQA:111        """Load chain from chain type."""112        _chain_type_kwargs = chain_type_kwargs or {}113        combine_documents_chain = load_qa_chain(114            llm,115            chain_type=chain_type,116            **_chain_type_kwargs,117        )118        return cls(combine_documents_chain=combine_documents_chain, **kwargs)119 120    @abstractmethod121    def _get_docs(122        self,123        question: str,124        *,125        run_manager: CallbackManagerForChainRun,126    ) -> list[Document]:127        """Get documents to do question answering over."""128 129    def _call(130        self,131        inputs: dict[str, Any],132        run_manager: CallbackManagerForChainRun | None = None,133    ) -> dict[str, Any]:134        """Run get_relevant_text and llm on input query.135 136        If chain has 'return_source_documents' as 'True', returns137        the retrieved documents as well under the key 'source_documents'.138 139        Example:140        ```python141        res = indexqa({"query": "This is my query"})142        answer, docs = res["result"], res["source_documents"]143        ```144        """145        _run_manager = run_manager or CallbackManagerForChainRun.get_noop_manager()146        question = inputs[self.input_key]147        accepts_run_manager = (148            "run_manager" in inspect.signature(self._get_docs).parameters149        )150        if accepts_run_manager:151            docs = self._get_docs(question, run_manager=_run_manager)152        else:153            docs = self._get_docs(question)  # type: ignore[call-arg]154        answer = self.combine_documents_chain.run(155            input_documents=docs,156            question=question,157            callbacks=_run_manager.get_child(),158        )159 160        if self.return_source_documents:161            return {self.output_key: answer, "source_documents": docs}162        return {self.output_key: answer}163 164    @abstractmethod165    async def _aget_docs(166        self,167        question: str,168        *,169        run_manager: AsyncCallbackManagerForChainRun,170    ) -> list[Document]:171        """Get documents to do question answering over."""172 173    async def _acall(174        self,175        inputs: dict[str, Any],176        run_manager: AsyncCallbackManagerForChainRun | None = None,177    ) -> dict[str, Any]:178        """Run get_relevant_text and llm on input query.179 180        If chain has 'return_source_documents' as 'True', returns181        the retrieved documents as well under the key 'source_documents'.182 183        Example:184        ```python185        res = indexqa({"query": "This is my query"})186        answer, docs = res["result"], res["source_documents"]187        ```188        """189        _run_manager = run_manager or AsyncCallbackManagerForChainRun.get_noop_manager()190        question = inputs[self.input_key]191        accepts_run_manager = (192            "run_manager" in inspect.signature(self._aget_docs).parameters193        )194        if accepts_run_manager:195            docs = await self._aget_docs(question, run_manager=_run_manager)196        else:197            docs = await self._aget_docs(question)  # type: ignore[call-arg]198        answer = await self.combine_documents_chain.arun(199            input_documents=docs,200            question=question,201            callbacks=_run_manager.get_child(),202        )203 204        if self.return_source_documents:205            return {self.output_key: answer, "source_documents": docs}206        return {self.output_key: answer}207 208 209@deprecated(210    since="0.1.17",211    removal="2.0.0",212    alternative="langchain.agents.create_agent",213    addendum=(214        "Build new RAG flows with `create_agent` and a retrieval tool. See "215        "https://docs.langchain.com/oss/python/langchain/rag"216    ),217)218class RetrievalQA(BaseRetrievalQA):219    """Chain for question-answering against an index.220 221    This class is deprecated. See below for an example implementation using222    `create_retrieval_chain`:223 224        ```python225        from langchain_classic.chains import create_retrieval_chain226        from langchain_classic.chains.combine_documents import (227            create_stuff_documents_chain,228        )229        from langchain_core.prompts import ChatPromptTemplate230        from langchain_openai import ChatOpenAI231 232 233        retriever = ...  # Your retriever234        model = ChatOpenAI()235 236        system_prompt = (237            "Use the given context to answer the question. "238            "If you don't know the answer, say you don't know. "239            "Use three sentence maximum and keep the answer concise. "240            "Context: {context}"241        )242        prompt = ChatPromptTemplate.from_messages(243            [244                ("system", system_prompt),245                ("human", "{input}"),246            ]247        )248        question_answer_chain = create_stuff_documents_chain(model, prompt)249        chain = create_retrieval_chain(retriever, question_answer_chain)250 251        chain.invoke({"input": query})252        ```253 254    Example:255        ```python256        from langchain_openai import OpenAI257        from langchain_classic.chains import RetrievalQA258        from langchain_community.vectorstores import FAISS259        from langchain_core.vectorstores import VectorStoreRetriever260 261        retriever = VectorStoreRetriever(vectorstore=FAISS(...))262        retrievalQA = RetrievalQA.from_llm(llm=OpenAI(), retriever=retriever)263        ```264    """265 266    retriever: BaseRetriever = Field(exclude=True)267 268    def _get_docs(269        self,270        question: str,271        *,272        run_manager: CallbackManagerForChainRun,273    ) -> list[Document]:274        """Get docs."""275        return self.retriever.invoke(276            question,277            config={"callbacks": run_manager.get_child()},278        )279 280    async def _aget_docs(281        self,282        question: str,283        *,284        run_manager: AsyncCallbackManagerForChainRun,285    ) -> list[Document]:286        """Get docs."""287        return await self.retriever.ainvoke(288            question,289            config={"callbacks": run_manager.get_child()},290        )291 292    @property293    def _chain_type(self) -> str:294        """Return the chain type."""295        return "retrieval_qa"296 297 298@deprecated(299    since="0.2.13",300    removal="2.0.0",301    alternative="langchain.agents.create_agent",302    addendum=(303        "Build new RAG flows with `create_agent` and a retrieval tool. See "304        "https://docs.langchain.com/oss/python/langchain/rag"305    ),306)307class VectorDBQA(BaseRetrievalQA):308    """Chain for question-answering against a vector database."""309 310    vectorstore: VectorStore = Field(exclude=True, alias="vectorstore")311    """Vector Database to connect to."""312    k: int = 4313    """Number of documents to query for."""314    search_type: str = "similarity"315    """Search type to use over vectorstore. `similarity` or `mmr`."""316    search_kwargs: dict[str, Any] = Field(default_factory=dict)317    """Extra search args."""318 319    @model_validator(mode="before")320    @classmethod321    def validate_search_type(cls, values: dict) -> Any:322        """Validate search type."""323        if "search_type" in values:324            search_type = values["search_type"]325            if search_type not in ("similarity", "mmr"):326                msg = f"search_type of {search_type} not allowed."327                raise ValueError(msg)328        return values329 330    @override331    def _get_docs(332        self,333        question: str,334        *,335        run_manager: CallbackManagerForChainRun,336    ) -> list[Document]:337        """Get docs."""338        if self.search_type == "similarity":339            docs = self.vectorstore.similarity_search(340                question,341                k=self.k,342                **self.search_kwargs,343            )344        elif self.search_type == "mmr":345            docs = self.vectorstore.max_marginal_relevance_search(346                question,347                k=self.k,348                **self.search_kwargs,349            )350        else:351            msg = f"search_type of {self.search_type} not allowed."352            raise ValueError(msg)353        return docs354 355    async def _aget_docs(356        self,357        question: str,358        *,359        run_manager: AsyncCallbackManagerForChainRun,360    ) -> list[Document]:361        """Get docs."""362        msg = "VectorDBQA does not support async"363        raise NotImplementedError(msg)364 365    @property366    def _chain_type(self) -> str:367        """Return the chain type."""368        return "vector_db_qa"369 
codekingpro/portable-devtools · Team Ai