diff --git a/src/memos/reranker/cosine_local.py b/src/memos/reranker/cosine_local.py index 140074b1d..0b09f2f7a 100644 --- a/src/memos/reranker/cosine_local.py +++ b/src/memos/reranker/cosine_local.py @@ -25,6 +25,9 @@ def _cosine_one_to_many(q: list[float], m: list[list[float]]) -> list[float]: """ Compute cosine similarities between a single vector q and a matrix m (rows are candidates). + Returns a list of floats in [-1, 1], with numerical-stability safeguards: + zero vectors produce similarity 0 instead of NaN/Inf, and any remaining edge cases + (NaN, Inf) are clamped to 0. """ if not _HAS_NUMPY: @@ -38,15 +41,27 @@ def norm(a): # lowercase per N806 sims = [] for v in m: vn = norm(v) or 1e-10 - sims.append(dot(q, v) / (qn * vn)) + score = dot(q, v) / (qn * vn) + if score != score or score in (float("inf"), float("-inf")): + score = 0.0 + sims.append(score) return sims qv = _np.asarray(q, dtype=float) # lowercase + if qv.ndim > 1: + qv = qv.reshape(-1) mv = _np.asarray(m, dtype=float) # lowercase - qn = _np.linalg.norm(qv) or 1e-10 + if mv.ndim != 2 or mv.shape[1] == 0: + return [0.0] * len(m) + qn = _np.linalg.norm(qv) mn = _np.linalg.norm(mv, axis=1) # lowercase + if qn < 1e-10: + return [0.0] * len(m) + denom = mn * qn + 1e-10 dots = mv @ qv - return (dots / (mn * qn + 1e-10)).tolist() + scores = dots / denom + scores = _np.nan_to_num(scores, nan=0.0, posinf=0.0, neginf=0.0) + return scores.tolist() class CosineLocalReranker(BaseReranker): diff --git a/tests/reranker/test_cosine_local.py b/tests/reranker/test_cosine_local.py new file mode 100644 index 000000000..f4c6a157a --- /dev/null +++ b/tests/reranker/test_cosine_local.py @@ -0,0 +1,131 @@ +"""Tests for CosineLocalReranker numerical stability.""" + +import math + +from types import SimpleNamespace + +from memos.reranker.cosine_local import CosineLocalReranker, _cosine_one_to_many + + +class MemoryStub: + def __init__(self, memory_id, embedding, background="fact"): + self.id = memory_id + self.memory = f"memory-{memory_id}" + self.metadata = SimpleNamespace(embedding=embedding, background=background) + + +def _make_items(embeddings, backgrounds=None): + backgrounds = backgrounds or ["fact"] * len(embeddings) + return [ + MemoryStub(i, emb, bg) + for i, (emb, bg) in enumerate(zip(embeddings, backgrounds, strict=False)) + ] + + +class TestCosineOneToManyNumerical: + def test_basic_identical_vector_returns_one(self): + q = [1.0, 0.0, 0.0] + m = [[1.0, 0.0, 0.0], [0.0, 1.0, 0.0]] + result = _cosine_one_to_many(q, m) + assert len(result) == 2 + assert abs(result[0] - 1.0) < 1e-6 + assert abs(result[1]) < 1e-6 + + def test_orthogonal_returns_zero(self): + q = [1.0, 0.0] + m = [[0.0, 1.0]] + result = _cosine_one_to_many(q, m) + assert abs(result[0]) < 1e-6 + + def test_opposite_direction(self): + q = [1.0, 0.0] + m = [[-1.0, 0.0]] + result = _cosine_one_to_many(q, m) + assert abs(result[0] - (-1.0)) < 1e-6 + + def test_empty_matrix_returns_empty(self): + q = [1.0, 2.0, 3.0] + m = [] + result = _cosine_one_to_many(q, m) + assert result == [] + + def test_zero_query_vector(self): + q = [0.0, 0.0, 0.0] + m = [[1.0, 0.0, 0.0], [0.0, 1.0, 0.0]] + result = _cosine_one_to_many(q, m) + assert len(result) == 2 + for r in result: + assert r == 0.0 + + def test_zero_candidate_vectors(self): + q = [1.0, 0.0, 0.0] + m = [[0.0, 0.0, 0.0], [1.0, 0.0, 0.0]] + result = _cosine_one_to_many(q, m) + assert len(result) == 2 + assert result[0] == 0.0 + assert abs(result[1] - 1.0) < 1e-6 + + def test_2d_query_vector_flattened(self): + q = [[1.0, 0.0, 0.0]] + m = [[1.0, 0.0, 0.0]] + result = _cosine_one_to_many(q, m) + assert len(result) == 1 + assert abs(result[0] - 1.0) < 1e-6 + + def test_no_nan_in_output_normalized_vectors(self): + q = [1.0] + m = [[0.0]] + result = _cosine_one_to_many(q, m) + assert not any(math.isnan(r) for r in result) + assert result[0] == 0.0 + + def test_multiple_equal_length_vectors(self): + q = [1.0, 2.0] + m = [[3.0, 4.0], [1.0, 2.0], [0.5, 1.0]] + result = _cosine_one_to_many(q, m) + assert len(result) == 3 + assert all(-1.0 - 1e-6 <= r <= 1.0 + 1e-6 for r in result) + + +class TestCosineLocalReranker: + def test_empty_graph_results(self): + reranker = CosineLocalReranker() + assert reranker.rerank("query", [], top_k=5) == [] + + def test_rerank_preserves_top_items(self): + embeddings = [ + [1.0, 0.0, 0.0], + [0.9, 0.1, 0.0], + [0.0, 1.0, 0.0], + [0.0, 0.9, 0.1], + ] + items = _make_items(embeddings) + query_emb = [1.0, 0.0, 0.0] + reranker = CosineLocalReranker() + result = reranker.rerank("any", items, top_k=2, query_embedding=query_emb) + ids = [it.id for it, _ in result] + assert len(result) == 2 + assert ids[0] == 0 or ids[0] == 1 + assert ids[1] == 0 or ids[1] == 1 + + def test_rerank_without_embeddings_falls_back(self): + items = [MemoryStub(i, None) for i in range(3)] + query_emb = [1.0, 0.0, 0.0] + reranker = CosineLocalReranker() + result = reranker.rerank("any", items, top_k=5, query_embedding=query_emb) + assert len(result) == 3 + + def test_level_weights_applied(self): + embeddings = [[1.0, 0.0], [1.0, 0.0]] + backgrounds = ["topic", "fact"] + items = _make_items(embeddings, backgrounds) + query_emb = [1.0, 0.0] + reranker = CosineLocalReranker( + level_weights={"topic": 2.0, "concept": 1.0, "fact": 0.5}, + level_field="background", + ) + result = reranker.rerank("any", items, top_k=2, query_embedding=query_emb) + ids = [it.id for it, _ in result] + scores = {it.id: sc for it, sc in result} + assert ids[0] == 0 + assert scores[0] > scores[1]