Serie: RAG AI
python
158 righe
· Aggiornato 2026-05-08
rag.py
RAG_AI/src/rag.py
"""RAG 主邏輯:Hybrid Search(BM25 + Vector)+ 可選 Re-ranker + 來源追溯。
關鍵設計:
- 用 Hybrid Search 結合稀疏(BM25)與密集(Vector)檢索,提升召回率
- Re-ranker 預設關閉(避免下載 2GB 模型),可在 .env 開啟
- 中文 BM25 用 jieba 斷詞
- 回傳 chain 與 retriever 兩個物件,UI 可分別使用
"""
from __future__ import annotations
import sys
from typing import Iterable
import jieba
from langchain.retrievers import ContextualCompressionRetriever, EnsembleRetriever
from langchain_chroma import Chroma
from langchain_community.retrievers import BM25Retriever
from langchain_core.documents import Document
from langchain_core.output_parsers import StrOutputParser
from langchain_core.runnables import Runnable, RunnablePassthrough
from langchain_openai import ChatOpenAI, OpenAIEmbeddings
from src.config import config
from src.prompts import RAG_PROMPT
def _chinese_tokenizer(text: str) -> list[str]:
"""中文 BM25 必須先斷詞,否則整段被當成一個 token,效果極差。"""
return [t for t in jieba.cut(text) if t.strip()]
def _load_all_chunks(vectorstore: Chroma) -> list[Document]:
"""從 Chroma 讀出所有已索引的 chunk,給 BM25 建立 inverted index 用。
BM25 需要看到全部資料才能算 IDF;Vector retriever 則靠 ANN,不必載入全部。
"""
raw = vectorstore.get()
return [
Document(page_content=text, metadata=meta or {})
for text, meta in zip(raw["documents"], raw["metadatas"])
]
def build_retriever() -> Runnable:
"""建立完整檢索管線:Hybrid Search → (可選) Re-rank。"""
embeddings = OpenAIEmbeddings(model=config.embedding_model)
vectorstore = Chroma(
collection_name=config.collection_name,
embedding_function=embeddings,
persist_directory=str(config.db_dir),
)
chunks = _load_all_chunks(vectorstore)
if not chunks:
print(
"[錯誤] 向量資料庫是空的。請先跑 `python -m src.ingest` 索引文件。",
file=sys.stderr,
)
sys.exit(1)
bm25 = BM25Retriever.from_documents(chunks, preprocess_func=_chinese_tokenizer)
bm25.k = config.top_k_retrieve
vector_retriever = vectorstore.as_retriever(
search_kwargs={"k": config.top_k_retrieve}
)
hybrid = EnsembleRetriever(
retrievers=[bm25, vector_retriever],
weights=[config.bm25_weight, config.vector_weight],
)
if not config.use_reranker:
return hybrid
# Re-ranker 啟用:套上 Cross-Encoder 精排,把 top_k 縮成 top_n
try:
from langchain.retrievers.document_compressors import CrossEncoderReranker
from langchain_community.cross_encoders import HuggingFaceCrossEncoder
except ImportError:
print(
"[警告] Re-ranker 啟用但缺少 sentence-transformers。"
"請執行 `pip install sentence-transformers torch`,或在 .env 設 USE_RERANKER=false。",
file=sys.stderr,
)
return hybrid
print(f" 載入 Re-ranker:{config.reranker_model}(首次會下載 ~2GB)")
reranker = CrossEncoderReranker(
model=HuggingFaceCrossEncoder(model_name=config.reranker_model),
top_n=config.top_n_rerank,
)
return ContextualCompressionRetriever(
base_compressor=reranker, base_retriever=hybrid
)
def format_docs(docs: Iterable[Document]) -> str:
"""把多個 chunks 組成一段給 LLM 看的文字,附上來源資訊。"""
return "\n\n".join(
f"[來源: {d.metadata.get('source', '?')} | 頁: {d.metadata.get('page', '-')}]\n{d.page_content}"
for d in docs
)
def build_chain() -> tuple[Runnable, Runnable]:
"""組裝完整 RAG chain。回傳 (chain, retriever) 方便 UI 分別呼叫。"""
if not config.has_openai_key:
print("[錯誤] 找不到 OPENAI_API_KEY,請設定 .env", file=sys.stderr)
sys.exit(1)
retriever = build_retriever()
llm = ChatOpenAI(model=config.chat_model, temperature=0)
chain = (
{
"context": retriever | format_docs,
"question": RunnablePassthrough(),
}
| RAG_PROMPT
| llm
| StrOutputParser()
)
return chain, retriever
def cli() -> None:
"""互動式 CLI,可串流回應。"""
print("\n=== RAG 智能問答 CLI ===")
print(f"模型:{config.chat_model} | Re-ranker:{'on' if config.use_reranker else 'off'}")
print("輸入 'exit' 或 Ctrl-C 離開\n")
chain, retriever = build_chain()
try:
while True:
q = input("\n問題> ").strip()
if q.lower() in {"exit", "quit", ""}:
break
print("\n回答:", end="", flush=True)
for token in chain.stream(q):
print(token, end="", flush=True)
print()
print("\n--- 引用來源 ---")
for i, doc in enumerate(retriever.invoke(q), 1):
src = doc.metadata.get("source", "?")
page = doc.metadata.get("page", "-")
preview = doc.page_content[:80].replace("\n", " ")
print(f" [{i}] {src} (頁 {page}): {preview}...")
except (KeyboardInterrupt, EOFError):
print("\n再見!")
if __name__ == "__main__":
cli()
Articoli correlati
RAG AI
python
Aggiornato 2026-05-08
main.py
main.py — python source code from the RAG AI learning materials (RAG_AI/main.py).
Leggi l'articolo →
RAG AI
python
Aggiornato 2026-05-08
__init__.py
__init__.py — python source code from the RAG AI learning materials (RAG_AI/src/__init__.py).
Leggi l'articolo →
RAG AI
python
Aggiornato 2026-05-08
app.py
app.py — python source code from the RAG AI learning materials (RAG_AI/src/app.py).
Leggi l'articolo →
RAG AI
python
Aggiornato 2026-05-08
config.py
config.py — python source code from the RAG AI learning materials (RAG_AI/src/config.py).
Leggi l'articolo →
RAG AI
python
Aggiornato 2026-05-08
evaluate.py
evaluate.py — python source code from the RAG AI learning materials (RAG_AI/src/evaluate.py).
Leggi l'articolo →
RAG AI
python
Aggiornato 2026-05-08
ingest.py
ingest.py — python source code from the RAG AI learning materials (RAG_AI/src/ingest.py).
Leggi l'articolo →