From fa2205115df24acf983da9b57bc3f623eca0b2aa Mon Sep 17 00:00:00 2001 From: n0x29a <15330763+n0x29a@users.noreply.github.com> Date: Fri, 6 Sep 2024 11:42:00 +0200 Subject: [PATCH] Fix: Normalize tokens to lowercase before checking stopwords in BM25 (#337) * Fix: Normalize tokens to lowercase before checking stopwords in BM25 * Test: Normalize tokens to lowercase before checking stopwords in BM25 * Test Fix test_multilanguage: in "Je suis au lit", the "Je" should be skipped because it in the stopwords. * chore: apply ruff --------- Co-authored-by: H4-8ZSI --- fastembed/sparse/bm25.py | 2 +- tests/test_attention_embeddings.py | 4 +-- tests/test_sparse_embeddings.py | 46 ++++++++++++++++++++++++++---- 3 files changed, 43 insertions(+), 9 deletions(-) diff --git a/fastembed/sparse/bm25.py b/fastembed/sparse/bm25.py index 7c0a2d6..1dfaa39 100644 --- a/fastembed/sparse/bm25.py +++ b/fastembed/sparse/bm25.py @@ -219,7 +219,7 @@ class Bm25(SparseTextEmbeddingBase): if token in self.punctuation: continue - if token in self.stopwords: + if token.lower() in self.stopwords: continue stemmed_token = self.stemmer.stemWord(token) diff --git a/tests/test_attention_embeddings.py b/tests/test_attention_embeddings.py index efd9d1e..31dfcf2 100644 --- a/tests/test_attention_embeddings.py +++ b/tests/test_attention_embeddings.py @@ -92,8 +92,8 @@ def test_multilanguage(model_name): assert embeddings[0].values.shape == (3,) assert embeddings[0].indices.shape == (3,) - assert embeddings[1].values.shape == (2,) - assert embeddings[1].indices.shape == (2,) + assert embeddings[1].values.shape == (1,) + assert embeddings[1].indices.shape == (1,) model = SparseTextEmbedding(model_name=model_name, language="english") embeddings = list(model.embed(docs))[:2] diff --git a/tests/test_sparse_embeddings.py b/tests/test_sparse_embeddings.py index 06a5366..b99f51b 100644 --- a/tests/test_sparse_embeddings.py +++ b/tests/test_sparse_embeddings.py @@ -1,5 +1,6 @@ import pytest +from fastembed.sparse.bm25 import Bm25 from fastembed.sparse.sparse_text_embedding import SparseTextEmbedding CANONICAL_COLUMN_VALUES = { @@ -93,9 +94,42 @@ def test_parallel_processing(): == sparse_embedding_duo.indices.tolist() == sparse_embedding_all.indices.tolist() ) - assert np.allclose( - sparse_embedding.values, sparse_embedding_duo.values, atol=1e-3 - ) - assert np.allclose( - sparse_embedding.values, sparse_embedding_all.values, atol=1e-3 - ) + assert np.allclose(sparse_embedding.values, sparse_embedding_duo.values, atol=1e-3) + assert np.allclose(sparse_embedding.values, sparse_embedding_all.values, atol=1e-3) + + +@pytest.fixture +def bm25_instance(): + return Bm25("Qdrant/bm25", language="english") + + +def test_stem_with_stopwords_and_punctuation(bm25_instance): + # Setup + bm25_instance.stopwords = set(["the", "is", "a"]) + bm25_instance.punctuation = set([".", ",", "!"]) + + # Test data + tokens = ["The", "quick", "brown", "fox", "is", "a", "test", "sentence", ".", "!"] + + # Execute + result = bm25_instance._stem(tokens) + + # Assert + expected = ["quick", "brown", "fox", "test", "sentenc"] + assert result == expected, f"Expected {expected}, but got {result}" + + +def test_stem_case_insensitive_stopwords(bm25_instance): + # Setup + bm25_instance.stopwords = set(["the", "is", "a"]) + bm25_instance.punctuation = set([".", ",", "!"]) + + # Test data + tokens = ["THE", "Quick", "Brown", "Fox", "IS", "A", "Test", "Sentence", ".", "!"] + + # Execute + result = bm25_instance._stem(tokens) + + # Assert + expected = ["Quick", "Brown", "Fox", "Test", "Sentenc"] + assert result == expected, f"Expected {expected}, but got {result}"