引入rag
This commit is contained in:
+147
@@ -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("完成")
|
||||
Reference in New Issue
Block a user