rag_cli.py 1.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748
  1. from agent.rag.document import load_pdf_document, clean_pdf_text
  2. from agent.rag import splitter
  3. from agent.rag.embedding import get_embedding_model, get_vectorstore
  4. from agent.llm import create_llm
  5. from agent.config import load_config
  6. from langchain_core.prompts import ChatPromptTemplate
  7. #加载文档
  8. pages = load_pdf_document("car_info.pdf")
  9. #清洗
  10. #documents = [clean_pdf_text(page.page_content) for page in pages]
  11. #分块
  12. #chunks=[splitter.embedding_split(document)for document in documents]
  13. docs=splitter.recursive_split(pages)
  14. #向量化+构建向量数据库
  15. embedding_model = get_embedding_model()
  16. db = get_vectorstore(embedding_model, docs)
  17. #创建检索器
  18. retriever = db.as_retriever(search_kwargs={"k": 3})
  19. context=""
  20. config=load_config()
  21. llm = create_llm(config)
  22. prompt=ChatPromptTemplate.from_messages([
  23. ("system", "你是⼀个专业的知识库助⼿。请根据以下上下⽂回答问题。"),
  24. ("user", "根据以下内容回答用户问题,如果无法从中获取答案,请说“抱歉,我无法回答这个问题。”\n\n{context}\n\n用户问题: {question}")
  25. ])
  26. while True:
  27. question = input("请输入: ").strip()
  28. if not question:
  29. continue
  30. if question.lower() in ["exit", "quit", "q"]:
  31. break
  32. #查询向量数据库
  33. relevant_docs = retriever.invoke(question)
  34. for doc in relevant_docs:
  35. context += doc.page_content + "\n"
  36. #生成回答
  37. chain=prompt|llm
  38. answer = chain.invoke({"context": context, "question": question})
  39. print("AI回答:", answer)