S SmartDocs
시리즈: RAG AI python 158 줄 · 업데이트 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()

관련 글