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}"