S SmartDocs
Chuỗi bài: RAG AI python 152 dòng · Cập nhật 2026-05-08

ingest.py

RAG_AI/src/ingest.py

"""文件索引腳本:把 data/ 下的所有文件載入、切塊、向量化、存入 Chroma。

支援格式:PDF (.pdf)、Markdown (.md)、純文字 (.txt)、Word (.docx)
使用方式:
    python -m src.ingest                # 增量索引(新檔加入、舊檔更新)
    python -m src.ingest --rebuild      # 全量重建(清空再建)
"""

from __future__ import annotations

import argparse
import sys
from pathlib import Path
from typing import Iterable

from langchain.embeddings import CacheBackedEmbeddings
from langchain.indexes import SQLRecordManager, index
from langchain.storage import LocalFileStore
from langchain_chroma import Chroma
from langchain_community.document_loaders import (
    Docx2txtLoader,
    PyPDFLoader,
    TextLoader,
)
from langchain_core.documents import Document
from langchain_openai import OpenAIEmbeddings
from langchain_text_splitters import RecursiveCharacterTextSplitter

from src.config import config, ensure_dirs

LOADER_MAP = {
    ".pdf": PyPDFLoader,
    ".md": lambda p: TextLoader(p, encoding="utf-8"),
    ".txt": lambda p: TextLoader(p, encoding="utf-8"),
    ".docx": Docx2txtLoader,
}


def build_embeddings() -> CacheBackedEmbeddings:
    """帶本地快取的 embedding,重複內容不會再打 OpenAI API(省錢、更快)。"""
    base = OpenAIEmbeddings(model=config.embedding_model)
    store = LocalFileStore(str(config.cache_dir))
    return CacheBackedEmbeddings.from_bytes_store(
        base, store, namespace=config.embedding_model
    )


def load_documents(data_dir: Path) -> list[Document]:
    """掃描 data_dir 載入所有支援格式的檔案。"""
    docs: list[Document] = []
    files = [f for f in data_dir.rglob("*") if f.is_file() and f.suffix.lower() in LOADER_MAP]

    if not files:
        print(f"[警告] {data_dir} 找不到任何可索引的文件")
        return docs

    for path in files:
        loader_cls = LOADER_MAP[path.suffix.lower()]
        try:
            loaded = loader_cls(str(path)).load()
        except Exception as exc:
            print(f"[錯誤] 無法載入 {path.name}: {exc}")
            continue

        # 正規化 metadata:用相對路徑當 source 才能在不同機器一致
        rel_source = str(path.relative_to(data_dir))
        for d in loaded:
            d.metadata["source"] = rel_source
            d.metadata.setdefault("file_type", path.suffix.lower().lstrip("."))
        docs.extend(loaded)
        print(f"  載入 {path.name} ({len(loaded)} 段)")

    return docs


def split_documents(docs: Iterable[Document]) -> list[Document]:
    """用 token-aware 的遞迴切塊器,避免超過 embedding 上限。"""
    splitter = RecursiveCharacterTextSplitter.from_tiktoken_encoder(
        model_name=config.embedding_model,
        chunk_size=config.chunk_size,
        chunk_overlap=config.chunk_overlap,
        # 中文友善的分隔符優先序
        separators=["\n## ", "\n### ", "\n\n", "\n", "。", "!", "?", ".", " ", ""],
    )
    return splitter.split_documents(list(docs))


def ingest(rebuild: bool = False) -> dict:
    """主流程。回傳 LangChain index() 的結果摘要。"""
    if not config.has_openai_key:
        print("[錯誤] 找不到 OPENAI_API_KEY,請複製 .env.example 為 .env 並填入 API Key")
        sys.exit(1)

    ensure_dirs()

    print(f"\n=== 階段 1:載入文件(從 {config.data_dir})===")
    docs = load_documents(config.data_dir)
    if not docs:
        print("沒有文件可索引,結束。")
        return {}
    print(f"共載入 {len(docs)} 份段落")

    print("\n=== 階段 2:切塊 ===")
    chunks = split_documents(docs)
    print(f"切成 {len(chunks)} 個 chunks(chunk_size={config.chunk_size})")

    print("\n=== 階段 3:建立 Vector Store ===")
    embeddings = build_embeddings()
    vectorstore = Chroma(
        collection_name=config.collection_name,
        embedding_function=embeddings,
        persist_directory=str(config.db_dir),
    )

    record_manager = SQLRecordManager(
        f"chroma/{config.collection_name}",
        db_url=f"sqlite:///{config.record_db}",
    )
    record_manager.create_schema()

    cleanup_mode = "full" if rebuild else "incremental"
    print(f"  使用 cleanup={cleanup_mode}")

    result = index(
        chunks,
        record_manager,
        vectorstore,
        cleanup=cleanup_mode,
        source_id_key="source",
    )

    print("\n=== 完成 ===")
    print(f"  新增:{result['num_added']}")
    print(f"  更新:{result['num_updated']}")
    print(f"  刪除:{result['num_deleted']}")
    print(f"  跳過:{result['num_skipped']}")
    return result


def main() -> None:
    parser = argparse.ArgumentParser(description="索引文件到向量資料庫")
    parser.add_argument(
        "--rebuild",
        action="store_true",
        help="清空既有索引重新建立(預設為增量更新)",
    )
    args = parser.parse_args()
    ingest(rebuild=args.rebuild)


if __name__ == "__main__":
    main()

Bài viết liên quan