Files
Financial_Analysis_for_Stocks/build_rag.py
T
2026-08-07 17:13:43 +08:00

147 lines
2.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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("完成")