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:
soumil-rathi
2026-04-14 05:30:58 -07:00
committed by GitHub
parent 57f944e18a
commit a488e19044
120 changed files with 10107 additions and 17135 deletions
View File
+101
View File
@@ -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])
+67
View File
@@ -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
+142
View File
@@ -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