feat(oss): port v3 pipeline with hybrid search, entity extraction, and additive scoring (#4805)
Co-authored-by: Soumil Rathi <soumilrathi@gmail.com> Co-authored-by: Saket Aryan <saketaryan2002@gmail.com> Co-authored-by: chaithanyak42 <chaithanya.kumar42a@gmail.com> Co-authored-by: kartik-mem0 <kartik.labhshetwar@mem0.ai>
This commit is contained in:
@@ -0,0 +1,101 @@
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _ensure_spacy():
|
||||
"""Skip tests if spaCy model is not available."""
|
||||
try:
|
||||
import spacy
|
||||
spacy.load("en_core_web_sm")
|
||||
except Exception:
|
||||
pytest.skip("spaCy en_core_web_sm model not available")
|
||||
|
||||
|
||||
class TestExtractEntities:
|
||||
def test_proper_nouns(self):
|
||||
from mem0.utils.entity_extraction import extract_entities
|
||||
|
||||
entities = extract_entities("John Smith works at Google on machine learning projects")
|
||||
entity_texts = [e[1] for e in entities]
|
||||
# Should extract proper nouns
|
||||
found_proper = any("John" in t or "Google" in t for t in entity_texts)
|
||||
assert found_proper, f"Expected proper nouns, got {entities}"
|
||||
|
||||
def test_quoted_text(self):
|
||||
from mem0.utils.entity_extraction import extract_entities
|
||||
|
||||
entities = extract_entities('She is reading "The Great Gatsby" this week')
|
||||
entity_texts = [e[1] for e in entities]
|
||||
assert any("Great Gatsby" in t for t in entity_texts), f"Expected quoted text, got {entities}"
|
||||
|
||||
def test_compound_nouns(self):
|
||||
from mem0.utils.entity_extraction import extract_entities
|
||||
|
||||
entities = extract_entities("The machine learning engineer built a neural network")
|
||||
entity_texts = [e[1].lower() for e in entities]
|
||||
has_compound = any("machine" in t and "learning" in t for t in entity_texts) or \
|
||||
any("neural" in t and "network" in t for t in entity_texts)
|
||||
assert has_compound, f"Expected compound nouns, got {entities}"
|
||||
|
||||
def test_empty_string(self):
|
||||
from mem0.utils.entity_extraction import extract_entities
|
||||
|
||||
entities = extract_entities("")
|
||||
assert entities == []
|
||||
|
||||
def test_no_entities(self):
|
||||
from mem0.utils.entity_extraction import extract_entities
|
||||
|
||||
entities = extract_entities("I like things and stuff")
|
||||
# Generic words should be filtered out
|
||||
entity_texts = [e[1].lower() for e in entities]
|
||||
assert "things" not in entity_texts
|
||||
assert "stuff" not in entity_texts
|
||||
|
||||
def test_deduplication(self):
|
||||
from mem0.utils.entity_extraction import extract_entities
|
||||
|
||||
entities = extract_entities("Google is great. I love working at Google.")
|
||||
google_count = sum(1 for _, t in entities if "Google" in t)
|
||||
assert google_count <= 1, f"Expected dedup, got {entities}"
|
||||
|
||||
def test_returns_tuples(self):
|
||||
from mem0.utils.entity_extraction import extract_entities
|
||||
|
||||
entities = extract_entities("John Smith lives in New York City")
|
||||
for entity in entities:
|
||||
assert isinstance(entity, tuple)
|
||||
assert len(entity) == 2
|
||||
assert entity[0] in ("PROPER", "QUOTED", "COMPOUND", "NOUN")
|
||||
assert isinstance(entity[1], str)
|
||||
|
||||
|
||||
class TestExtractEntitiesBatch:
|
||||
def test_batch_processing(self):
|
||||
from mem0.utils.entity_extraction import extract_entities_batch
|
||||
|
||||
texts = [
|
||||
"John works at Google",
|
||||
"Mary lives in Paris",
|
||||
"The cat sat on the mat",
|
||||
]
|
||||
results = extract_entities_batch(texts)
|
||||
assert len(results) == 3
|
||||
assert isinstance(results[0], list)
|
||||
assert isinstance(results[1], list)
|
||||
assert isinstance(results[2], list)
|
||||
|
||||
def test_empty_input(self):
|
||||
from mem0.utils.entity_extraction import extract_entities_batch
|
||||
|
||||
assert extract_entities_batch([]) == []
|
||||
|
||||
def test_consistency_with_single(self):
|
||||
from mem0.utils.entity_extraction import extract_entities, extract_entities_batch
|
||||
|
||||
text = "John Smith works at Google headquarters"
|
||||
single = extract_entities(text)
|
||||
batch = extract_entities_batch([text])
|
||||
assert len(batch) == 1
|
||||
# Both should extract the same entities
|
||||
assert set(t for _, t in single) == set(t for _, t in batch[0])
|
||||
@@ -0,0 +1,67 @@
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _ensure_spacy():
|
||||
"""Skip tests if spaCy model is not available."""
|
||||
try:
|
||||
import spacy
|
||||
spacy.load("en_core_web_sm")
|
||||
except Exception:
|
||||
pytest.skip("spaCy en_core_web_sm model not available")
|
||||
|
||||
|
||||
class TestLemmatizeForBm25:
|
||||
def test_basic_lemmatization(self):
|
||||
from mem0.utils.lemmatization import lemmatize_for_bm25
|
||||
|
||||
result = lemmatize_for_bm25("The cats are running quickly")
|
||||
assert "cat" in result
|
||||
assert "run" in result or "running" in result
|
||||
# Stop words and punctuation should be removed
|
||||
assert "the" not in result.split()
|
||||
|
||||
def test_verb_forms_normalized(self):
|
||||
from mem0.utils.lemmatization import lemmatize_for_bm25
|
||||
|
||||
result = lemmatize_for_bm25("she attended multiple meetings yesterday")
|
||||
assert "attend" in result or "attended" in result
|
||||
assert "meeting" in result # -ing form preserved alongside lemma
|
||||
# "multiple" is kept (not a spaCy stop word)
|
||||
|
||||
def test_ing_preservation(self):
|
||||
from mem0.utils.lemmatization import lemmatize_for_bm25
|
||||
|
||||
result = lemmatize_for_bm25("attending the morning meeting")
|
||||
tokens = result.split()
|
||||
# Should have both the lemma and the -ing form
|
||||
assert "attending" in tokens or "attend" in tokens
|
||||
|
||||
def test_empty_string(self):
|
||||
from mem0.utils.lemmatization import lemmatize_for_bm25
|
||||
|
||||
result = lemmatize_for_bm25("")
|
||||
assert result == ""
|
||||
|
||||
def test_punctuation_removed(self):
|
||||
from mem0.utils.lemmatization import lemmatize_for_bm25
|
||||
|
||||
result = lemmatize_for_bm25("Hello, world! How are you?")
|
||||
assert "," not in result
|
||||
assert "!" not in result
|
||||
assert "?" not in result
|
||||
|
||||
def test_lowercased(self):
|
||||
from mem0.utils.lemmatization import lemmatize_for_bm25
|
||||
|
||||
result = lemmatize_for_bm25("PYTHON Programming LANGUAGE")
|
||||
for token in result.split():
|
||||
assert token == token.lower()
|
||||
|
||||
def test_stop_words_removed(self):
|
||||
from mem0.utils.lemmatization import lemmatize_for_bm25
|
||||
|
||||
result = lemmatize_for_bm25("this is a very simple test of the system")
|
||||
tokens = result.split()
|
||||
for stop in ["this", "is", "a", "very", "of", "the"]:
|
||||
assert stop not in tokens
|
||||
@@ -0,0 +1,142 @@
|
||||
import pytest
|
||||
|
||||
from mem0.utils.scoring import (
|
||||
get_bm25_params,
|
||||
normalize_bm25,
|
||||
score_and_rank,
|
||||
ENTITY_BOOST_WEIGHT,
|
||||
)
|
||||
|
||||
|
||||
class TestGetBm25Params:
|
||||
def test_short_query(self):
|
||||
midpoint, steepness = get_bm25_params("hello world", lemmatized="hello world")
|
||||
assert midpoint == 5.0
|
||||
assert steepness == 0.7
|
||||
|
||||
def test_medium_query(self):
|
||||
midpoint, steepness = get_bm25_params("x", lemmatized="one two three four five")
|
||||
assert midpoint == 7.0
|
||||
assert steepness == 0.6
|
||||
|
||||
def test_long_query(self):
|
||||
words = " ".join(f"word{i}" for i in range(20))
|
||||
midpoint, steepness = get_bm25_params("x", lemmatized=words)
|
||||
assert midpoint == 12.0
|
||||
assert steepness == 0.5
|
||||
|
||||
def test_empty_lemmatized(self):
|
||||
midpoint, steepness = get_bm25_params("test", lemmatized="")
|
||||
# Empty string -> 1 term -> short query params
|
||||
assert midpoint == 5.0
|
||||
|
||||
|
||||
class TestNormalizeBm25:
|
||||
def test_at_midpoint(self):
|
||||
score = normalize_bm25(5.0, 5.0, 0.7)
|
||||
assert abs(score - 0.5) < 0.01 # Should be ~0.5 at midpoint
|
||||
|
||||
def test_high_score(self):
|
||||
score = normalize_bm25(20.0, 5.0, 0.7)
|
||||
assert score > 0.99 # Well above midpoint
|
||||
|
||||
def test_low_score(self):
|
||||
score = normalize_bm25(0.0, 5.0, 0.7)
|
||||
assert score < 0.05 # Well below midpoint
|
||||
|
||||
def test_range(self):
|
||||
for raw in [0, 1, 5, 10, 20, 50]:
|
||||
score = normalize_bm25(float(raw), 5.0, 0.7)
|
||||
assert 0.0 <= score <= 1.0
|
||||
|
||||
|
||||
class TestScoreAndRank:
|
||||
def test_semantic_only(self):
|
||||
results = [
|
||||
{"id": "a", "score": 0.9, "payload": {"data": "mem a"}},
|
||||
{"id": "b", "score": 0.5, "payload": {"data": "mem b"}},
|
||||
]
|
||||
scored = score_and_rank(results, {}, {}, threshold=0.1, top_k=10)
|
||||
assert len(scored) == 2
|
||||
# With no BM25/entity, max_possible=1.0, so scores stay the same
|
||||
assert scored[0]["score"] == pytest.approx(0.9)
|
||||
assert scored[1]["score"] == pytest.approx(0.5)
|
||||
|
||||
def test_semantic_plus_bm25(self):
|
||||
results = [
|
||||
{"id": "a", "score": 0.8, "payload": {"data": "mem a"}},
|
||||
{"id": "b", "score": 0.6, "payload": {"data": "mem b"}},
|
||||
]
|
||||
bm25 = {"a": 0.3, "b": 0.9}
|
||||
scored = score_and_rank(results, bm25, {}, threshold=0.1, top_k=10)
|
||||
# max_possible = 2.0 (semantic + bm25)
|
||||
# a: (0.8 + 0.3) / 2.0 = 0.55
|
||||
# b: (0.6 + 0.9) / 2.0 = 0.75
|
||||
assert scored[0]["id"] == "b" # b should rank higher due to BM25
|
||||
assert scored[0]["score"] == pytest.approx(0.75)
|
||||
assert scored[1]["id"] == "a"
|
||||
assert scored[1]["score"] == pytest.approx(0.55)
|
||||
|
||||
def test_all_three_signals(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)
|
||||
# max_possible = 2.5
|
||||
expected = (0.8 + 0.6 + 0.3) / 2.5
|
||||
assert scored[0]["score"] == pytest.approx(expected)
|
||||
|
||||
def test_threshold_gates_on_semantic(self):
|
||||
results = [
|
||||
{"id": "a", "score": 0.05, "payload": {"data": "mem a"}}, # Below threshold
|
||||
{"id": "b", "score": 0.5, "payload": {"data": "mem b"}},
|
||||
]
|
||||
bm25 = {"a": 0.99} # High BM25 shouldn't save it
|
||||
scored = score_and_rank(results, bm25, {}, threshold=0.1, top_k=10)
|
||||
assert len(scored) == 1
|
||||
assert scored[0]["id"] == "b"
|
||||
|
||||
def test_top_k_limit(self):
|
||||
results = [{"id": str(i), "score": 0.5, "payload": {}} for i in range(20)]
|
||||
scored = score_and_rank(results, {}, {}, threshold=0.1, top_k=5)
|
||||
assert len(scored) == 5
|
||||
|
||||
def test_score_breakdown_present(self):
|
||||
results = [{"id": "a", "score": 0.8, "payload": {"data": "x"}}]
|
||||
bm25 = {"a": 0.4}
|
||||
entity = {"a": 0.2}
|
||||
scored = score_and_rank(results, bm25, entity, threshold=0.1, top_k=10)
|
||||
breakdown = scored[0]["score_breakdown"]
|
||||
assert breakdown["semantic"] == 0.8
|
||||
assert breakdown["bm25"] == 0.4
|
||||
assert breakdown["entity_boost"] == 0.2
|
||||
|
||||
def test_adaptive_divisor_semantic_only(self):
|
||||
results = [{"id": "a", "score": 0.8, "payload": {}}]
|
||||
scored = score_and_rank(results, {}, {}, threshold=0.1, top_k=10)
|
||||
# max_possible = 1.0 (no bm25, no entity)
|
||||
assert scored[0]["score"] == pytest.approx(0.8)
|
||||
|
||||
def test_adaptive_divisor_semantic_plus_entity(self):
|
||||
results = [{"id": "a", "score": 0.8, "payload": {}}]
|
||||
entity = {"a": 0.3}
|
||||
scored = score_and_rank(results, {}, entity, threshold=0.1, top_k=10)
|
||||
# max_possible = 1.5 (semantic + entity)
|
||||
expected = (0.8 + 0.3) / 1.5
|
||||
assert scored[0]["score"] == pytest.approx(expected)
|
||||
|
||||
def test_empty_results(self):
|
||||
scored = score_and_rank([], {}, {}, threshold=0.1, top_k=10)
|
||||
assert scored == []
|
||||
|
||||
def test_score_clamped_to_1(self):
|
||||
results = [{"id": "a", "score": 1.0, "payload": {}}]
|
||||
bm25 = {"a": 1.0}
|
||||
entity = {"a": 0.5}
|
||||
scored = score_and_rank(results, bm25, entity, threshold=0.1, top_k=10)
|
||||
assert scored[0]["score"] <= 1.0
|
||||
|
||||
|
||||
class TestEntityBoostWeight:
|
||||
def test_weight_value(self):
|
||||
assert ENTITY_BOOST_WEIGHT == 0.5
|
||||
Reference in New Issue
Block a user