| |
| import os |
| |
| from dotenv import load_dotenv |
| load_dotenv() |
|
|
| |
| |
| |
| from langchain_text_splitters import CharacterTextSplitter |
|
|
| |
| |
| |
| from langchain_community.document_loaders import TextLoader |
|
|
| |
| |
| |
| |
| from langchain_openai import OpenAIEmbeddings |
| |
| |
|
|
| |
| |
| |
| from langchain_community.vectorstores import Chroma |
|
|
| |
| |
| |
| |
| from langchain_openai import ChatOpenAI |
| |
| |
|
|
| |
| |
| |
| from langchain_community.chains import RetrievalQA |
|
|
| |
| os.environ["OPENAI_API_KEY"] = os.getenv("OPENAI_API_KEY", "your-openai-api-key") |
|
|
| def load_and_split_documents(file_path="docs/knowledge.txt"): |
| """加载文档并分割为小片段""" |
| |
| os.makedirs(os.path.dirname(file_path), exist_ok=True) |
| |
| if not os.path.exists(file_path): |
| with open(file_path, "w", encoding="utf-8") as f: |
| f.write("LangChain 是一个用于构建 LLM 应用的框架,支持检索增强生成(RAG)等功能。\n") |
| f.write("RetrievalQA 是 LangChain 中用于检索问答的核心链,在 v0.1.0 后迁移到了 langchain_community 模块。\n") |
| f.write("CharacterTextSplitter 在 v0.1.0 后迁移到了 langchain_text_splitters 模块。\n") |
| |
| loader = TextLoader(file_path, encoding="utf-8") |
| documents = loader.load() |
| |
| |
| text_splitter = CharacterTextSplitter( |
| chunk_size=1000, |
| chunk_overlap=200, |
| separator="\n" |
| ) |
| split_docs = text_splitter.split_documents(documents) |
| return split_docs |
|
|
| def create_vector_store(docs): |
| """创建向量存储(使用 Chroma 本地向量库)""" |
| |
| |
| embeddings = OpenAIEmbeddings(model="text-embedding-3-small") |
| |
| |
| |
| |
| vector_store = Chroma.from_documents( |
| documents=docs, |
| embedding=embeddings, |
| persist_directory="./chroma_db" |
| ) |
| |
| return vector_store |
|
|
| def create_retrieval_qa_chain(vector_store): |
| """创建 RetrievalQA 检索问答链""" |
| |
| |
| llm = ChatOpenAI( |
| model_name="gpt-3.5-turbo", |
| temperature=0.1 |
| ) |
| |
| |
| |
| |
| retriever = vector_store.as_retriever( |
| search_kwargs={"k": 3} |
| ) |
| |
| |
| qa_chain = RetrievalQA.from_chain_type( |
| llm=llm, |
| chain_type="stuff", |
| retriever=retriever, |
| return_source_documents=True |
| ) |
| return qa_chain |
|
|
| if __name__ == "__main__": |
| |
| docs = load_and_split_documents() |
| |
| vector_store = create_vector_store(docs) |
| |
| qa_chain = create_retrieval_qa_chain(vector_store) |
| |
| |
| query = "CharacterTextSplitter 在 LangChain 新版本中被迁移到了哪里?" |
| result = qa_chain.invoke({"query": query}) |
| |
| |
| print("===== 问题 =====") |
| print(query) |
| print("\n===== 回答 =====") |
| print(result["result"]) |
| print("\n===== 参考文档 =====") |
| for i, doc in enumerate(result["source_documents"]): |
| print(f"\n文档 {i+1}:") |
| print(f"内容:{doc.page_content}") |
| print(f"来源:{doc.metadata['source']}") |