rag_cli.py 2.1 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364
  1. """RAG 命令行入口。"""
  2. from pathlib import Path
  3. from langchain_core.prompts import ChatPromptTemplate
  4. from agent.config import load_config
  5. from agent.llm import create_llm
  6. from agent.rag.document import load_pdf_document
  7. from agent.rag.embedding import get_embedding_model
  8. from agent.rag.multi_retriever import MultiRetrieverSystem
  9. from agent.rag.splitter import recursive_split
  10. from query_optimization import get_prompt_template
  11. DOCUMENT_PATH = Path("car_info.pdf")
  12. def create_rag_system() -> tuple[MultiRetrieverSystem, object, ChatPromptTemplate]:
  13. """加载文档并创建检索器和问答链所需组件。"""
  14. pages = load_pdf_document(str(DOCUMENT_PATH))
  15. documents = recursive_split(pages)
  16. config = load_config()
  17. llm = create_llm(config)
  18. retriever_system = MultiRetrieverSystem(
  19. documents=documents,
  20. embedding_model=get_embedding_model(),
  21. llm=llm,
  22. )
  23. return retriever_system, llm
  24. def run_cli() -> None:
  25. """启动交互式 RAG 命令行。"""
  26. retriever_system, llm = create_rag_system()
  27. prompt =get_prompt_template("rewrite")
  28. mode = "ensemble"
  29. valid_modes = {"bm25", "vector", "ensemble"}
  30. print("检索模式:ensemble(可输入 /mode bm25、/mode vector 或 /mode ensemble 切换)")
  31. while True:
  32. question = input("请输入问题:").strip()
  33. if not question:
  34. continue
  35. if question.lower() in {"exit", "quit", "q"}:
  36. break
  37. if question.startswith("/mode "):
  38. new_mode = question.removeprefix("/mode ").strip().lower()
  39. if new_mode in valid_modes:
  40. mode = new_mode
  41. print(f"已切换为 {mode} 检索。")
  42. else:
  43. print("不支持的检索模式,可选:bm25、vector、ensemble")
  44. continue
  45. relevant_docs = retriever_system.search(question, mode)
  46. context = "\n\n".join(document.page_content for document in relevant_docs)
  47. answer = (prompt | llm).invoke({"context": context, "question": question})
  48. print(f"AI 回答:{answer.content}")
  49. if __name__ == "__main__":
  50. run_cli()