feat: Add configurable embedding similarity threshold for graph store node matching (#3593)
This commit is contained in:
@@ -64,6 +64,7 @@
|
||||
"platform/features/async-client",
|
||||
"platform/features/async-mode-default-change",
|
||||
"platform/features/graph-memory",
|
||||
"platform/features/graph-threshold",
|
||||
"platform/features/advanced-retrieval",
|
||||
"platform/features/criteria-retrieval",
|
||||
"platform/features/multimodal-support",
|
||||
|
||||
@@ -0,0 +1,209 @@
|
||||
---
|
||||
title: Configurable Graph Threshold
|
||||
---
|
||||
|
||||
## Overview
|
||||
|
||||
The graph store threshold parameter controls how strictly nodes are matched during graph data ingestion based on embedding similarity. This feature allows you to customize the matching behavior to prevent false matches or enable entity merging based on your specific use case.
|
||||
|
||||
## Configuration
|
||||
|
||||
Add the `threshold` parameter to your graph store configuration:
|
||||
|
||||
```python
|
||||
from mem0 import Memory
|
||||
|
||||
config = {
|
||||
"graph_store": {
|
||||
"provider": "neo4j", # or memgraph, neptune, kuzu
|
||||
"config": {
|
||||
"url": "bolt://localhost:7687",
|
||||
"username": "neo4j",
|
||||
"password": "password"
|
||||
},
|
||||
"threshold": 0.7 # Default value, range: 0.0 to 1.0
|
||||
}
|
||||
}
|
||||
|
||||
memory = Memory.from_config(config)
|
||||
```
|
||||
|
||||
## Parameters
|
||||
|
||||
| Parameter | Type | Default | Range | Description |
|
||||
|-----------|------|---------|-------|-------------|
|
||||
| `threshold` | float | 0.7 | 0.0 - 1.0 | Minimum embedding similarity score required to match existing nodes during graph ingestion |
|
||||
|
||||
## Use Cases
|
||||
|
||||
### Strict Matching (UUIDs, IDs)
|
||||
|
||||
Use higher thresholds (0.95-0.99) when working with identifiers that should remain distinct:
|
||||
|
||||
```python
|
||||
config = {
|
||||
"graph_store": {
|
||||
"provider": "neo4j",
|
||||
"config": {...},
|
||||
"threshold": 0.95 # Strict matching
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Example:** Prevents UUID collisions like `MXxBUE18QVBQTElDQVRJT058MjM3MTM4NjI5` being matched with `MXxBUE18QVBQTElDQVRJT058MjA2OTYxMzM`
|
||||
|
||||
### Permissive Matching (Natural Language)
|
||||
|
||||
Use lower thresholds (0.6-0.7) when entity variations should be merged:
|
||||
|
||||
```python
|
||||
config = {
|
||||
"graph_store": {
|
||||
"threshold": 0.6 # Permissive matching
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Example:** Merges similar entities like "Bob" and "Robert" as the same person.
|
||||
|
||||
## Threshold Guidelines
|
||||
|
||||
| Use Case | Recommended Threshold | Behavior |
|
||||
|----------|----------------------|----------|
|
||||
| UUIDs, IDs, Keys | 0.95 - 0.99 | Prevent false matches between similar identifiers |
|
||||
| Structured Data | 0.85 - 0.9 | Balanced precision and recall |
|
||||
| General Purpose | 0.7 - 0.8 | Default recommendation |
|
||||
| Natural Language | 0.6 - 0.7 | Allow entity variations to merge |
|
||||
|
||||
## Examples
|
||||
|
||||
### Example 1: Preventing Data Loss with UUIDs
|
||||
|
||||
```python
|
||||
from mem0 import Memory
|
||||
|
||||
config = {
|
||||
"graph_store": {
|
||||
"provider": "neo4j",
|
||||
"config": {
|
||||
"url": "bolt://localhost:7687",
|
||||
"username": "neo4j",
|
||||
"password": "password"
|
||||
},
|
||||
"threshold": 0.98 # Very strict for UUIDs
|
||||
}
|
||||
}
|
||||
|
||||
memory = Memory.from_config(config)
|
||||
|
||||
# These UUIDs create separate nodes instead of being incorrectly merged
|
||||
memory.add(
|
||||
[{"role": "user", "content": "MXxBUE18QVBQTElDQVRJT058MjM3MTM4NjI5 relates to Project A"}],
|
||||
user_id="user1"
|
||||
)
|
||||
|
||||
memory.add(
|
||||
[{"role": "user", "content": "MXxBUE18QVBQTElDQVRJT058MjA2OTYxMzM relates to Project B"}],
|
||||
user_id="user1"
|
||||
)
|
||||
```
|
||||
|
||||
### Example 2: Merging Entity Variations
|
||||
|
||||
```python
|
||||
config = {
|
||||
"graph_store": {
|
||||
"provider": "neo4j",
|
||||
"config": {...},
|
||||
"threshold": 0.6 # More permissive
|
||||
}
|
||||
}
|
||||
|
||||
memory = Memory.from_config(config)
|
||||
|
||||
# These will be merged as the same entity
|
||||
memory.add([{"role": "user", "content": "Bob works at Google"}], user_id="user1")
|
||||
memory.add([{"role": "user", "content": "Robert works at Google"}], user_id="user1")
|
||||
```
|
||||
|
||||
### Example 3: Different Thresholds for Different Clients
|
||||
|
||||
```python
|
||||
# Client 1: Strict matching for transactional data
|
||||
memory_strict = Memory.from_config({
|
||||
"graph_store": {"threshold": 0.95}
|
||||
})
|
||||
|
||||
# Client 2: Permissive matching for conversational data
|
||||
memory_permissive = Memory.from_config({
|
||||
"graph_store": {"threshold": 0.6}
|
||||
})
|
||||
```
|
||||
|
||||
## Supported Graph Providers
|
||||
|
||||
The threshold parameter works with all graph store providers:
|
||||
|
||||
- ✅ Neo4j
|
||||
- ✅ Memgraph
|
||||
- ✅ Kuzu
|
||||
- ✅ Neptune (both Analytics and DB)
|
||||
|
||||
## How It Works
|
||||
|
||||
When adding a relation to the graph:
|
||||
|
||||
1. **Embedding Generation**: The system generates embeddings for source and destination entities
|
||||
2. **Node Search**: Searches for existing nodes with similar embeddings
|
||||
3. **Threshold Comparison**: Compares similarity scores against the configured threshold
|
||||
4. **Decision**:
|
||||
- If similarity ≥ threshold: Uses the existing node
|
||||
- If similarity < threshold: Creates a new node
|
||||
|
||||
```python
|
||||
# Pseudocode
|
||||
if node_similarity >= threshold:
|
||||
use_existing_node()
|
||||
else:
|
||||
create_new_node()
|
||||
```
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
### Issue: Duplicate nodes being created
|
||||
|
||||
**Symptom**: Expected nodes to merge but they're created separately
|
||||
|
||||
**Solution**: Lower the threshold
|
||||
```python
|
||||
config = {"graph_store": {"threshold": 0.6}}
|
||||
```
|
||||
|
||||
### Issue: Unrelated entities being merged
|
||||
|
||||
**Symptom**: Different entities incorrectly matched as the same node
|
||||
|
||||
**Solution**: Raise the threshold
|
||||
```python
|
||||
config = {"graph_store": {"threshold": 0.95}}
|
||||
```
|
||||
|
||||
### Issue: Validation error
|
||||
|
||||
**Symptom**: `ValidationError: threshold must be between 0.0 and 1.0`
|
||||
|
||||
**Solution**: Ensure threshold is in valid range
|
||||
```python
|
||||
config = {"graph_store": {"threshold": 0.7}} # Valid: 0.0 ≤ x ≤ 1.0
|
||||
```
|
||||
|
||||
## Backward Compatibility
|
||||
|
||||
- **Default Value**: 0.7 (maintains existing behavior)
|
||||
- **Optional Parameter**: Existing code works without any changes
|
||||
- **No Breaking Changes**: Graceful fallback if not specified
|
||||
|
||||
## Related
|
||||
|
||||
- [Graph Store Configuration](/features/custom-categories)
|
||||
- [Issue #3590](https://github.com/mem0ai/mem0/issues/3590)
|
||||
@@ -89,6 +89,15 @@ class GraphStoreConfig(BaseModel):
|
||||
custom_prompt: Optional[str] = Field(
|
||||
description="Custom prompt to fetch entities from the given text", default=None
|
||||
)
|
||||
threshold: float = Field(
|
||||
description="Threshold for embedding similarity when matching nodes during graph ingestion. "
|
||||
"Range: 0.0 to 1.0. Higher values require closer matches. "
|
||||
"Use lower values (e.g., 0.5-0.7) for distinct entities with similar embeddings. "
|
||||
"Use higher values (e.g., 0.9+) when you want stricter matching.",
|
||||
default=0.7,
|
||||
ge=0.0,
|
||||
le=1.0,
|
||||
)
|
||||
|
||||
@field_validator("config")
|
||||
def validate_config(cls, v, values):
|
||||
|
||||
@@ -234,8 +234,8 @@ class NeptuneBase(ABC):
|
||||
dest_embedding = self.embedding_model.embed(destination)
|
||||
|
||||
# search for the nodes with the closest embeddings
|
||||
source_node_search_result = self._search_source_node(source_embedding, user_id, threshold=0.9)
|
||||
destination_node_search_result = self._search_destination_node(dest_embedding, user_id, threshold=0.9)
|
||||
source_node_search_result = self._search_source_node(source_embedding, user_id, threshold=self.threshold)
|
||||
destination_node_search_result = self._search_destination_node(dest_embedding, user_id, threshold=self.threshold)
|
||||
|
||||
cypher, params = self._add_entities_cypher(
|
||||
source_node_search_result,
|
||||
|
||||
@@ -56,7 +56,8 @@ class MemoryGraph(NeptuneBase):
|
||||
|
||||
self.llm = NeptuneBase._create_llm(self.config, self.llm_provider)
|
||||
self.user_id = None
|
||||
self.threshold = 0.7
|
||||
# Use threshold from graph_store config, default to 0.7 for backward compatibility
|
||||
self.threshold = self.config.graph_store.threshold if hasattr(self.config.graph_store, 'threshold') else 0.7
|
||||
self.vector_store_limit=5
|
||||
|
||||
def _delete_entities_cypher(self, source, destination, relationship, user_id):
|
||||
|
||||
@@ -39,7 +39,8 @@ class MemoryGraph(NeptuneBase):
|
||||
|
||||
self.llm = NeptuneBase._create_llm(self.config, self.llm_provider)
|
||||
self.user_id = None
|
||||
self.threshold = 0.7
|
||||
# Use threshold from graph_store config, default to 0.7 for backward compatibility
|
||||
self.threshold = self.config.graph_store.threshold if hasattr(self.config.graph_store, 'threshold') else 0.7
|
||||
|
||||
def _delete_entities_cypher(self, source, destination, relationship, user_id):
|
||||
"""
|
||||
|
||||
@@ -70,7 +70,8 @@ class MemoryGraph:
|
||||
llm_config = self.config.llm.config
|
||||
self.llm = LlmFactory.create(self.llm_provider, llm_config)
|
||||
self.user_id = None
|
||||
self.threshold = 0.7
|
||||
# Use threshold from graph_store config, default to 0.7 for backward compatibility
|
||||
self.threshold = self.config.graph_store.threshold if hasattr(self.config.graph_store, 'threshold') else 0.7
|
||||
|
||||
def add(self, data, filters):
|
||||
"""
|
||||
@@ -434,8 +435,8 @@ class MemoryGraph:
|
||||
dest_embedding = self.embedding_model.embed(destination)
|
||||
|
||||
# search for the nodes with the closest embeddings
|
||||
source_node_search_result = self._search_source_node(source_embedding, filters, threshold=0.9)
|
||||
destination_node_search_result = self._search_destination_node(dest_embedding, filters, threshold=0.9)
|
||||
source_node_search_result = self._search_source_node(source_embedding, filters, threshold=self.threshold)
|
||||
destination_node_search_result = self._search_destination_node(dest_embedding, filters, threshold=self.threshold)
|
||||
|
||||
# TODO: Create a cypher query and common params for all the cases
|
||||
if not destination_node_search_result and source_node_search_result:
|
||||
|
||||
@@ -62,7 +62,8 @@ class MemoryGraph:
|
||||
self.llm = LlmFactory.create(self.llm_provider, llm_config)
|
||||
|
||||
self.user_id = None
|
||||
self.threshold = 0.7
|
||||
# Use threshold from graph_store config, default to 0.7 for backward compatibility
|
||||
self.threshold = self.config.graph_store.threshold if hasattr(self.config.graph_store, 'threshold') else 0.7
|
||||
|
||||
def kuzu_create_schema(self):
|
||||
self.kuzu_execute(
|
||||
@@ -452,8 +453,8 @@ class MemoryGraph:
|
||||
dest_embedding = self.embedding_model.embed(destination)
|
||||
|
||||
# search for the nodes with the closest embeddings
|
||||
source_node_search_result = self._search_source_node(source_embedding, filters, threshold=0.9)
|
||||
destination_node_search_result = self._search_destination_node(dest_embedding, filters, threshold=0.9)
|
||||
source_node_search_result = self._search_source_node(source_embedding, filters, threshold=self.threshold)
|
||||
destination_node_search_result = self._search_destination_node(dest_embedding, filters, threshold=self.threshold)
|
||||
|
||||
if not destination_node_search_result and source_node_search_result:
|
||||
params = {
|
||||
|
||||
@@ -55,7 +55,8 @@ class MemoryGraph:
|
||||
llm_config = self.config.llm.config
|
||||
self.llm = LlmFactory.create(self.llm_provider, llm_config)
|
||||
self.user_id = None
|
||||
self.threshold = 0.7
|
||||
# Use threshold from graph_store config, default to 0.7 for backward compatibility
|
||||
self.threshold = self.config.graph_store.threshold if hasattr(self.config.graph_store, 'threshold') else 0.7
|
||||
|
||||
# Setup Memgraph:
|
||||
# 1. Create vector index (created Entity label on all nodes)
|
||||
@@ -422,8 +423,8 @@ class MemoryGraph:
|
||||
dest_embedding = self.embedding_model.embed(destination)
|
||||
|
||||
# search for the nodes with the closest embeddings
|
||||
source_node_search_result = self._search_source_node(source_embedding, filters, threshold=0.9)
|
||||
destination_node_search_result = self._search_destination_node(dest_embedding, filters, threshold=0.9)
|
||||
source_node_search_result = self._search_source_node(source_embedding, filters, threshold=self.threshold)
|
||||
destination_node_search_result = self._search_destination_node(dest_embedding, filters, threshold=self.threshold)
|
||||
|
||||
# Prepare agent_id for node creation
|
||||
agent_id_clause = ""
|
||||
|
||||
Reference in New Issue
Block a user