multi_retriever.py 2.3 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465
  1. """BM25、Milvus 向量和混合检索。"""
  2. from langchain_classic.retrievers import EnsembleRetriever
  3. from langchain_community.retrievers import BM25Retriever
  4. from langchain_core.documents import Document
  5. from langchain_milvus import Milvus
  6. class MultiRetrieverSystem:
  7. """多路召回检索系统,默认使用混合检索。"""
  8. def __init__(self, documents, embedding_model, llm):
  9. self.documents = documents
  10. self.embedding_model = embedding_model
  11. self.llm = llm
  12. self.setup_retrievers()
  13. def setup_retrievers(self):
  14. """初始化 BM25、向量和混合检索器。"""
  15. self.bm25 = BM25Retriever.from_documents(self.documents)
  16. self.bm25.k = 10
  17. self.vectorstore = Milvus(
  18. embedding_function=self.embedding_model,
  19. connection_args={"uri": "http://localhost:19530"},
  20. collection_name="car_info_collection",
  21. primary_field="id",
  22. text_field="content",
  23. vector_field="embedding",
  24. auto_id=True,
  25. search_params={"metric_type": "COSINE", "params": {}},
  26. )
  27. self._index_documents_if_empty()
  28. self.vector = self.vectorstore.as_retriever(search_kwargs={"k": 10})
  29. # LangChain 当前版本使用 RRF 融合排序,不需要 normalize_scores 参数。
  30. self.ensemble = EnsembleRetriever(
  31. retrievers=[self.bm25, self.vector],
  32. weights=[0.4, 0.6],
  33. )
  34. def _index_documents_if_empty(self):
  35. """Milvus 集合为空时才写入分块,避免重复入库。"""
  36. rows = self.vectorstore.client.query(
  37. collection_name="car_info_collection",
  38. filter="id >= 0",
  39. output_fields=["id"],
  40. limit=1,
  41. consistency_level="Strong",
  42. )
  43. if not rows:
  44. self.vectorstore.add_documents(self.documents)
  45. def search(self, query, mode="ensemble") -> list[Document]:
  46. """执行检索,mode 可选 bm25、vector、ensemble。"""
  47. retriever_map = {
  48. "bm25": self.bm25,
  49. "vector": self.vector,
  50. "ensemble": self.ensemble,
  51. }
  52. if mode not in retriever_map:
  53. choices = ", ".join(retriever_map)
  54. raise ValueError(f"不支持的检索模式: {mode},可选: {choices}")
  55. return retriever_map[mode].invoke(query)[:10]