jdwh08s/Autodoc-Lifter
0
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 