diff --git a/dataflow/operators/general_text/filter/minhash_deduplicate_filter.py b/dataflow/operators/general_text/filter/minhash_deduplicate_filter.py index db73df92..7925a47d 100644 --- a/dataflow/operators/general_text/filter/minhash_deduplicate_filter.py +++ b/dataflow/operators/general_text/filter/minhash_deduplicate_filter.py @@ -45,8 +45,12 @@ def get_desc(lang: str = "zh"): def create_minhash(self, data): minhash = MinHash(num_perm=self.num_perm) if self.use_n_gram: - for i in range(len(data) - self.n_gram + 1): - minhash.update(data[i:i + self.n_gram].encode('utf8')) + if 0 < len(data) < self.n_gram: + # A non-empty text shorter than n still needs a distinguishing shingle. + minhash.update(data.encode('utf8')) + else: + for i in range(len(data) - self.n_gram + 1): + minhash.update(data[i:i + self.n_gram].encode('utf8')) else: for d in data: minhash.update(d.encode('utf8')) diff --git a/test/cpu_only/test_minhash_short_texts.py b/test/cpu_only/test_minhash_short_texts.py new file mode 100644 index 00000000..f693ece5 --- /dev/null +++ b/test/cpu_only/test_minhash_short_texts.py @@ -0,0 +1,89 @@ +import pandas as pd +import pytest + +from dataflow.operators.general_text import MinHashDeduplicateFilter +from dataflow.utils.storage import DataFlowStorage + + +class MemoryStorage(DataFlowStorage): + def __init__(self, texts): + self.dataframe = pd.DataFrame({"text": texts}) + self.result = None + + def get_keys_from_dataframe(self): + return self.dataframe.columns.tolist() + + def read(self, output_type="dataframe"): + assert output_type == "dataframe" + return self.dataframe.copy() + + def write(self, dataframe): + self.result = dataframe.copy() + return "memory://minhash-short-texts" + + +@pytest.mark.cpu +@pytest.mark.parametrize( + "texts, ngram", + [ + (["cat", "dog", "cat"], 5), + (["北京", "上海", "北京"], 5), + (["apples", "oranges", "apples"], 8), + ], +) +def test_preserves_distinct_texts_shorter_than_ngram(texts, ngram): + storage = MemoryStorage(texts) + + result_keys = MinHashDeduplicateFilter(ngram=ngram).run( + storage, input_key="text" + ) + + assert storage.result["text"].tolist() == texts[:2] + assert storage.result.index.tolist() == [0, 1] + assert result_keys == ["minhash_deduplicated_label"] + assert storage.result[result_keys[0]].tolist() == [1, 1] + + +@pytest.mark.cpu +def test_empty_texts_are_deduplicated_separately_from_short_texts(): + storage = MemoryStorage(["", "cat", "", "dog", "cat", ""]) + + MinHashDeduplicateFilter().run(storage, input_key="text") + + assert storage.result["text"].tolist() == ["", "cat", "dog"] + assert storage.result.index.tolist() == [0, 1, 3] + + +@pytest.mark.cpu +@pytest.mark.parametrize( + "texts", + [ + ["aaaaa", "zzzzz", "aaaaa"], + ["aaaaaaaaaa", "zzzzzzzzzz", "aaaaaaaaaa"], + ], +) +def test_texts_at_or_above_ngram_keep_existing_deduplication(texts): + storage = MemoryStorage(texts) + + MinHashDeduplicateFilter().run(storage, input_key="text") + + assert storage.result["text"].tolist() == texts[:2] + + +@pytest.mark.cpu +def test_character_mode_keeps_existing_deduplication(): + storage = MemoryStorage(["cat", "dog", "cat"]) + + MinHashDeduplicateFilter(use_n_gram=False).run(storage, input_key="text") + + assert storage.result["text"].tolist() == ["cat", "dog"] + + +@pytest.mark.cpu +def test_reusing_operator_does_not_share_deduplication_state(): + operator = MinHashDeduplicateFilter() + + for _ in range(2): + storage = MemoryStorage(["cat", "dog", "cat"]) + operator.run(storage, input_key="text") + assert storage.result["text"].tolist() == ["cat", "dog"]