You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何为MultiVectorRetriever实现持久化数据库并解决无检索结果问题

基于PDF的RAG系统:MultiVectorRetriever检索无结果问题

我正尝试基于PDF构建RAG系统,提取其中的文本与表格,通过持久化数据库存储分片、表格、嵌入向量等数据,重新加载后使用MultiVectorRetriever。由于该Retriever需要docstore,我采用LocalFileStore批量存储父文档(原始PDF相关内容),虽能成功初始化MultiVectorRetriever,但检索时无法得到任何结果。以下是实现vectorstore和Retriever的代码:

utils.py

import os
from typing import List, Dict
from langchain.prompts import PromptTemplate
from langchain_openai import AzureChatOpenAI
from langchain.text_splitter import RecursiveCharacterTextSplitter
from unstructured.partition.pdf import partition_pdf
from dotenv import load_dotenv
import pytesseract

# Load environment variables from .env file
load_dotenv()

def setup_azure_llm() -> AzureChatOpenAI:
    """Initialize Azure OpenAI LLM for context generation."""
    return AzureChatOpenAI(
        deployment_name=os.getenv("DEPLOYMENT_NAME"),
        azure_endpoint=os.getenv("AZURE_OPENAI_ENDPOINT"),
        api_key=os.getenv("AZURE_OPENAI_API_KEY"),
        api_version=os.getenv("OPENAI_API_VERSION"),
        temperature=0.3
    )

def setup_embedding_model():
    """Initialize embedding model."""
    from sentence_transformers import SentenceTransformer
    return SentenceTransformer('multi-qa-MiniLM-L6-cos-v1')

def extract_pdf_text_and_tables(pdf_path: str) -> List[Dict]:
    """Extract text and tables from PDF using partition_pdf."""
    print(f"Attempting to load PDF: {pdf_path}")
    if not os.path.exists(pdf_path):
        print(f"Error: PDF file not found at {pdf_path}")
        return []
    
    try:
        # Use partition_pdf with optimized chunking
        elements = partition_pdf(
            filename=pdf_path,
            extract_images_in_pdf=False,  # Disable image extraction
            infer_table_structure=True,
            chunking_strategy="by_title",
            max_characters=2000,  # Reduced for faster processing
            new_after_n_chars=1800,
            combine_text_under_n_chars=1000
        )
        print(f"Successfully loaded {len(elements)} elements from {pdf_path}")
        if not elements:
            print(f"No elements extracted from {pdf_path}. Check PDF format or OCR dependencies.")
    except Exception as e:
        print(f"Failed to load PDF {pdf_path}: {str(e)}")
        print(f"Possible causes: Corrupted PDF, encryption, or missing Tesseract/Poppler.")
        return []

    items = []
    category_counts = {"text": 0, "table": 0}
    page_number = 1
    for element in elements:
        # Access page_number directly with hasattr
        metadata = getattr(element, "metadata", None)
        if metadata and hasattr(metadata, "page_number") and metadata.page_number:
            page_number = metadata.page_number
        
        content = str(element).strip()
        if not content:
            continue
        
        # Categorize based on element type
        if "unstructured.documents.elements.Table" in str(type(element)):
            category_counts["table"] += 1
            item = {
                'content': content,
                'page_number': page_number,
                'source': os.path.basename(pdf_path),
                'source_type': 'table'
            }
            items.append(item)
            print(f"Extracted table on page {page_number}: {content[:100] + '...' if len(content) > 100 else content}")
        elif "unstructured.documents.elements.CompositeElement" in str(type(element)):
            category_counts["text"] += 1
            item = {
                'content': content,
                'page_number': page_number,
                'source': os.path.basename(pdf_path),
                'source_type': 'text'
            }
            items.append(item)
            print(f"Extracted text on page {page_number}: {content[:100] + '...' if len(content) > 100 else content}")
        else:
            # Skip other element types (e.g., Title, Footer)
            continue
    
    print(f"Category counts for {pdf_path}: {category_counts}")
    if not items:
        print(f"No text or tables extracted from {pdf_path}")
    else:
        text_count = category_counts["text"]
        table_count = category_counts["table"]
        print(f"Summary for {pdf_path}: {text_count} text items, {table_count} table items, {len(items)} total items")
    return items

def split_text(text: str, chunk_size: int = 300, chunk_overlap: int = 50) -> List[str]:
    """Split text into chunks using RecursiveCharacterTextSplitter."""
    text_splitter = RecursiveCharacterTextSplitter(
        chunk_size=chunk_size,
        chunk_overlap=chunk_overlap,
        length_function=len,
        separators=["\n\n", "\n", ". ", " ", ""]
    )
    return text_splitter.split_text(text)

def generate_context(document: str, chunk: str, llm: AzureChatOpenAI) -> str:
    """Generate contextual prefix for a text or table chunk."""
    context_prompt = PromptTemplate(
        input_variables=["document", "chunk"],
        template="""<document>
{document}
</document>
Here is the chunk we want to situate within the whole document:
<chunk>
{chunk}
</chunk>
Please provide a short, succinct context (50-100 tokens) to situate this chunk within the overall document for improved search retrieval. Return only the context."""
    )
    try:
        prompt = context_prompt.format(document=document, chunk=chunk)
        context = llm.invoke(prompt).content
        return f"{context}\n\n{chunk}"
    except:
        return chunk

create_vectorstore.py

import os
import json
from collections import Counter
from langchain_community.vectorstores import Chroma
from langchain.retrievers.multi_vector import MultiVectorRetriever
from langchain.storage import LocalFileStore
from langchain_community.embeddings import HuggingFaceEmbeddings
from langchain_core.documents import Document
from utils import extract_pdf_text_and_tables, split_text, setup_azure_llm, generate_context
from tqdm import tqdm
import uuid

def create_vectorstore(pdf_dir="./pdfs",
                       vectorstore_dir="./vectorstore",
                       parent_store_dir="./parent_store",
                       use_llm_context=False, batch_size=5000):
    """Create a Chroma vector store with parent-child document relationship, persisting parent docs."""
    print("Initializing models...")
    llm = setup_azure_llm() if use_llm_context else None
    embeddings = HuggingFaceEmbeddings(model_name="multi-qa-MiniLM-L6-cos-v1")

    # Initialize Chroma vector store (persistent)
    vectorstore = Chroma(
        collection_name="rag_collection",
        embedding_function=embeddings,
        persist_directory=vectorstore_dir
    )

    # Persistent parent store
    parent_store = LocalFileStore(parent_store_dir)
    id_key = "doc_id"
    retriever = MultiVectorRetriever(
        vectorstore=vectorstore,
        docstore=parent_store,
        id_key=id_key,
        search_type="mmr",
        search_kwargs={"k": 5, "fetch_k": 20, "lambda_mult": 0.5}
    )

    parent_docs_serialized = []  # store as (key, bytes)
    child_chunks = []
    child_ids = []

    pdf_files = [f for f in os.listdir(pdf_dir) if f.endswith(".pdf")]
    if not pdf_files:
        print(f"No PDF files found in {pdf_dir}")
        return retriever
    print(f"📄 Found {len(pdf_files)} PDF(s) in {pdf_dir}")

    all_extracted_items = []
    for pdf_file in tqdm(pdf_files, desc="Processing PDFs"):
        pdf_path = os.path.join(pdf_dir, pdf_file)
        extracted_items = extract_pdf_text_and_tables(pdf_path)
        all_extracted_items.extend(extracted_items)

        for item in extracted_items:
            content = item["content"]
            metadata = {
                "source": item["source"],
                "page_number": item["page_number"],
                "type": item["source_type"]
            }
            parent_id = str(uuid.uuid4())
            parent_doc = Document(page_content=content, metadata={**metadata, id_key: parent_id})

            # Serialize Document to JSON bytes
            parent_docs_serialized.append((parent_id, json.dumps(parent_doc.dict()).encode("utf-8")))

            if item["source_type"] == "text":
                chunks = split_text(content, chunk_size=300, chunk_overlap=50)
                for chunk in chunks:
                    chunk_id = str(uuid.uuid4())
                    child_doc = Document(page_content=chunk, metadata={**metadata, id_key: chunk_id, "parent_id": parent_id})
                    child_chunks.append(child_doc)
                    child_ids.append(chunk_id)
            elif item["source_type"] == "table":
                chunk_id = str(uuid.uuid4())
                chunk_content = content
                if use_llm_context and llm:
                    chunk_content = generate_context(content, content, llm)
                child_doc = Document(page_content=chunk_content, metadata={**metadata, id_key: chunk_id, "parent_id": parent_id})
                child_chunks.append(child_doc)
                child_ids.append(chunk_id)

    # Store serialized parent documents
    if parent_docs_serialized:
        parent_store.mset(parent_docs_serialized)
        print(f"Stored {len(parent_docs_serialized)} parent documents in persistent docstore at {parent_store_dir}")

    # Add child documents to Chroma
    if child_chunks:
        total_chunks = len(child_chunks)
        print(f" Adding {total_chunks} child chunks to Chroma in batches of {batch_size}...")
        for i in tqdm(range(0, total_chunks, batch_size), desc="Indexing child chunks"):
            batch_chunks = child_chunks[i:i + batch_size]
            batch_ids = child_ids[i:i + batch_size]
            vectorstore.add_documents(batch_chunks, ids=batch_ids)
        print(f"Embedded {total_chunks} child chunks in Chroma vector store.")

    # Summary
    type_counts = Counter(item["source_type"] for item in all_extracted_items)
    print("\nDocument type counts (all PDFs):")
    print(f"   text: {type_counts.get('text', 0)}")
    print(f"   table: {type_counts.get('table', 0)}")
    print(f"Extracted {len(child_chunks)} child chunks and {len(parent_docs_serialized)} parent documents in total.")

    print(f"Vector store automatically persisted to {vectorstore_dir}")

    return retriever

if __name__ == "__main__":
    create_vectorstore()

retriever.py

import time
from typing import List, Tuple
from langchain_community.vectorstores import Chroma
from langchain_core.documents import Document
from langchain_community.embeddings import HuggingFaceEmbeddings
from dotenv import load_dotenv
from langchain.storage import LocalFileStore
from langchain.retrievers.multi_vector import MultiVectorRetriever

load_dotenv()

def setup_retriever(vector_store_dir: str = "./vectorstore", parent_store_dir: str = "./parent_store"):
    """Load a simple Chroma retriever without parent/child mapping."""
    try:
        embeddings = HuggingFaceEmbeddings(model_name="multi-qa-MiniLM-L6-cos-v1")
        vectorstore = Chroma(
            collection_name="rag_collection",
            embedding_function=embeddings,
            persist_directory=vector_store_dir
        )
        parent_store = LocalFileStore(parent_store_dir)
        retriever = MultiVectorRetriever(
            vectorstore=vectorstore,
            docstore=parent_store,
            id_key="doc_id",
            search_kwargs={"k": 10}
        )
        print(f"Retriever set up successfully with {len(vectorstore.get()['ids'])} child documents")
        return retriever
    except Exception as e:
        print(f"Failed to load retriever: {e}")
        return None

def perform_retrieval(retriever, query: str) -> Tuple[List[Document], float]:
    """Perform retrieval using the Chroma retriever."""
    start_time = time.time()
    try:
        print(f"Performing retrieval for query: {query}")
        results = retriever.invoke(query)  # LangChain now prefers invoke()

        if not results:
            print("Retriever returned no results.")
        else:
            print(f"Retrieved {len(results)} documents.")
            for i, doc in enumerate(results[:5], 1):
                content = doc.page_content.strip().replace("\n", " ") if doc.page_content else "[No content]"
                if len(content) > 100:
                    content = content[:100] + "..."
                print(f"  Top {i}: {content} "
                      f"[Source: {doc.metadata.get('source', 'Unknown')}, "
                      f"Page: {doc.metadata.get('page_number', 'Unknown')}, "
                      f"Type: {doc.metadata.get('type', 'Unknown')}]")

        return results, time.time() - start_time
    except Exception as e:
        print(f"Error during retrieval: {e}")
        return [], time.time() - start_time

def evaluate_retrieval(results: List[Document], ground_truth: dict, query: str, k: int = 2) -> dict:
    """Evaluate retrieval performance."""
    try:
        relevant_docs = ground_truth.get(query, [])
        retrieved_docs = [
            f"{doc.metadata.get('source', 'Unknown')}_page_{doc.metadata.get('page_number', 'Unknown')}"
            for doc in results[:k]
        ]
        true_positives = len(set(retrieved_docs).intersection(set(relevant_docs)))
        precision = true_positives / len(retrieved_docs) if retrieved_docs else 0.0
        recall = true_positives / len(relevant_docs) if relevant_docs else 0.0
        f1 = 2 * (precision * recall) / (precision + recall) if (precision + recall) > 0 else 0.0
        mrr = 0.0
        for i, doc in enumerate(retrieved_docs, 1):
            if doc in relevant_docs:
                mrr = 1.0 / i
                break
        return {
            "precision": precision,
            "recall": recall,
            "mrr": mrr,
            "f1": f1,
            "ndcg": 0.0,
            "map": precision,
            "precision@2": precision,
            "recall@2": recall,
            "ap": precision,
            "hit_rate": 1.0 if true_positives > 0 else 0.0
        }
    except Exception as e:
        print(f"Error evaluating retrieval: {e}")
        return {
            "precision": 0.0,
            "recall": 0.0,
            "mrr": 0.0,
            "f1": 0.0,
            "ndcg": 0.0,
            "map": 0.0,
            "precision@2": 0.0,
            "recall@2": 0.0,
            "ap": 0.0,
            "hit_rate": 0.0
        }
问题分析与解决方案

核心问题定位

  1. 子父文档关联错误:初始化Retriever时指定id_key="doc_id",但子文档的doc_id是自身UUID,而非父文档ID。MultiVectorRetriever需要通过子文档的id_key字段获取父文档ID,才能从docstore中取出对应父文档返回。
  2. 序列化/反序列化不匹配:用JSON序列化Document后,加载时未正确反序列化为Document实例,导致docstore数据无法被Retriever识别。

修复步骤

1. 修正子文档的id_key关联

在create_vectorstore.py中,将子文档的id_key值改为父文档ID:

# 原错误代码
child_doc = Document(page_content=chunk, metadata={**metadata, id_key: chunk_id, "parent_id": parent_id})

# 修改后
child_doc = Document(page_content=chunk, metadata={**metadata, id_key: parent_id, "child_id": chunk_id})

2. 完善Document序列化与反序列化

在utils.py中添加序列化工具函数:

from langchain_core.documents import Document
import json

def serialize_doc(doc: Document) -> bytes:
    return json.dumps({
        "page_content": doc.page_content,
        "metadata": doc.metadata
    }).encode("utf-8")

def deserialize_doc(data: bytes) -> Document:
    doc
相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.12 11:24:50