import os import requests from typing import List from llama_index.core.node_parser import SentenceSplitter from llama_index.retrievers.bm25 import BM25Retriever from llama_index.core.retrievers import QueryFusionRetriever from llama_index.core.query_engine import RetrieverQueryEngine from llama_index.core.postprocessor import SentenceTransformerRerank from llama_index.core.embeddings import BaseEmbedding from llama_index.llms.openai_like import OpenAILike from llama_index.core import ( VectorStoreIndex, SimpleDirectoryReader, Settings, StorageContext, load_index_from_storage, )
class LocalLlamaServerEmbedding(BaseEmbedding): api_base: str api_key: str = "dummy" max_tokens: int = 400
def _get_embedding(self, text: str) -> List[float]: url = f"{self.api_base}/embeddings" payload = { "input": text, "model": "Qwen3-Embedding-0.6B-Q8_0.gguf", } headers = {"Authorization": f"Bearer {self.api_key}"} resp = requests.post(url, json=payload, headers=headers, timeout=120) if resp.status_code != 200: raise RuntimeError( f"Embedding 接口请求失败 status={resp.status_code} body={resp.text}" ) data = resp.json() return data["data"][0]["embedding"]
def _get_text_embedding(self, text: str) -> List[float]: return self._get_embedding(text)
def _get_query_embedding(self, query: str) -> List[float]: return self._get_embedding(query)
async def _aget_query_embedding(self, query: str) -> List[float]: return self._get_embedding(query)
async def _aget_text_embedding(self, text: str) -> List[float]: return self._get_embedding(text)
def setup_base_env(): os.environ["OPENAI_API_KEY"] = "dummy" os.environ["OPENAI_BASE_URL"] = "http://127.0.0.1:11433/v1"
llm = OpenAILike( model="qwen2.5-1.5b-instruct-q4_k_m.gguf", api_base=os.environ["OPENAI_BASE_URL"], api_key=os.environ["OPENAI_API_KEY"], is_chat_model=True, context_window=1024, temperature=0.1, system_prompt="你是企业文档助手,严格依据检索文档回答,文档没有相关信息请明确告知。", )
Settings.node_parser = SentenceSplitter(chunk_size=380, chunk_overlap=50)
Settings.llm = llm Settings.embed_model = LocalLlamaServerEmbedding( api_base="http://127.0.0.1:11434/v1" )
class HybridRAG: """Hybrid-RAG:向量检索 + BM25 + RRF 倒数排序融合 + Rerank"""
def __init__(self, data_dir, persist_dir="./storage_hybrid"): self.data_dir = data_dir self.persist_dir = persist_dir setup_base_env() self.index = self._load_or_build_index() self.query_engine = self._build_engine()
def _load_or_build_index(self): if os.path.exists(self.persist_dir): try: storage_context = StorageContext.from_defaults(persist_dir=self.persist_dir) idx = load_index_from_storage(storage_context) print("HybridRAG 索引加载成功") return idx except Exception as e: print(f"索引加载失败: {e},删除旧索引并重建") import shutil shutil.rmtree(self.persist_dir, ignore_errors=True) try: documents = SimpleDirectoryReader( self.data_dir, required_exts=[".pdf", ".docx", ".txt"], recursive=False, ).load_data() except Exception as err: print(f"[文档读取解析失败] {err}") raise err
print(f"读取文档片段 {len(documents)}") idx = VectorStoreIndex.from_documents(documents, show_progress=True) idx.storage_context.persist(persist_dir=self.persist_dir) return idx
def _build_engine(self): vector_retriever = self.index.as_retriever(similarity_top_k=10) bm25_retriever = BM25Retriever.from_defaults( docstore=self.index.docstore, similarity_top_k=10 ) fusion_retriever = QueryFusionRetriever( retrievers=[vector_retriever, bm25_retriever], similarity_top_k=8, mode="reciprocal_rerank", num_queries=1, ) post_processors = [ SentenceTransformerRerank( model="E:/llamacpp/models/BAAI--bge-reranker-v2-m3", top_n=4 ) ] engine = RetrieverQueryEngine.from_args( retriever=fusion_retriever, node_postprocessors=post_processors, response_mode="compact", ) return engine
def query(self, question: str): resp = self.query_engine.query(question) return { "answer": str(resp), "sources": [n.metadata for n in resp.source_nodes], }
if __name__ == "__main__":
rag = HybridRAG("./company_docs")
res = rag.query("公司报销流程是什么?") print("=" * 60) print(res["answer"]) print("=" * 60) print("引用来源:") for s in res["sources"]: print(" -", s.get("file_name", s))
|