mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-06 02:48:22 +00:00
feat(vector-store): add threshold-based filtering to vector search
This commit is contained in:
parent
9e50e71b52
commit
b0b800fbd8
7 changed files with 304 additions and 27 deletions
|
|
@ -73,8 +73,21 @@ class BaseVectorStore(ABC):
|
|||
"""Add one or more vector nodes into the current collection."""
|
||||
|
||||
@abstractmethod
|
||||
async def search(self, query: str, limit: int = 5, filters: dict | None = None, **kwargs) -> list[VectorNode]:
|
||||
"""Find the most similar vector nodes based on a text query."""
|
||||
async def search(
|
||||
self,
|
||||
query: str,
|
||||
limit: int = 5,
|
||||
candidates: int | None = None,
|
||||
filters: dict | None = None,
|
||||
threshold: float | None = None,
|
||||
**kwargs,
|
||||
) -> list[VectorNode]:
|
||||
"""Find the most similar vector nodes based on a text query.
|
||||
|
||||
When threshold is None, uses default behavior.
|
||||
When threshold is set, searches max(candidates, limit) nodes,
|
||||
filters by threshold, then returns top limit results.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def delete(self, vector_ids: str | list[str], **kwargs) -> None:
|
||||
|
|
|
|||
|
|
@ -307,21 +307,35 @@ class ChromaVectorStore(BaseVectorStore):
|
|||
self,
|
||||
query: str,
|
||||
limit: int = 5,
|
||||
candidates: int | None = None,
|
||||
filters: dict | None = None,
|
||||
threshold: float | None = None,
|
||||
**kwargs,
|
||||
) -> list[VectorNode]:
|
||||
"""Search for the most similar vector nodes based on a text query."""
|
||||
"""Search for the most similar vector nodes based on a text query.
|
||||
|
||||
When threshold is None, uses default behavior.
|
||||
When threshold is set, searches max(candidates, limit) nodes,
|
||||
filters by threshold, then returns top limit results.
|
||||
"""
|
||||
query_vector = await self.get_embedding(query)
|
||||
where_clause = self._generate_where_clause(filters)
|
||||
include_embeddings = kwargs.get("include_embeddings", False)
|
||||
|
||||
# When threshold is set, search more candidates
|
||||
if threshold is not None:
|
||||
effective_candidates = candidates if candidates is not None else limit * 2
|
||||
search_limit = max(effective_candidates, limit)
|
||||
else:
|
||||
search_limit = limit
|
||||
|
||||
def _search():
|
||||
include: list = ["documents", "metadatas", "distances"]
|
||||
if include_embeddings:
|
||||
include.append("embeddings")
|
||||
return self.collection.query(
|
||||
query_embeddings=[query_vector],
|
||||
n_results=limit,
|
||||
n_results=search_limit,
|
||||
where=where_clause,
|
||||
include=include,
|
||||
)
|
||||
|
|
@ -329,9 +343,11 @@ class ChromaVectorStore(BaseVectorStore):
|
|||
results = await self._run_sync_in_executor(_search)
|
||||
nodes = self._parse_results(results, include_score=True)
|
||||
|
||||
score_threshold = kwargs.get("score_threshold")
|
||||
if score_threshold is not None:
|
||||
nodes = [n for n in nodes if n.metadata.get("score", 0) >= score_threshold]
|
||||
# Apply threshold filtering if specified
|
||||
if threshold is not None:
|
||||
nodes = [n for n in nodes if n.metadata.get("score", 0) >= threshold]
|
||||
nodes = nodes[:limit]
|
||||
|
||||
return nodes
|
||||
|
||||
async def delete(self, vector_ids: str | list[str], **kwargs):
|
||||
|
|
|
|||
|
|
@ -238,31 +238,48 @@ class ESVectorStore(BaseVectorStore):
|
|||
self,
|
||||
query: str,
|
||||
limit: int = 5,
|
||||
candidates: int | None = None,
|
||||
filters: dict | None = None,
|
||||
threshold: float | None = None,
|
||||
**kwargs,
|
||||
) -> list[VectorNode]:
|
||||
"""Perform a kNN similarity search based on a text query.
|
||||
|
||||
When threshold is None, uses default behavior.
|
||||
When threshold is set, searches max(candidates, limit) nodes,
|
||||
filters by threshold, then returns top limit results.
|
||||
|
||||
Args:
|
||||
query: The text to search for.
|
||||
limit: Maximum number of nearest neighbors to return.
|
||||
filters: Metadata filters for exact match or 'IN' operations.
|
||||
**kwargs: Search parameters like num_candidates or score_threshold.
|
||||
candidates: Number of candidates to consider during search.
|
||||
filters: Metadata filters for exact match or range operations.
|
||||
threshold: Minimum score threshold for results.
|
||||
**kwargs: Additional search parameters.
|
||||
|
||||
Returns:
|
||||
List of VectorNode objects ordered by similarity.
|
||||
"""
|
||||
query_vector = await self.get_embedding(query)
|
||||
num_candidates = kwargs.get("num_candidates", limit * 2)
|
||||
|
||||
# When threshold is set, search more candidates
|
||||
if threshold is not None:
|
||||
effective_candidates = candidates if candidates is not None else limit * 2
|
||||
search_candidates = max(effective_candidates, limit)
|
||||
search_limit = search_candidates
|
||||
else:
|
||||
effective_candidates = candidates if candidates is not None else limit * 2
|
||||
search_candidates = effective_candidates
|
||||
search_limit = limit
|
||||
|
||||
search_query: dict = {
|
||||
"knn": {
|
||||
"field": "vector",
|
||||
"query_vector": query_vector,
|
||||
"k": limit,
|
||||
"num_candidates": num_candidates,
|
||||
"k": search_limit,
|
||||
"num_candidates": search_candidates,
|
||||
},
|
||||
"size": limit,
|
||||
"size": search_limit,
|
||||
}
|
||||
|
||||
if filters:
|
||||
|
|
@ -300,6 +317,11 @@ class ESVectorStore(BaseVectorStore):
|
|||
node.metadata["score"] = hit["_score"]
|
||||
results.append(node)
|
||||
|
||||
# Apply threshold filtering if specified
|
||||
if threshold is not None:
|
||||
results = [n for n in results if n.metadata.get("score", 0) >= threshold]
|
||||
results = results[:limit]
|
||||
|
||||
return results
|
||||
|
||||
async def delete(self, vector_ids: str | list[str], refresh: bool = True, **kwargs):
|
||||
|
|
|
|||
|
|
@ -193,10 +193,17 @@ class LocalVectorStore(BaseVectorStore):
|
|||
self,
|
||||
query: str,
|
||||
limit: int = 5,
|
||||
candidates: int | None = None,
|
||||
filters: dict | None = None,
|
||||
threshold: float | None = None,
|
||||
**kwargs,
|
||||
) -> list[VectorNode]:
|
||||
"""Search for nodes similar to the query using brute-force cosine similarity."""
|
||||
"""Search for nodes similar to the query using brute-force cosine similarity.
|
||||
|
||||
When threshold is None, uses default behavior.
|
||||
When threshold is set, searches max(candidates, limit) nodes,
|
||||
filters by threshold, then returns top limit results.
|
||||
"""
|
||||
query_vector = await self.get_embedding(query)
|
||||
all_nodes = self._load_all_nodes()
|
||||
filtered_nodes = [node for node in all_nodes if self._match_filters(node, filters)]
|
||||
|
|
@ -215,11 +222,17 @@ class LocalVectorStore(BaseVectorStore):
|
|||
|
||||
scored_nodes.sort(key=lambda x: x[1], reverse=True)
|
||||
|
||||
score_threshold = kwargs.get("score_threshold")
|
||||
if score_threshold is not None:
|
||||
scored_nodes = [(node, score) for node, score in scored_nodes if score >= score_threshold]
|
||||
# Apply threshold filtering if specified
|
||||
if threshold is not None:
|
||||
# When threshold is set, consider more candidates
|
||||
effective_candidates = candidates if candidates is not None else limit * 2
|
||||
search_limit = max(effective_candidates, limit)
|
||||
scored_nodes = scored_nodes[:search_limit]
|
||||
scored_nodes = [(node, score) for node, score in scored_nodes if score >= threshold]
|
||||
scored_nodes = scored_nodes[:limit]
|
||||
else:
|
||||
scored_nodes = scored_nodes[:limit]
|
||||
|
||||
scored_nodes = scored_nodes[:limit]
|
||||
results = []
|
||||
for node, score in scored_nodes:
|
||||
node.metadata["score"] = score
|
||||
|
|
|
|||
|
|
@ -333,10 +333,17 @@ class PGVectorStore(BaseVectorStore):
|
|||
self,
|
||||
query: str,
|
||||
limit: int = 5,
|
||||
candidates: int | None = None,
|
||||
filters: dict | None = None,
|
||||
threshold: float | None = None,
|
||||
**kwargs,
|
||||
) -> list[VectorNode]:
|
||||
"""Perform vector similarity search with optional metadata filtering."""
|
||||
"""Perform vector similarity search with optional metadata filtering.
|
||||
|
||||
When threshold is None, uses default behavior.
|
||||
When threshold is set, searches max(candidates, limit) nodes,
|
||||
filters by threshold, then returns top limit results.
|
||||
"""
|
||||
await self._ensure_collection_exists()
|
||||
|
||||
query_vector = await self.get_embedding(query)
|
||||
|
|
@ -351,6 +358,13 @@ class PGVectorStore(BaseVectorStore):
|
|||
new_placeholder = f"${i + 1}"
|
||||
filter_clause = re.sub(rf"\${i}\b", new_placeholder, filter_clause)
|
||||
|
||||
# When threshold is set, search more candidates
|
||||
if threshold is not None:
|
||||
effective_candidates = candidates if candidates is not None else limit * 2
|
||||
search_limit = max(effective_candidates, limit)
|
||||
else:
|
||||
search_limit = limit
|
||||
|
||||
async with pool.acquire() as conn:
|
||||
sql = f"""
|
||||
SELECT id, content, vector, metadata, vector <=> $1::vector AS distance
|
||||
|
|
@ -359,14 +373,17 @@ class PGVectorStore(BaseVectorStore):
|
|||
ORDER BY distance
|
||||
LIMIT ${len(filter_params) + 2}
|
||||
"""
|
||||
rows = await conn.fetch(sql, vector_str, *filter_params, limit)
|
||||
rows = await conn.fetch(sql, vector_str, *filter_params, search_limit)
|
||||
|
||||
results = []
|
||||
score_threshold = kwargs.get("score_threshold")
|
||||
|
||||
for row in rows:
|
||||
distance = row["distance"]
|
||||
if score_threshold is not None and distance > score_threshold:
|
||||
# Convert distance to score (1 - distance for cosine)
|
||||
score = 1 - distance
|
||||
|
||||
# Apply threshold filtering if specified
|
||||
if threshold is not None and score < threshold:
|
||||
continue
|
||||
|
||||
vector_data = None
|
||||
|
|
@ -379,7 +396,7 @@ class PGVectorStore(BaseVectorStore):
|
|||
if isinstance(metadata, str):
|
||||
metadata = json.loads(metadata)
|
||||
|
||||
metadata["score"] = 1 - distance
|
||||
metadata["score"] = score
|
||||
metadata["_distance"] = distance
|
||||
|
||||
node = VectorNode(
|
||||
|
|
@ -390,6 +407,10 @@ class PGVectorStore(BaseVectorStore):
|
|||
)
|
||||
results.append(node)
|
||||
|
||||
# Apply limit after threshold filtering
|
||||
if threshold is not None:
|
||||
results = results[:limit]
|
||||
|
||||
return results
|
||||
|
||||
async def delete(self, vector_ids: str | list[str], **kwargs):
|
||||
|
|
|
|||
|
|
@ -325,20 +325,33 @@ class QdrantVectorStore(BaseVectorStore):
|
|||
self,
|
||||
query: str,
|
||||
limit: int = 5,
|
||||
candidates: int | None = None,
|
||||
filters: dict | None = None,
|
||||
threshold: float | None = None,
|
||||
**kwargs: Any,
|
||||
) -> list[VectorNode]:
|
||||
"""Search for the most similar vectors based on a text query."""
|
||||
"""Search for the most similar vectors based on a text query.
|
||||
|
||||
When threshold is None, uses default behavior.
|
||||
When threshold is set, searches max(candidates, limit) nodes,
|
||||
filters by threshold, then returns top limit results.
|
||||
"""
|
||||
query_vector = await self.get_embedding(query)
|
||||
query_filter = self._create_filter(filters) if filters else None
|
||||
score_threshold = kwargs.get("score_threshold", None)
|
||||
|
||||
# When threshold is set, search more candidates
|
||||
if threshold is not None:
|
||||
effective_candidates = candidates if candidates is not None else limit * 2
|
||||
search_limit = max(effective_candidates, limit)
|
||||
else:
|
||||
search_limit = limit
|
||||
|
||||
results = await self.client.query_points(
|
||||
collection_name=self.collection_name,
|
||||
query=query_vector,
|
||||
query_filter=query_filter,
|
||||
limit=limit,
|
||||
score_threshold=score_threshold,
|
||||
limit=search_limit,
|
||||
score_threshold=threshold,
|
||||
)
|
||||
|
||||
nodes = []
|
||||
|
|
@ -353,6 +366,10 @@ class QdrantVectorStore(BaseVectorStore):
|
|||
node.metadata["score"] = point.score
|
||||
nodes.append(node)
|
||||
|
||||
# Apply limit after threshold filtering (Qdrant already filters by threshold)
|
||||
if threshold is not None:
|
||||
nodes = nodes[:limit]
|
||||
|
||||
return nodes
|
||||
|
||||
async def delete(self, vector_ids: str | list[str], **kwargs: Any):
|
||||
|
|
|
|||
|
|
@ -414,6 +414,180 @@ async def test_search_with_multiple_filters(store: BaseVectorStore, _store_name:
|
|||
logger.info("✓ Multiple filters search test passed")
|
||||
|
||||
|
||||
async def test_search_with_threshold(store: BaseVectorStore, _store_name: str):
|
||||
"""Test vector search with threshold parameter.
|
||||
|
||||
When threshold is None, uses default behavior.
|
||||
When threshold is set, searches max(candidates, limit) nodes,
|
||||
filters by threshold, then returns top limit results.
|
||||
"""
|
||||
logger.info("=" * 20 + " SEARCH WITH THRESHOLD TEST " + "=" * 20)
|
||||
|
||||
# Clean up any existing test data first
|
||||
try:
|
||||
existing_nodes = await store.list(filters={"test_type": "threshold_test"})
|
||||
if existing_nodes:
|
||||
await store.delete([node.vector_id for node in existing_nodes])
|
||||
logger.info(f"Cleaned up {len(existing_nodes)} existing test nodes")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to clean up existing nodes: {e}")
|
||||
|
||||
# Insert test nodes with varied content for different similarity scores
|
||||
threshold_nodes = [
|
||||
VectorNode(
|
||||
vector_id="threshold_node_1",
|
||||
content="Machine learning and artificial intelligence are the future of technology.",
|
||||
metadata={"test_type": "threshold_test", "relevance": "high"},
|
||||
),
|
||||
VectorNode(
|
||||
vector_id="threshold_node_2",
|
||||
content="Deep learning uses neural networks for artificial intelligence applications.",
|
||||
metadata={"test_type": "threshold_test", "relevance": "high"},
|
||||
),
|
||||
VectorNode(
|
||||
vector_id="threshold_node_3",
|
||||
content="Natural language processing is a branch of AI technology.",
|
||||
metadata={"test_type": "threshold_test", "relevance": "medium"},
|
||||
),
|
||||
VectorNode(
|
||||
vector_id="threshold_node_4",
|
||||
content="Computer vision enables machines to interpret visual information.",
|
||||
metadata={"test_type": "threshold_test", "relevance": "medium"},
|
||||
),
|
||||
VectorNode(
|
||||
vector_id="threshold_node_5",
|
||||
content="Software development requires good coding practices.",
|
||||
metadata={"test_type": "threshold_test", "relevance": "low"},
|
||||
),
|
||||
VectorNode(
|
||||
vector_id="threshold_node_6",
|
||||
content="Database management systems store and retrieve data efficiently.",
|
||||
metadata={"test_type": "threshold_test", "relevance": "low"},
|
||||
),
|
||||
VectorNode(
|
||||
vector_id="threshold_node_7",
|
||||
content="Cloud computing provides scalable infrastructure for applications.",
|
||||
metadata={"test_type": "threshold_test", "relevance": "low"},
|
||||
),
|
||||
VectorNode(
|
||||
vector_id="threshold_node_8",
|
||||
content="Cooking pasta requires boiling water and adding salt.",
|
||||
metadata={"test_type": "threshold_test", "relevance": "none"},
|
||||
),
|
||||
]
|
||||
|
||||
await store.insert(threshold_nodes)
|
||||
logger.info(f"✓ Inserted {len(threshold_nodes)} test nodes")
|
||||
|
||||
# Test 1: Search without threshold (baseline)
|
||||
results_baseline = await store.search(
|
||||
query="artificial intelligence and machine learning",
|
||||
limit=3,
|
||||
candidates=10,
|
||||
filters={"test_type": "threshold_test"},
|
||||
)
|
||||
logger.info(f"Test 1 - Search without threshold: {len(results_baseline)} results")
|
||||
for i, r in enumerate(results_baseline, 1):
|
||||
score = r.metadata.get("score", 0)
|
||||
logger.info(f" Result {i}: {r.vector_id}, score={score:.4f}")
|
||||
|
||||
assert len(results_baseline) <= 3, "Should return at most limit results"
|
||||
logger.info("✓ Baseline search validated")
|
||||
|
||||
# Test 2: Search with threshold (should filter low-score results)
|
||||
threshold_value = 0.5
|
||||
results_with_threshold = await store.search(
|
||||
query="artificial intelligence and machine learning",
|
||||
limit=3,
|
||||
candidates=10,
|
||||
filters={"test_type": "threshold_test"},
|
||||
threshold=threshold_value,
|
||||
)
|
||||
logger.info(
|
||||
f"Test 2 - Search with threshold={threshold_value}: {len(results_with_threshold)} results",
|
||||
)
|
||||
for i, r in enumerate(results_with_threshold, 1):
|
||||
score = r.metadata.get("score", 0)
|
||||
logger.info(f" Result {i}: {r.vector_id}, score={score:.4f}")
|
||||
# Verify all results meet threshold
|
||||
assert score >= threshold_value, f"Score {score} should be >= threshold {threshold_value}"
|
||||
|
||||
# Should return at most limit results
|
||||
assert len(results_with_threshold) <= 3, "Should return at most limit results"
|
||||
logger.info("✓ Threshold filtering validated")
|
||||
|
||||
# Test 3: Search with high threshold (should filter more strictly)
|
||||
high_threshold = 0.7
|
||||
results_high_threshold = await store.search(
|
||||
query="artificial intelligence and machine learning",
|
||||
limit=10, # Use same limit as candidates to get all results above threshold
|
||||
candidates=10,
|
||||
filters={"test_type": "threshold_test"},
|
||||
threshold=high_threshold,
|
||||
)
|
||||
logger.info(
|
||||
f"Test 3 - Search with high threshold={high_threshold}: {len(results_high_threshold)} results",
|
||||
)
|
||||
for i, r in enumerate(results_high_threshold, 1):
|
||||
score = r.metadata.get("score", 0)
|
||||
logger.info(f" Result {i}: {r.vector_id}, score={score:.4f}")
|
||||
assert score >= high_threshold, f"Score {score} should be >= threshold {high_threshold}"
|
||||
|
||||
# Compare with same limit and lower threshold to verify stricter filtering
|
||||
results_low_threshold = await store.search(
|
||||
query="artificial intelligence and machine learning",
|
||||
limit=10,
|
||||
candidates=10,
|
||||
filters={"test_type": "threshold_test"},
|
||||
threshold=0.3, # Lower threshold
|
||||
)
|
||||
# High threshold should return same or fewer results than low threshold
|
||||
assert len(results_high_threshold) <= len(
|
||||
results_low_threshold,
|
||||
), (
|
||||
f"High threshold ({len(results_high_threshold)}) should return fewer results than low threshold "
|
||||
f"({len(results_low_threshold)})"
|
||||
)
|
||||
logger.info("✓ High threshold filtering validated")
|
||||
|
||||
# Test 4: Search with threshold and large candidates
|
||||
results_large_candidates = await store.search(
|
||||
query="artificial intelligence and machine learning",
|
||||
limit=2,
|
||||
candidates=20,
|
||||
filters={"test_type": "threshold_test"},
|
||||
threshold=0.5,
|
||||
)
|
||||
logger.info(
|
||||
f"Test 4 - Search with threshold and large candidates: {len(results_large_candidates)} results",
|
||||
)
|
||||
# Should search max(candidates, limit) = 20 nodes, filter by threshold, return top 2
|
||||
assert len(results_large_candidates) <= 2, "Should return at most limit results"
|
||||
for r in results_large_candidates:
|
||||
score = r.metadata.get("score", 0)
|
||||
assert score >= 0.5, f"Score {score} should be >= 0.5"
|
||||
logger.info("✓ Large candidates with threshold validated")
|
||||
|
||||
# Test 5: Verify that threshold=None behaves same as no threshold
|
||||
results_none_threshold = await store.search(
|
||||
query="artificial intelligence and machine learning",
|
||||
limit=3,
|
||||
candidates=10,
|
||||
filters={"test_type": "threshold_test"},
|
||||
threshold=None,
|
||||
)
|
||||
logger.info(f"Test 5 - Search with threshold=None: {len(results_none_threshold)} results")
|
||||
# Should behave same as baseline (Test 1)
|
||||
assert len(results_none_threshold) <= 3, "Should return at most limit results"
|
||||
logger.info("✓ threshold=None validated")
|
||||
|
||||
# Clean up test data
|
||||
await store.delete([node.vector_id for node in threshold_nodes])
|
||||
logger.info("Cleaned up test nodes")
|
||||
|
||||
logger.info("✓ Search with threshold test passed")
|
||||
|
||||
|
||||
async def test_get_by_id(store: BaseVectorStore, _store_name: str):
|
||||
"""Test retrieving nodes by vector_id (single and batch)."""
|
||||
logger.info("=" * 20 + " GET BY ID TEST " + "=" * 20)
|
||||
|
|
@ -1671,6 +1845,7 @@ async def run_all_tests_for_store(store_type: str, store_name: str):
|
|||
await test_search_with_single_filter(store, store_name)
|
||||
await test_search_with_exact_match_filter(store, store_name)
|
||||
await test_search_with_multiple_filters(store, store_name)
|
||||
await test_search_with_threshold(store, store_name)
|
||||
await test_get_by_id(store, store_name)
|
||||
await test_list_all(store, store_name)
|
||||
await test_list_with_filters(store, store_name)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue