Spaces:
Running on Zero
Running on Zero
File size: 4,860 Bytes
0828c2c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 | """Vector-based retrieval strategy using embeddings."""
import logging
from pathlib import Path
from typing import Any
from langchain_core.documents import Document
from langchain_core.retrievers import BaseRetriever
from ..base import BaseRetrieverStrategy
from ..factory import RetrieverFactory
logger = logging.getLogger(__name__)
@RetrieverFactory.register("vector")
class VectorStrategy(BaseRetrieverStrategy):
"""Vector-based retrieval using semantic similarity.
This strategy uses embedding models to convert documents and queries
into dense vectors, then retrieves documents based on cosine similarity.
It wraps the existing VectorStoreManager for backward compatibility.
Configuration:
retrieval.vector.search_type: 'similarity' or 'mmr'
retrieval.vector.k: Number of documents to retrieve
retrieval.vector.search_kwargs: Additional search parameters
"""
def __init__(self, config: dict[str, Any]):
"""Initialize vector strategy.
Args:
config: Configuration dictionary
"""
super().__init__(config)
# Lazy import to avoid circular dependencies
from src.vectorstore import VectorStoreManager
self.vectorstore_manager = VectorStoreManager()
self._retriever: BaseRetriever | None = None
# Get vector-specific config
vector_config = config.get("retrieval", {}).get("vector", {})
self.search_type = vector_config.get("search_type", "similarity")
self.k = vector_config.get("k", 4)
self.search_kwargs = vector_config.get("search_kwargs", {})
@property
def name(self) -> str:
"""Return strategy identifier."""
return "vector"
def build_index(self, documents: list[Document]) -> None:
"""Build vector index from documents.
Args:
documents: List of Document objects to index
"""
logger.info(f"Building vector index with {len(documents)} documents...")
self.vectorstore_manager.create_vectorstore(documents)
self._is_initialized = True
logger.info("Vector index built successfully")
def load_index(self) -> bool:
"""Load existing vector index.
Returns:
True if loaded successfully
"""
try:
persist_dir = Path(self.vectorstore_manager.persist_directory)
if not persist_dir.exists():
logger.warning(f"Vector store not found at {persist_dir}")
return False
self.vectorstore_manager.load_vectorstore()
self._is_initialized = True
logger.info("Vector index loaded successfully")
return True
except Exception as e:
logger.error(f"Error loading vector index: {e}")
return False
def retrieve(self, query: str, k: int | None = None) -> list[Document]:
"""Retrieve relevant documents using vector similarity.
Args:
query: Search query
k: Number of documents to retrieve (uses config default if None)
Returns:
List of relevant documents
"""
if not self._is_initialized:
self.load_index()
k = k or self.k
return self.vectorstore_manager.similarity_search(query, k=k)
def as_retriever(self, **kwargs: Any) -> BaseRetriever:
"""Get LangChain-compatible retriever.
Args:
**kwargs: Override search parameters
Returns:
BaseRetriever instance
"""
if not self._is_initialized:
self.load_index()
# Merge config with kwargs
search_kwargs = {**self.search_kwargs, "k": self.k}
search_kwargs.update(kwargs.get("search_kwargs", {}))
return self.vectorstore_manager.get_retriever(
search_type=kwargs.get("search_type", self.search_type),
**search_kwargs,
)
def get_index_stats(self) -> dict[str, Any]:
"""Get vector index statistics.
Returns:
Dictionary with index statistics
"""
stats = super().get_index_stats()
if self._is_initialized and self.vectorstore_manager.vectorstore:
try:
collection_count = self.vectorstore_manager.vectorstore._collection.count()
stats.update(
{
"num_documents": collection_count,
"persist_directory": self.vectorstore_manager.persist_directory,
"collection_name": self.vectorstore_manager.collection_name,
"search_type": self.search_type,
}
)
except Exception as e:
logger.warning(f"Could not get collection stats: {e}")
return stats
|