CoolFace
Apppublic

jdwh08s/Autodoc-Lifter

sourceHugging Faceagpl-3.0updated 2y agoView on Hugging Face
0likes
engine.py127 linesDownload Raw Back to root
1#####################################################2### DOCUMENT PROCESSOR [ENGINE]3#####################################################4# Jonathan Wang5 6# ABOUT:7# This project creates an app to chat with PDFs.8 9# This is the ENGINE10# which defines how LLMs handle processing.11#####################################################12## TODO Board:13 14#####################################################15## IMPORTS16from __future__ import annotations17 18import gc19from typing import TYPE_CHECKING, Callable, List, Optional, cast20 21from llama_index.core.query_engine import CustomQueryEngine22from llama_index.core.schema import NodeWithScore, QueryBundle23from llama_index.core.settings import (24    Settings,25)26from torch.cuda import empty_cache27 28if TYPE_CHECKING:29    from llama_index.core.base.response.schema import Response30    from llama_index.core.callbacks import CallbackManager31    from llama_index.core.postprocessor.types import BaseNodePostprocessor32    from llama_index.core.response_synthesizers import (33        BaseSynthesizer,34    )35    from llama_index.core.retrievers import BaseRetriever36 37# Own Modules38 39#####################################################40## CODE41class RAGQueryEngine(CustomQueryEngine):42    """Custom RAG Query Engine."""43 44    retriever: BaseRetriever45    response_synthesizer: BaseSynthesizer46    node_postprocessors: Optional[List[BaseNodePostprocessor]] = []47 48    # def __init__(49    #     self,50    #     retriever: BaseRetriever,51    #     response_synthesizer: Optional[BaseSynthesizer] = None,52    #     node_postprocessors: Optional[List[BaseNodePostprocessor]] = None,53    #     callback_manager: Optional[CallbackManager] = None,54    # ) -> None:55    #     self._retriever = retriever56    #     # callback_manager = (57    #     #     callback_manager58    #     #     Settings.callback_manager59    #     # )60    #     # llm = llm or Settings.llm61 62    #     self._response_synthesizer = response_synthesizer or get_response_synthesizer(63    #         # llm=llm,64    #         # service_context=service_context,65    #         # callback_manager=callback_manager,66    #     )67    #     self._node_postprocessors = node_postprocessors or []68    #     self._metadata_mode = metadata_mode69 70    #     for node_postprocessor in self._node_postprocessors:71    #         node_postprocessor.callback_manager = callback_manager72 73    #     super().__init__(callback_manager=callback_manager)74 75    @classmethod76    def class_name(cls) -> str:77        """Class name."""78        return "RAGQueryEngine"79 80    # taken from Llamaindex CustomEngine:81    # https://github.com/run-llama/llama_index/blob/main/llama-index-core/llama_index/core/query_engine/retriever_query_engine.py#L13482    def _apply_node_postprocessors(83        self, nodes: list[NodeWithScore], query_bundle: QueryBundle84    ) -> list[NodeWithScore]:85        if self.node_postprocessors is None:86            return nodes87 88        for node_postprocessor in self.node_postprocessors:89            nodes = node_postprocessor.postprocess_nodes(90                nodes, query_bundle=query_bundle91            )92        return nodes93 94    def retrieve(self, query_bundle: QueryBundle) -> list[NodeWithScore]:95        nodes = self.retriever.retrieve(query_bundle)96        return self._apply_node_postprocessors(nodes, query_bundle=query_bundle)97 98    async def aretrieve(self, query_bundle: QueryBundle) -> list[NodeWithScore]:99        nodes = await self.retriever.aretrieve(query_bundle)100        return self._apply_node_postprocessors(nodes, query_bundle=query_bundle)101 102    def custom_query(self, query_str: str) -> Response:103        # Convert query string into query bundle104        query_bundle = QueryBundle(query_str=query_str)105        nodes = self.retrieve(query_bundle)  # also does the postprocessing.106 107        response_obj = self.response_synthesizer.synthesize(query_bundle, nodes)108 109        empty_cache()110        gc.collect()111        return cast(Response, response_obj)  # type: ignore112 113 114# @st.cache_resource  # none of these can be hashable or cached :(115def get_engine(116    retriever: BaseRetriever,117    response_synthesizer: BaseSynthesizer,118    node_postprocessors: list[BaseNodePostprocessor] | None = None,119    callback_manager: CallbackManager | None = None,120) -> RAGQueryEngine:121    return RAGQueryEngine(122        retriever=retriever,123        response_synthesizer=response_synthesizer,124        node_postprocessors=node_postprocessors,125        callback_manager=callback_manager or Settings.callback_manager,126    )127