147 lines
2.5 KiB
Python
147 lines
2.5 KiB
Python
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("完成") |