引入rag

This commit is contained in:
lzybetter
2026-08-07 17:13:43 +08:00
parent b9ae90d4aa
commit b43afe4b4a
15 changed files with 1330 additions and 52 deletions
+147
View File
@@ -0,0 +1,147 @@
import os
from dotenv import load_dotenv
from langchain_community.document_loaders import DirectoryLoader, TextLoader
from langchain_text_splitters import RecursiveCharacterTextSplitter
from langchain_huggingface import HuggingFaceEmbeddings
from langchain_chroma import Chroma
# ==========================
# 配置
# ==========================
load_dotenv()
# ======================
# 配置
# ======================
KNOWLEDGE_PATH = "./knowledge"
VECTOR_DB_PATH = "./rag_db"
COLLECTION_NAME = "investment_knowledge"
EMBEDDING_MODEL = "BAAI/bge-small-zh-v1.5"
# ======================
# Embedding
# ======================
print("加载Embedding模型")
embeddings = HuggingFaceEmbeddings(
model_name=EMBEDDING_MODEL,
model_kwargs={
"device":"cpu"
}
)
# ======================
# 加载Markdown
# ======================
print("加载知识文件")
loader = DirectoryLoader(
KNOWLEDGE_PATH,
glob="**/*.md",
loader_cls=TextLoader,
loader_kwargs={
"encoding":"utf-8"
}
)
documents = loader.load()
print(
f"加载 {len(documents)} 个文件"
)
# ======================
# 自动生成metadata
# ======================
def build_metadata(doc):
path = doc.metadata["source"]
metadata={
"source":path
}
# rules
if "/rules/" in path:
metadata["category"]="rule"
metadata["knowledge_type"]="analysis_logic"
# industry
elif "/industry/" in path:
metadata["category"]="industry"
metadata["knowledge_type"]="industry_logic"
# company
elif "/company_cases/" in path:
metadata["category"]="company_case"
metadata["knowledge_type"]="company_fact"
# framework
elif "/investment_framework/" in path:
metadata["category"]="framework"
metadata["knowledge_type"]="workflow"
else:
metadata["category"]="unknown"
return metadata
for doc in documents:
doc.metadata.update(
build_metadata(doc)
)
# ======================
# 文本切分
# ======================
splitter = RecursiveCharacterTextSplitter(
chunk_size=500,
chunk_overlap=80,
separators=[
"\n\n",
"\n",
"",
""
]
)
texts = splitter.split_documents(
documents
)
print(
f"生成 {len(texts)} 个chunk"
)
# ======================
# 创建向量库
# ======================
print("生成Chroma")
vector_store = Chroma.from_documents(
documents=texts,
embedding=embeddings,
persist_directory=VECTOR_DB_PATH,
collection_name=COLLECTION_NAME
)
print("完成")