Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
31 changes: 30 additions & 1 deletion fastembed/common/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
import tempfile
import unicodedata
from pathlib import Path
from functools import lru_cache
from itertools import islice
from typing import Iterable, TypeVar

Expand Down Expand Up @@ -93,5 +94,33 @@ def get_all_punctuation() -> set[str]:
)


@lru_cache(maxsize=None)
def get_all_marks() -> str:
"""Return the combining marks (Unicode category M) as regex character class ranges.

Marks include Tamil and Devanagari vowel signs and Arabic harakat. The regex word class
does not match them, so a pattern that only keeps word characters splits words in these
scripts at every mark. Ranges keep the class short, which keeps matching fast.
"""
ranges: list[str] = []
start = None
for i in range(sys.maxunicode + 2):
is_mark = i <= sys.maxunicode and unicodedata.category(chr(i)).startswith("M")
if is_mark and start is None:
start = i
elif not is_mark and start is not None:
ranges.append(f"{re.escape(chr(start))}-{re.escape(chr(i - 1))}")
start = None
return "".join(ranges)


@lru_cache(maxsize=None)
def _non_alphanumeric_pattern() -> re.Pattern[str]:
return re.compile(rf"[^\w\s{get_all_marks()}]")


def remove_non_alphanumeric(text: str) -> str:
return re.sub(r"[^\w\s]", " ", text, flags=re.UNICODE)
# ASCII text has no combining marks, and the plain class is faster to match.
if text.isascii():
return re.sub(r"[^\w\s]", " ", text)
return _non_alphanumeric_pattern().sub(" ", text)
15 changes: 14 additions & 1 deletion fastembed/sparse/utils/tokenizer.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,25 @@
# This code is a modified copy of the `NLTKWordTokenizer` class from `NLTK` library.

import re
from functools import lru_cache

from fastembed.common.utils import get_all_marks


@lru_cache(maxsize=None)
def _non_word_pattern() -> re.Pattern[str]:
return re.compile(rf"[^\w{get_all_marks()}]")


class SimpleTokenizer:
@staticmethod
def tokenize(text: str) -> list[str]:
text = re.sub(r"[^\w]", " ", text.lower())
text = text.lower()
# ASCII text has no combining marks, and the plain class is faster to match.
if text.isascii():
text = re.sub(r"[^\w]", " ", text)
else:
text = _non_word_pattern().sub(" ", text)
text = re.sub(r"\s+", " ", text)

return text.strip().split()
Expand Down
31 changes: 31 additions & 0 deletions tests/test_sparse_embeddings.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,10 @@
import pytest
import numpy as np

from fastembed.common.utils import remove_non_alphanumeric
from fastembed.sparse.bm25 import Bm25
from fastembed.sparse.sparse_text_embedding import SparseTextEmbedding
from fastembed.sparse.utils.tokenizer import SimpleTokenizer
from tests.utils import delete_model_cache, is_manual_run, should_test_model


Expand Down Expand Up @@ -311,6 +313,35 @@ def test_disable_stemmer_behavior(disable_stemmer: bool) -> None:
assert result == expected, f"Expected {expected}, but got {result}"


def test_combining_marks_do_not_split_words() -> None:
# Tamil and Devanagari vowel signs and Arabic harakat are combining marks, which the regex
# word class does not match.
text = "தமிழ் மொழி, हिन्दी भाषा! ذَهَبَ الطَّالِبُ"
assert SimpleTokenizer.tokenize(text) == [
"தமிழ்",
"மொழி",
"हिन्दी",
"भाषा",
"ذَهَبَ",
"الطَّالِبُ",
]
assert remove_non_alphanumeric("தமிழ், हिन्दी!") == "தமிழ் हिन्दी "


@pytest.mark.parametrize(
"language,text,expected",
[
("arabic", "ذَهَبَ الطَّالِبُ إِلَى المَدْرَسَةِ", ["ذهب", "طالب", "الي", "مدرس"]),
("tamil", "சென்னை தமிழ்நாட்டின் தலைநகரம் ஆகும்", ["சென்", "தமிழ்நாடு", "தலைநகரம்", "ஆக்"]),
],
ids=["arabic", "tamil"],
)
def test_stem_words_with_combining_marks(language: str, text: str, expected: list[str]) -> None:
model = Bm25("Qdrant/bm25", language=language)
tokens = model.tokenizer.tokenize(remove_non_alphanumeric(text))
assert model._stem(tokens) == expected


class _PolishStemmer:
def stem_word(self, word: str) -> str:
for suffix in ("ami", "ach"):
Expand Down