如何为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 }
问题分析与解决方案
核心问题定位
- 子父文档关联错误:初始化Retriever时指定
id_key="doc_id",但子文档的doc_id是自身UUID,而非父文档ID。MultiVectorRetriever需要通过子文档的id_key字段获取父文档ID,才能从docstore中取出对应父文档返回。 - 序列化/反序列化不匹配:用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
相关产品推荐
相关产品推荐

