feat(memory): add search score explanations (#5102)
Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
@@ -156,6 +156,33 @@ const memories = memory.search("food preferences", {
|
||||
On Mem0 Platform v3, time-aware queries use Temporal Reasoning internally while preserving the normal search response shape. See <Link href="/platform/features/temporal-reasoning">Temporal Reasoning</Link>.
|
||||
</Note>
|
||||
|
||||
### Explain OSS search scores
|
||||
|
||||
OSS search combines semantic similarity with optional keyword and entity signals. Pass `explain=True` when tuning retrieval quality or debugging why a memory ranked where it did:
|
||||
|
||||
<CodeGroup>
|
||||
```python Python
|
||||
results = m.search(
|
||||
"food preferences",
|
||||
filters={"user_id": "alice"},
|
||||
explain=True,
|
||||
)
|
||||
|
||||
print(results["results"][0]["score_details"])
|
||||
```
|
||||
|
||||
```javascript JavaScript
|
||||
const results = await memory.search("food preferences", {
|
||||
filters: { user_id: "alice" },
|
||||
explain: true,
|
||||
});
|
||||
|
||||
console.log(results.results[0].score_details);
|
||||
```
|
||||
</CodeGroup>
|
||||
|
||||
Each result includes `score_details` with the semantic score, normalized BM25 score, entity boost, raw combined score, maximum possible score, final score, and threshold used for filtering. The field is omitted unless `explain` is enabled, so existing response shapes stay unchanged.
|
||||
|
||||
## Filter patterns
|
||||
|
||||
Filters help narrow down search results. Common use cases:
|
||||
|
||||
@@ -226,6 +226,20 @@ curl -X POST http://localhost:8888/search \
|
||||
}'
|
||||
```
|
||||
|
||||
Set `explain` to inspect the scoring signals used by OSS hybrid search:
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:8888/search \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"query": "vegetable pizza",
|
||||
"user_id": "alice",
|
||||
"explain": true
|
||||
}'
|
||||
```
|
||||
|
||||
Each returned memory includes `score_details` only when explanation mode is enabled.
|
||||
|
||||
### Explore with OpenAPI docs
|
||||
|
||||
1. Navigate to `http://localhost:8888/docs` (Compose) or `http://localhost:8000/docs` (raw Docker / uvicorn).
|
||||
|
||||
@@ -1064,7 +1064,7 @@ export class Memory {
|
||||
: {};
|
||||
|
||||
await this._ensureInitialized();
|
||||
const { topK = 20, threshold = 0.1 } = config;
|
||||
const { topK = 20, threshold = 0.1, explain = false } = config;
|
||||
|
||||
await this._captureEvent("search", {
|
||||
query_length: query.length,
|
||||
@@ -1243,6 +1243,7 @@ export class Memory {
|
||||
entityBoosts,
|
||||
threshold ?? 0.1,
|
||||
topK,
|
||||
explain,
|
||||
);
|
||||
|
||||
// Step 9: Format results
|
||||
@@ -1275,6 +1276,7 @@ export class Memory {
|
||||
...(payload.user_id && { user_id: payload.user_id }),
|
||||
...(payload.agent_id && { agent_id: payload.agent_id }),
|
||||
...(payload.run_id && { run_id: payload.run_id }),
|
||||
...(scored.scoreDetails && { score_details: scored.scoreDetails }),
|
||||
};
|
||||
});
|
||||
|
||||
|
||||
@@ -17,6 +17,7 @@ export interface SearchMemoryOptions {
|
||||
topK?: number;
|
||||
filters?: SearchFilters;
|
||||
threshold?: number;
|
||||
explain?: boolean;
|
||||
}
|
||||
|
||||
export interface GetAllMemoryOptions {
|
||||
|
||||
@@ -56,10 +56,21 @@ export function normalizeBm25(
|
||||
return 1.0 / (1.0 + Math.exp(-steepness * (rawScore - midpoint)));
|
||||
}
|
||||
|
||||
export interface ScoreDetails {
|
||||
semanticScore: number;
|
||||
bm25Score: number;
|
||||
entityBoost: number;
|
||||
rawScore: number;
|
||||
maxPossibleScore: number;
|
||||
finalScore: number;
|
||||
threshold: number;
|
||||
}
|
||||
|
||||
export interface ScoredResult {
|
||||
id: string;
|
||||
score: number;
|
||||
payload: Record<string, any>;
|
||||
scoreDetails?: ScoreDetails;
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -82,6 +93,7 @@ export interface ScoredResult {
|
||||
* @param entityBoosts - Map of memory ID to entity boost score.
|
||||
* @param threshold - Minimum semantic score to include a candidate.
|
||||
* @param topK - Maximum number of results to return.
|
||||
* @param explain - Include scoreDetails in each result when true.
|
||||
* @returns Sorted list of scored results, highest score first.
|
||||
*/
|
||||
export function scoreAndRank(
|
||||
@@ -94,6 +106,7 @@ export function scoreAndRank(
|
||||
entityBoosts: Record<string, number>,
|
||||
threshold: number,
|
||||
topK: number,
|
||||
explain: boolean = false,
|
||||
): ScoredResult[] {
|
||||
const hasBm25 = Object.keys(bm25Scores).length > 0;
|
||||
const hasEntity = Object.keys(entityBoosts).length > 0;
|
||||
@@ -126,11 +139,23 @@ export function scoreAndRank(
|
||||
const rawCombined = semanticScore + bm25Score + entityBoost;
|
||||
const combined = Math.min(rawCombined / maxPossible, 1.0);
|
||||
|
||||
scored.push({
|
||||
const entry: ScoredResult = {
|
||||
id: memIdStr,
|
||||
score: combined,
|
||||
payload: result.payload,
|
||||
});
|
||||
};
|
||||
if (explain) {
|
||||
entry.scoreDetails = {
|
||||
semanticScore,
|
||||
bm25Score,
|
||||
entityBoost,
|
||||
rawScore: rawCombined,
|
||||
maxPossibleScore: maxPossible,
|
||||
finalScore: combined,
|
||||
threshold,
|
||||
};
|
||||
}
|
||||
scored.push(entry);
|
||||
}
|
||||
|
||||
scored.sort((a, b) => b.score - a.score);
|
||||
|
||||
@@ -0,0 +1,49 @@
|
||||
/// <reference types="jest" />
|
||||
|
||||
import { scoreAndRank } from "../src/utils/scoring";
|
||||
|
||||
describe("scoreAndRank", () => {
|
||||
const results = [
|
||||
{ id: "a", score: 0.8, payload: { data: "mem a" } },
|
||||
{ id: "b", score: 0.5, payload: { data: "mem b" } },
|
||||
];
|
||||
|
||||
it("omits scoreDetails by default", () => {
|
||||
const scored = scoreAndRank(results, {}, {}, 0.1, 10);
|
||||
expect(scored[0].scoreDetails).toBeUndefined();
|
||||
expect(scored[1].scoreDetails).toBeUndefined();
|
||||
});
|
||||
|
||||
it("omits scoreDetails when explain is false", () => {
|
||||
const scored = scoreAndRank(results, {}, {}, 0.1, 10, false);
|
||||
expect(scored[0].scoreDetails).toBeUndefined();
|
||||
});
|
||||
|
||||
it("includes scoreDetails when explain is true", () => {
|
||||
const bm25 = { a: 0.6 };
|
||||
const entity = { a: 0.3 };
|
||||
const scored = scoreAndRank(results, bm25, entity, 0.1, 10, true);
|
||||
|
||||
const details = scored[0].scoreDetails!;
|
||||
expect(details).toBeDefined();
|
||||
expect(details.semanticScore).toBe(0.8);
|
||||
expect(details.bm25Score).toBe(0.6);
|
||||
expect(details.entityBoost).toBe(0.3);
|
||||
expect(details.rawScore).toBeCloseTo(1.7);
|
||||
expect(details.maxPossibleScore).toBe(2.5);
|
||||
expect(details.finalScore).toBeCloseTo(0.68);
|
||||
expect(details.threshold).toBe(0.1);
|
||||
});
|
||||
|
||||
it("includes scoreDetails for results without bm25/entity signals", () => {
|
||||
const scored = scoreAndRank(results, {}, {}, 0.1, 10, true);
|
||||
|
||||
const details = scored[0].scoreDetails!;
|
||||
expect(details.semanticScore).toBe(0.8);
|
||||
expect(details.bm25Score).toBe(0);
|
||||
expect(details.entityBoost).toBe(0);
|
||||
expect(details.rawScore).toBe(0.8);
|
||||
expect(details.maxPossibleScore).toBe(1.0);
|
||||
expect(details.finalScore).toBe(0.8);
|
||||
});
|
||||
});
|
||||
+16
-4
@@ -1132,6 +1132,7 @@ class Memory(MemoryBase):
|
||||
filters: Optional[Dict[str, Any]] = None,
|
||||
threshold: float = 0.1,
|
||||
rerank: bool = False,
|
||||
explain: bool = False,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
@@ -1162,6 +1163,7 @@ class Memory(MemoryBase):
|
||||
- {"NOT": [filter1]} - logical NOT
|
||||
threshold (float, optional): Minimum score for a memory to be included. Defaults to 0.1.
|
||||
rerank (bool, optional): Whether to rerank results. Defaults to False.
|
||||
explain (bool, optional): Whether to include score_details for each result. Defaults to False.
|
||||
|
||||
Returns:
|
||||
dict: A dictionary containing the search results under a "results" key.
|
||||
@@ -1221,11 +1223,12 @@ class Memory(MemoryBase):
|
||||
"encoded_ids": encoded_ids,
|
||||
"sync_type": "sync",
|
||||
"threshold": threshold,
|
||||
"explain": explain,
|
||||
"advanced_filters": bool(filters and self._has_advanced_operators(filters)),
|
||||
},
|
||||
)
|
||||
|
||||
original_memories = self._search_vector_store(query, effective_filters, limit, threshold)
|
||||
original_memories = self._search_vector_store(query, effective_filters, limit, threshold, explain=explain)
|
||||
|
||||
# Apply reranking if enabled and reranker is available
|
||||
if rerank and self.reranker and original_memories:
|
||||
@@ -1341,7 +1344,7 @@ class Memory(MemoryBase):
|
||||
return True
|
||||
return False
|
||||
|
||||
def _search_vector_store(self, query, filters, limit, threshold=0.1):
|
||||
def _search_vector_store(self, query, filters, limit, threshold=0.1, explain=False):
|
||||
# Guard against None threshold (backward compat)
|
||||
if threshold is None:
|
||||
threshold = 0.1
|
||||
@@ -1396,6 +1399,7 @@ class Memory(MemoryBase):
|
||||
entity_boosts=entity_boosts,
|
||||
threshold=threshold,
|
||||
top_k=limit,
|
||||
explain=explain,
|
||||
)
|
||||
|
||||
# Step 9: Format results
|
||||
@@ -1433,6 +1437,8 @@ class Memory(MemoryBase):
|
||||
if not memory_item_dict.get("metadata"):
|
||||
memory_item_dict["metadata"] = {}
|
||||
memory_item_dict["metadata"].update(additional_metadata)
|
||||
if explain and "score_details" in scored:
|
||||
memory_item_dict["score_details"] = scored["score_details"]
|
||||
|
||||
original_memories.append(memory_item_dict)
|
||||
|
||||
@@ -2568,6 +2574,7 @@ class AsyncMemory(MemoryBase):
|
||||
filters: Optional[Dict[str, Any]] = None,
|
||||
threshold: float = 0.1,
|
||||
rerank: bool = False,
|
||||
explain: bool = False,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
@@ -2598,6 +2605,7 @@ class AsyncMemory(MemoryBase):
|
||||
- {"NOT": [filter1]} - logical NOT
|
||||
threshold (float, optional): Minimum score for a memory to be included. Defaults to 0.1.
|
||||
rerank (bool, optional): Whether to rerank results. Defaults to False.
|
||||
explain (bool, optional): Whether to include score_details for each result. Defaults to False.
|
||||
|
||||
Returns:
|
||||
dict: A dictionary containing the search results under a "results" key.
|
||||
@@ -2659,11 +2667,12 @@ class AsyncMemory(MemoryBase):
|
||||
"encoded_ids": encoded_ids,
|
||||
"sync_type": "async",
|
||||
"threshold": threshold,
|
||||
"explain": explain,
|
||||
"advanced_filters": bool(filters and self._has_advanced_operators(filters)),
|
||||
},
|
||||
)
|
||||
|
||||
original_memories = await self._search_vector_store(query, effective_filters, limit, threshold)
|
||||
original_memories = await self._search_vector_store(query, effective_filters, limit, threshold, explain=explain)
|
||||
|
||||
# Apply reranking if enabled and reranker is available
|
||||
if rerank and self.reranker and original_memories:
|
||||
@@ -2782,7 +2791,7 @@ class AsyncMemory(MemoryBase):
|
||||
return True
|
||||
return False
|
||||
|
||||
async def _search_vector_store(self, query, filters, limit, threshold=0.1):
|
||||
async def _search_vector_store(self, query, filters, limit, threshold=0.1, explain=False):
|
||||
if threshold is None:
|
||||
threshold = 0.1
|
||||
|
||||
@@ -2836,6 +2845,7 @@ class AsyncMemory(MemoryBase):
|
||||
entity_boosts=entity_boosts,
|
||||
threshold=threshold,
|
||||
top_k=limit,
|
||||
explain=explain,
|
||||
)
|
||||
|
||||
# Step 9: Format results
|
||||
@@ -2872,6 +2882,8 @@ class AsyncMemory(MemoryBase):
|
||||
if not memory_item_dict.get("metadata"):
|
||||
memory_item_dict["metadata"] = {}
|
||||
memory_item_dict["metadata"].update(additional_metadata)
|
||||
if explain and "score_details" in scored:
|
||||
memory_item_dict["score_details"] = scored["score_details"]
|
||||
|
||||
original_memories.append(memory_item_dict)
|
||||
|
||||
|
||||
+21
-3
@@ -63,6 +63,7 @@ def score_and_rank(
|
||||
entity_boosts: Dict[str, float],
|
||||
threshold: float,
|
||||
top_k: int,
|
||||
explain: bool = False,
|
||||
) -> List[Dict[str, Any]]:
|
||||
"""Score candidates additively and return top-k results.
|
||||
|
||||
@@ -79,6 +80,14 @@ def score_and_rank(
|
||||
- Semantic + BM25 + entity: max_possible = 2.5
|
||||
- Semantic + entity (no BM25): max_possible = 1.5
|
||||
|
||||
Args:
|
||||
semantic_results: Candidate memories from vector search.
|
||||
bm25_scores: Normalized keyword scores keyed by memory ID.
|
||||
entity_boosts: Entity-link boosts keyed by memory ID.
|
||||
threshold: Minimum semantic score required before hybrid scoring.
|
||||
top_k: Maximum number of results to return.
|
||||
explain: Include score_details in each result when true.
|
||||
|
||||
Returns:
|
||||
List of scored result dicts sorted by combined score descending.
|
||||
"""
|
||||
@@ -109,13 +118,22 @@ def score_and_rank(
|
||||
raw_combined = semantic_score + bm25_score + entity_boost
|
||||
combined = min(raw_combined / max_possible, 1.0)
|
||||
|
||||
scored.append(
|
||||
{
|
||||
scored_result = {
|
||||
"id": mem_id_str,
|
||||
"score": combined,
|
||||
"payload": result.get("payload"),
|
||||
}
|
||||
)
|
||||
if explain:
|
||||
scored_result["score_details"] = {
|
||||
"semantic_score": semantic_score,
|
||||
"bm25_score": bm25_score,
|
||||
"entity_boost": entity_boost,
|
||||
"raw_score": raw_combined,
|
||||
"max_possible_score": max_possible,
|
||||
"final_score": combined,
|
||||
"threshold": threshold,
|
||||
}
|
||||
scored.append(scored_result)
|
||||
|
||||
scored.sort(key=lambda x: x["score"], reverse=True)
|
||||
return scored[:top_k]
|
||||
|
||||
@@ -199,6 +199,7 @@ class SearchRequest(BaseModel):
|
||||
filters: Optional[Dict[str, Any]] = None
|
||||
top_k: Optional[int] = Field(None, description="Maximum number of results to return.")
|
||||
threshold: Optional[float] = Field(None, description="Minimum similarity score for results.")
|
||||
explain: Optional[bool] = Field(None, description="Include score details for each search result.")
|
||||
|
||||
|
||||
class GenerateInstructionsRequest(BaseModel):
|
||||
|
||||
+36
-1
@@ -143,6 +143,42 @@ def test_search_handles_incomplete_payloads(mock_sqlite, mock_llm_factory, mock_
|
||||
assert result[0]["memory"] == "content"
|
||||
|
||||
|
||||
@patch('mem0.memory.main.extract_entities', return_value=[])
|
||||
@patch('mem0.utils.factory.EmbedderFactory.create')
|
||||
@patch('mem0.utils.factory.VectorStoreFactory.create')
|
||||
@patch('mem0.utils.factory.LlmFactory.create')
|
||||
@patch('mem0.memory.storage.SQLiteManager')
|
||||
def test_search_explain_includes_score_details(
|
||||
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory, _mock_extract_entities
|
||||
):
|
||||
mock_embedder = MagicMock()
|
||||
mock_embedder.embed.return_value = [0.1, 0.2, 0.3]
|
||||
mock_embedder_factory.return_value = mock_embedder
|
||||
|
||||
mock_vector_store = MagicMock()
|
||||
mock_vector_store.search.return_value = [
|
||||
MockVectorMemory("mem_1", {"data": "content", "user_id": "test"}, score=0.8)
|
||||
]
|
||||
mock_vector_store.keyword_search.return_value = [
|
||||
MockVectorMemory("mem_1", {"data": "content", "user_id": "test"}, score=5.0)
|
||||
]
|
||||
mock_vector_factory.return_value = mock_vector_store
|
||||
mock_llm_factory.return_value = MagicMock()
|
||||
mock_sqlite.return_value = MagicMock()
|
||||
|
||||
from mem0.memory.main import Memory as MemoryClass
|
||||
memory = MemoryClass(MemoryConfig())
|
||||
|
||||
result = memory.search("test query", filters={"user_id": "test"}, explain=True)
|
||||
|
||||
details = result["results"][0]["score_details"]
|
||||
assert details["semantic_score"] == 0.8
|
||||
assert details["bm25_score"] > 0
|
||||
assert details["entity_boost"] == 0.0
|
||||
assert details["final_score"] == result["results"][0]["score"]
|
||||
assert details["threshold"] == 0.1
|
||||
|
||||
|
||||
@patch('mem0.utils.factory.EmbedderFactory.create')
|
||||
@patch('mem0.utils.factory.VectorStoreFactory.create')
|
||||
@patch('mem0.utils.factory.LlmFactory.create')
|
||||
@@ -934,4 +970,3 @@ async def test_async_create_memory_stores_text_lemmatized(mock_sqlite, mock_llm_
|
||||
"AsyncMemory._create_memory must store text_lemmatized for BM25 keyword search"
|
||||
)
|
||||
assert payload[0]["text_lemmatized"] != "", "text_lemmatized must not be empty"
|
||||
|
||||
|
||||
@@ -126,6 +126,28 @@ class TestScoreAndRank:
|
||||
scored = score_and_rank(results, bm25, entity, threshold=0.1, top_k=10)
|
||||
assert scored[0]["score"] <= 1.0
|
||||
|
||||
def test_explain_includes_score_details(self):
|
||||
results = [{"id": "a", "score": 0.8, "payload": {"data": "mem a"}}]
|
||||
bm25 = {"a": 0.6}
|
||||
entity = {"a": 0.3}
|
||||
scored = score_and_rank(results, bm25, entity, threshold=0.1, top_k=10, explain=True)
|
||||
|
||||
details = scored[0]["score_details"]
|
||||
assert details == {
|
||||
"semantic_score": 0.8,
|
||||
"bm25_score": 0.6,
|
||||
"entity_boost": 0.3,
|
||||
"raw_score": pytest.approx(1.7),
|
||||
"max_possible_score": 2.5,
|
||||
"final_score": pytest.approx(0.68),
|
||||
"threshold": 0.1,
|
||||
}
|
||||
|
||||
def test_score_details_are_omitted_by_default(self):
|
||||
results = [{"id": "a", "score": 0.8, "payload": {"data": "mem a"}}]
|
||||
scored = score_and_rank(results, {}, {}, threshold=0.1, top_k=10)
|
||||
assert "score_details" not in scored[0]
|
||||
|
||||
|
||||
class TestEntityBoostWeight:
|
||||
def test_weight_value(self):
|
||||
|
||||
Reference in New Issue
Block a user