diff --git a/pyproject.toml b/pyproject.toml index 926f5a0..a5030c8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -23,7 +23,7 @@ classifiers = [ dependencies = [ "model2vec>=0.3.4", - "vicinity[usearch]>=0.4.3", + "vicinity[usearch]>=0.4.6", "frozendict", "pyversity", ] @@ -106,3 +106,4 @@ version = {attr = "semhash.version.__version__"} [tool.uv] exclude-newer = "1 week" +exclude-newer-package = { vicinity = false } diff --git a/semhash/datamodels.py b/semhash/datamodels.py index 9f98c4b..26480cf 100644 --- a/semhash/datamodels.py +++ b/semhash/datamodels.py @@ -1,37 +1,47 @@ from __future__ import annotations -import json +import warnings from collections import defaultdict from collections.abc import Hashable, Sequence from dataclasses import dataclass, field from functools import cached_property from typing import Generic +import numpy as np from frozendict import frozendict -from semhash.utils import DuplicateList, Record, to_frozendict +from semhash.utils import DuplicateList, Neighbors, Record, select_canonicals, to_frozendict @dataclass class DuplicateRecord(Generic[Record]): """ - A single record with its duplicates. + A record that was filtered as a duplicate of another record. Attributes ---------- record: The original record being deduplicated. exact: Whether the record was identified as an exact match. - duplicates: List of tuples consisting of duplicate records and their associated scores. + duplicate_of: The record that this record is a duplicate of. + score: The similarity score between record and duplicate_of. + duplicates: Deprecated, use duplicate_of and score instead. """ record: Record exact: bool - duplicates: DuplicateList = field(default_factory=list) + duplicate_of: Record + score: float - def _rethreshold(self, threshold: float) -> None: - """Rethreshold the duplicates.""" - self.duplicates = [(d, score) for d, score in self.duplicates if score >= threshold] + @property + def duplicates(self) -> DuplicateList: + """Deprecated, use duplicate_of and score instead.""" + warnings.warn( + "'duplicates' is deprecated and will be removed in a future release. Use 'duplicate_of' and 'score' instead.", + DeprecationWarning, + stacklevel=2, + ) + return [(self.duplicate_of, self.score)] @dataclass @@ -69,6 +79,36 @@ class DeduplicationResult(Generic[Record]): threshold: float = field(default=0.9) columns: Sequence[str] | None = field(default=None) + # Inputs of self_deduplicate, so rethreshold can rerun the selection at a higher threshold. + _selection_inputs: tuple[list[list[Record]], list[Neighbors], np.ndarray] | None = field( + default=None, init=False, repr=False, compare=False + ) + + @classmethod + def _from_groups( + cls, + groups: list[list[Record]], + neighbors: list[Neighbors], + vectors: np.ndarray, + threshold: float, + columns: Sequence[str] | None, + ) -> DeduplicationResult[Record]: + """Build a self-deduplication result where every filtered record points to a directly matching selected record.""" + result = cls(threshold=threshold, columns=columns) + canonicals = select_canonicals(vectors=vectors, neighbors=neighbors, threshold=threshold) + for i, (group, (canonical, score)) in enumerate(zip(groups, canonicals)): + is_selected = canonical == i + if is_selected: + result.selected.append(group[0]) + # The rest of a selected group are exact copies of its first record. + filtered_records = group[1:] if is_selected else group + result.filtered.extend( + DuplicateRecord(record=record, exact=is_selected, duplicate_of=groups[canonical][0], score=score) + for record in filtered_records + ) + result._selection_inputs = (groups, neighbors, vectors) + return result + @property def duplicate_ratio(self) -> float: """Return the percentage of records dropped.""" @@ -90,7 +130,7 @@ def get_least_similar_from_duplicates(self, n: int = 1) -> list[tuple[Record, Re :param n: The number of least similar pairs to return. :return: A list of tuples consisting of (original_record, duplicate_record, score). """ - all_pairs = [(dup.record, d, score) for dup in self.filtered for d, score in dup.duplicates] + all_pairs = [(dup.record, dup.duplicate_of, dup.score) for dup in self.filtered] sorted_pairs = sorted(all_pairs, key=lambda x: x[2]) # Sort by score return sorted_pairs[:n] @@ -100,12 +140,21 @@ def rethreshold(self, threshold: float) -> None: raise ValueError("Threshold is smaller than the given value.") # Invalidate cached property before modifying data self.__dict__.pop("selected_with_duplicates", None) - # Rethreshold duplicates and move records without duplicates to selected - for dup in list(self.filtered): - dup._rethreshold(threshold) - if not dup.duplicates: - self.filtered.remove(dup) - self.selected.append(dup.record) + if (inputs := self._selection_inputs) is not None: + # Rerun the selection, since a record that is no longer filtered can become the canonical of later records. + groups, neighbors, vectors = inputs + result = self._from_groups( + groups=groups, neighbors=neighbors, vectors=vectors, threshold=threshold, columns=self.columns + ) + self.selected, self.filtered = result.selected, result.filtered + else: + filtered = [] + for dup in self.filtered: + if dup.score >= threshold: + filtered.append(dup) + else: + self.selected.append(dup.record) + self.filtered = filtered self.threshold = threshold @cached_property @@ -127,24 +176,14 @@ def _to_hashable(record: Record) -> frozendict[str, str] | str: # Build a mapping from original-record to [(duplicate, score), …] buckets: defaultdict[Hashable, DuplicateList] = defaultdict(list) for duplicate_record in self.filtered: - for original_record, score in duplicate_record.duplicates: - buckets[_to_hashable(original_record)].append((duplicate_record.record, float(score))) + buckets[_to_hashable(duplicate_record.duplicate_of)].append( + (duplicate_record.record, duplicate_record.score) + ) result: list[SelectedWithDuplicates[Record]] = [] for selected in self.selected: - # Get the list of duplicates for the selected record - raw_list = buckets.get(_to_hashable(selected), []) - # Ensure we don't have duplicates in the list - # Use full-record canonical JSON for dicts so that unhashable values are handled correctly - deduped = { - ( - json.dumps(rec, sort_keys=True, separators=(",", ":"), ensure_ascii=False) - if isinstance(rec, dict) - else rec - ): (rec, score) - for rec, score in raw_list - } - result.append(SelectedWithDuplicates(record=selected, duplicates=list(deduped.values()))) + # Preserve occurrences, even when multiple input records have identical values. + result.append(SelectedWithDuplicates(record=selected, duplicates=buckets.get(_to_hashable(selected), []))) return result diff --git a/semhash/index.py b/semhash/index.py index 207a860..331ac6d 100644 --- a/semhash/index.py +++ b/semhash/index.py @@ -7,8 +7,8 @@ from vicinity.backends import AbstractBackend, get_backend_class from vicinity.datatypes import SingleQueryResult -DocScore = tuple[dict[str, str], float] -DocScores = list[DocScore] +from semhash.utils import MAX_NEIGHBORS, NEIGHBORS_PER_QUERY, Neighbors + DictItem = list[dict[str, str]] @@ -47,25 +47,26 @@ def from_vectors_and_items( return cls(vectors, items, backend) - def query_threshold(self, vectors: np.ndarray, threshold: float) -> list[DocScores]: + def query_threshold(self, vectors: np.ndarray, threshold: float) -> list[Neighbors]: """ Query the index with a threshold. :param vectors: The vectors to query. :param threshold: The similarity threshold. - :return: The query results. + :return: Arrays of group indices and cosine similarity scores for each query. """ - out: list[DocScores] = [] - for result in self.backend.threshold(vectors, threshold=1 - threshold, max_k=100): - intermediate = [] - for index, distance in zip(*result): - # Every item in the index contains one or more records that are exact duplicates of each other. - # Only the first is returned, since listing every copy grows with the size of the group. - # The score is the cosine similarity. The backend returns distances, so we need to convert. - intermediate.append((self.items[index][0], 1 - distance)) - out.append(intermediate) - - return out + neighbors = [ + (indices, 1 - distances) + for indices, distances in self.backend.threshold( + vectors, threshold=1 - threshold, max_k=NEIGHBORS_PER_QUERY + ) + ] + # Query rows that hit the limit again with a larger one, so dense clusters rarely need a brute-force fallback. + if capped := [i for i, (indices, _) in enumerate(neighbors) if len(indices) >= NEIGHBORS_PER_QUERY]: + results = self.backend.threshold(vectors[capped], threshold=1 - threshold, max_k=MAX_NEIGHBORS) + for i, (indices, distances) in zip(capped, results): + neighbors[i] = (indices, 1 - distances) + return neighbors def query_top_k(self, vectors: np.ndarray, k: int, vectors_are_in_index: bool) -> list[SingleQueryResult]: """ diff --git a/semhash/records.py b/semhash/records.py index ab66f89..b209fd0 100644 --- a/semhash/records.py +++ b/semhash/records.py @@ -153,17 +153,12 @@ def map_deduplication_result_to_strings(result: DeduplicationResult, columns: Se mapped = [] for dup_rec in result.filtered: record_as_str = dict_to_string(dup_rec.record, columns) - duplicates_as_str = [(dict_to_string(r, columns), score) for r, score in dup_rec.duplicates] mapped.append( DuplicateRecord( record=record_as_str, - duplicates=duplicates_as_str, + duplicate_of=dict_to_string(dup_rec.duplicate_of, columns), + score=dup_rec.score, exact=dup_rec.exact, ) ) return DeduplicationResult(selected=deduplicated_str, filtered=mapped, threshold=result.threshold, columns=columns) - - -def add_scores_to_records(records: list[dict[str, str]]) -> list[tuple[dict[str, str], float]]: - """Add scores to records and return a DeduplicationResult.""" - return [(record, 1.0) for record in records] diff --git a/semhash/semhash.py b/semhash/semhash.py index 27e883d..a34088a 100644 --- a/semhash/semhash.py +++ b/semhash/semhash.py @@ -13,7 +13,7 @@ from semhash.datamodels import DeduplicationResult, DuplicateRecord, FilterResult from semhash.index import Index from semhash.records import ( - add_scores_to_records, + dict_to_string, group_records_by_key, map_deduplication_result_to_strings, prepare_records, @@ -25,6 +25,7 @@ coerce_value, compute_candidate_limit, featurize, + normalize, to_frozendict, ) @@ -173,26 +174,27 @@ def deduplicate( ) duplicate_records = [] for record, duplicates in exact_duplicates: - duplicated_with_score = add_scores_to_records(duplicates) - duplicate_record = DuplicateRecord(record=record, duplicates=duplicated_with_score, exact=True) + duplicate_record = DuplicateRecord(record=record, duplicate_of=duplicates[0], score=1.0, exact=True) duplicate_records.append(duplicate_record) # Only embed and query the records that are left after removing exact duplicates - results = [] + deduplicated_records = [] if dict_records: embeddings = featurize(records=dict_records, columns=self.columns, model=self.model) results = self.index.query_threshold(embeddings, threshold=threshold) - - deduplicated_records = [] - for record, similar_items in zip(dict_records, results): - if not similar_items: - # No duplicates found, keep this record - deduplicated_records.append(record) - else: + for record, embedding, (indices, _) in zip(dict_records, embeddings, results): + # Rescore the neighbors with exact cosine similarity, like self_deduplicate does. + scores = normalize(self.index.vectors[indices]) @ normalize(embedding) + if not len(indices) or scores.max() < threshold: + # No duplicates found, keep this record + deduplicated_records.append(record) + continue + best = int(np.argmax(scores)) duplicate_records.append( DuplicateRecord( record=record, - duplicates=[(item, score) for item, score in similar_items], + duplicate_of=self.index.items[indices[best]][0], + score=float(scores[best]), exact=False, ) ) @@ -217,54 +219,14 @@ def self_deduplicate( :param threshold: Similarity threshold for deduplication. :return: A deduplicated list of records. """ - # Query the fitted index - results = self.index.query_threshold(self.index.vectors, threshold=threshold) - column_set = set(self.columns) - - duplicate_records = [] - - deduplicated_records = [] - seen_items: set[frozendict[str, str]] = set() - for item, similar_items in zip(self.index.items, results): - # Items is a list of items which are exact duplicates of each other. - # The first record is kept, and every other copy is an exact duplicate of it. Each copy only lists - # the kept record, since listing every other copy grows quadratically with the size of the group. - record, *duplicates = item - for curr_record in duplicates: - duplicate_records.append(DuplicateRecord(record=curr_record, duplicates=[(record, 1.0)], exact=True)) - - # If we don't see any similar_items, we know the record is not a duplicate. - # In rare cases, the item itself might not be returned by the index. - if not similar_items: # pragma: no cover - deduplicated_records.append(record) - continue - items, _ = zip(*similar_items) - frozen_items = [to_frozendict(item, column_set) for item in items] - # similar_items includes 'record' itself - # If we've seen any of these items before, this is a duplicate cluster. - if any(item in seen_items for item in frozen_items): - duplicate_records.append( - DuplicateRecord( - record=record, - duplicates=[(item, score) for item, score in similar_items if item != record], - exact=False, - ) - ) - continue - # This is the first time we see this cluster of similar items - deduplicated_records.append(record) - # Mark all items in this cluster as seen - seen_items.update(frozen_items) - - result = DeduplicationResult( - selected=deduplicated_records, filtered=duplicate_records, threshold=threshold, columns=self.columns - ) - + neighbors = self.index.query_threshold(self.index.vectors, threshold=threshold) + groups: list[list[Any]] = self.index.items if self._was_string: - # Convert records back to strings if the records were originally strings - return map_deduplication_result_to_strings(result, columns=self.columns) - - return result + # Convert before selection, so the result holds strings when rethreshold reruns it. + groups = [[dict_to_string(record, self.columns) for record in group] for group in groups] + return DeduplicationResult._from_groups( + groups=groups, neighbors=neighbors, vectors=self.index.vectors, threshold=threshold, columns=self.columns + ) def _validate_if_strings(self, records: Sequence[dict[str, Any] | str]) -> list[dict[str, Any]]: """ diff --git a/semhash/utils.py b/semhash/utils.py index c2f54e0..72ddc77 100644 --- a/semhash/utils.py +++ b/semhash/utils.py @@ -8,6 +8,10 @@ # Type definitions Record = TypeVar("Record", str, dict[str, Any]) DuplicateList: TypeAlias = list[tuple[Record, float]] +Neighbors: TypeAlias = tuple[np.ndarray, np.ndarray] + +NEIGHBORS_PER_QUERY = 100 +MAX_NEIGHBORS = 1000 class Encoder(Protocol): @@ -147,3 +151,46 @@ def featurize( embeddings_per_col.append(np.asarray(col_emb)) return np.concatenate(embeddings_per_col, axis=1) + + +def normalize(vectors: np.ndarray) -> np.ndarray: + """Scale vectors to unit length, leaving zero vectors (e.g. empty text) at zero so they are similar to nothing.""" + norms = np.linalg.norm(vectors, axis=-1, keepdims=True) + return vectors / np.where(norms == 0, 1.0, norms) + + +def select_canonicals(vectors: np.ndarray, neighbors: list[Neighbors], threshold: float) -> list[tuple[int, float]]: + """ + Greedily select groups in input order, assigning every other group to a directly matching selected group. + + :param vectors: The vector of each group. + :param neighbors: The approximate neighbors of each group, as arrays of group indices and similarity scores. + :param threshold: The similarity threshold. + :return: The canonical group index and exact similarity score for each group; selected groups point to themselves. + """ + unit_vectors = normalize(vectors) + canonical_indices = np.empty(len(neighbors), dtype=np.int64) + selected_indices: list[int] = [] + canonicals: list[tuple[int, float]] = [] + + def _closest(i: int, candidates: np.ndarray) -> tuple[int, float] | None: + """Return the candidate most similar to group i, if it reaches the threshold.""" + if not len(candidates): + return None + scores = unit_vectors[candidates] @ unit_vectors[i] + best = int(np.argmax(scores)) + return (int(candidates[best]), float(scores[best])) if scores[best] >= threshold else None + + for i, (indices, scores) in enumerate(neighbors): + matches = indices[scores >= threshold] + # Check the canonicals of earlier matching neighbors directly, to avoid transitive matches. + best_match = _closest(i, np.unique(canonical_indices[matches[matches < i]])) + if best_match is None and len(matches) >= MAX_NEIGHBORS: + # The neighbors may be truncated, so compare against every selected group directly. + best_match = _closest(i, np.asarray(selected_indices)) + if best_match is None: + selected_indices.append(i) + best_match = (i, 1.0) + canonical_indices[i] = best_match[0] + canonicals.append(best_match) + return canonicals diff --git a/tests/test_datamodels.py b/tests/test_datamodels.py index 307eeeb..220688a 100644 --- a/tests/test_datamodels.py +++ b/tests/test_datamodels.py @@ -7,7 +7,7 @@ def test_deduplication_scoring() -> None: """Test the deduplication scoring.""" d = DeduplicationResult( ["a", "b", "c"], - [DuplicateRecord("a", False, [("b", 0.9)]), DuplicateRecord("b", False, [("c", 0.8)])], + [DuplicateRecord("a", False, "b", 0.9), DuplicateRecord("b", False, "c", 0.8)], 0.8, ) assert d.duplicate_ratio == 0.4 @@ -17,7 +17,7 @@ def test_deduplication_scoring_exact() -> None: """Test the deduplication scoring.""" d = DeduplicationResult( ["a", "b", "c"], - [DuplicateRecord("a", True, [("b", 0.9)]), DuplicateRecord("b", False, [("c", 0.8)])], + [DuplicateRecord("a", True, "b", 0.9), DuplicateRecord("b", False, "c", 0.8)], 0.8, ) assert d.exact_duplicate_ratio == 0.2 @@ -30,23 +30,17 @@ def test_deduplication_scoring_empty() -> None: assert d.exact_duplicate_ratio == 0.0 -def test_rethreshold() -> None: - """Test rethresholding the duplicates, including empty case.""" - d = DuplicateRecord("a", False, [("b", 0.9), ("c", 0.8)]) - d._rethreshold(0.85) - assert d.duplicates == [("b", 0.9)] - - # Empty case - d_empty = DuplicateRecord("a", False, []) - d_empty._rethreshold(0.85) - assert d_empty.duplicates == [] +def test_duplicates_is_deprecated() -> None: + """Reading duplicates warns and returns the record it duplicates with its score.""" + with pytest.warns(DeprecationWarning, match="duplicate_of"): + assert DuplicateRecord("a", False, "b", 0.9).duplicates == [("b", 0.9)] def test_get_least_similar_from_duplicates() -> None: """Test getting the least similar duplicates, including empty case.""" d = DeduplicationResult( ["a", "b", "c"], - [DuplicateRecord("a", False, [("b", 0.9), ("c", 0.7)]), DuplicateRecord("b", False, [("c", 0.8)])], + [DuplicateRecord("a", False, "c", 0.7), DuplicateRecord("b", False, "c", 0.8)], 0.8, ) result = d.get_least_similar_from_duplicates(1) @@ -62,13 +56,13 @@ def test_rethreshold_deduplication_result() -> None: d = DeduplicationResult( ["a", "b", "c"], [ - DuplicateRecord("d", False, [("x", 0.9), ("y", 0.8)]), - DuplicateRecord("e", False, [("z", 0.8)]), + DuplicateRecord("d", False, "x", 0.9), + DuplicateRecord("e", False, "z", 0.8), ], 0.8, ) d.rethreshold(0.85) - assert d.filtered == [DuplicateRecord("d", False, [("x", 0.9)])] + assert d.filtered == [DuplicateRecord("d", False, "x", 0.9)] assert d.selected == ["a", "b", "c", "e"] @@ -77,8 +71,8 @@ def test_rethreshold_exception() -> None: d = DeduplicationResult( ["a", "b", "c"], [ - DuplicateRecord("d", False, [("x", 0.9), ("y", 0.8)]), - DuplicateRecord("e", False, [("z", 0.8)]), + DuplicateRecord("d", False, "x", 0.9), + DuplicateRecord("e", False, "z", 0.8), ], 0.7, ) @@ -91,8 +85,8 @@ def test_selected_with_duplicates_strings() -> None: d = DeduplicationResult( selected=["original"], filtered=[ - DuplicateRecord("duplicate_1", False, [("original", 0.9)]), - DuplicateRecord("duplicate_2", False, [("original", 0.8)]), + DuplicateRecord("duplicate_1", False, "original", 0.9), + DuplicateRecord("duplicate_2", False, "original", 0.8), ], threshold=0.8, ) @@ -112,8 +106,8 @@ def test_selected_with_duplicates_dicts() -> None: d = DeduplicationResult( selected=[selected], filtered=[ - DuplicateRecord({"id": 1, "text": "hello"}, True, [(selected, 1.0)]), - DuplicateRecord({"id": 2, "text": "helllo"}, False, [(selected, 0.1)]), + DuplicateRecord({"id": 1, "text": "hello"}, True, selected, 1.0), + DuplicateRecord({"id": 2, "text": "helllo"}, False, selected, 0.1), ], threshold=0.8, columns=["text"], @@ -133,8 +127,8 @@ def test_selected_with_duplicates_multi_column() -> None: d = DeduplicationResult( selected=[selected], filtered=[ - DuplicateRecord({"text": "hello", "text2": "world"}, True, [(selected, 1.0)]), - DuplicateRecord({"text": "helllo", "text2": "world"}, False, [(selected, 0.1)]), + DuplicateRecord({"text": "hello", "text2": "world"}, True, selected, 1.0), + DuplicateRecord({"text": "helllo", "text2": "world"}, False, selected, 0.1), ], threshold=0.8, columns=["text", "text2"], @@ -153,7 +147,7 @@ def test_selected_with_duplicates_unhashable_values() -> None: d = DeduplicationResult( selected=[selected], - filtered=[DuplicateRecord(filtered, exact=False, duplicates=[(selected, 1.0)])], + filtered=[DuplicateRecord(filtered, exact=False, duplicate_of=selected, score=1.0)], threshold=0.8, columns=["text"], ) @@ -162,16 +156,16 @@ def test_selected_with_duplicates_unhashable_values() -> None: assert items == [SelectedWithDuplicates(record=selected, duplicates=[(filtered, 1.0)])] -def test_selected_with_duplicates_removes_internal_duplicates() -> None: - """Test that selected_with_duplicates removes internal duplicates that have the same hash.""" +def test_selected_with_duplicates_preserves_occurrences() -> None: + """Identical record values must not erase distinct filtered occurrences.""" selected = {"id": 0, "text": "hello"} filtered = {"id": 1, "text": "hello"} d = DeduplicationResult( selected=[selected], filtered=[ - DuplicateRecord(filtered, exact=False, duplicates=[(selected, 0.95)]), - DuplicateRecord(filtered, exact=False, duplicates=[(selected, 0.90)]), + DuplicateRecord(filtered, exact=False, duplicate_of=selected, score=0.95), + DuplicateRecord(filtered, exact=False, duplicate_of=selected, score=0.90), ], threshold=0.8, columns=["text"], @@ -184,9 +178,7 @@ def test_selected_with_duplicates_removes_internal_duplicates() -> None: duplicate_list = items[0].duplicates # Should keep the kept record unchanged assert selected_record == selected - # The duplicate row must appear only once - assert len(duplicate_list) == 1 - assert duplicate_list[0][0] == filtered + assert duplicate_list == [(filtered, 0.95), (filtered, 0.90)] def test_selected_with_duplicates_caching() -> None: @@ -194,8 +186,8 @@ def test_selected_with_duplicates_caching() -> None: d = DeduplicationResult( selected=["original"], filtered=[ - DuplicateRecord("duplicate_1", False, [("original", 0.9)]), - DuplicateRecord("duplicate_2", False, [("original", 0.8)]), + DuplicateRecord("duplicate_1", False, "original", 0.9), + DuplicateRecord("duplicate_2", False, "original", 0.8), ], threshold=0.8, ) @@ -212,9 +204,9 @@ def test_selected_with_duplicates_cache_invalidation_on_rethreshold() -> None: d = DeduplicationResult( selected=["original"], filtered=[ - DuplicateRecord("duplicate_1", False, [("original", 0.9)]), - DuplicateRecord("duplicate_2", False, [("original", 0.8)]), - DuplicateRecord("duplicate_3", False, [("original", 0.7)]), + DuplicateRecord("duplicate_1", False, "original", 0.9), + DuplicateRecord("duplicate_2", False, "original", 0.8), + DuplicateRecord("duplicate_3", False, "original", 0.7), ], threshold=0.7, ) diff --git a/tests/test_semhash.py b/tests/test_semhash.py index bf3ffb2..416c071 100644 --- a/tests/test_semhash.py +++ b/tests/test_semhash.py @@ -1,3 +1,6 @@ +from collections.abc import Sequence +from typing import Any + import numpy as np import pytest @@ -132,15 +135,17 @@ def test_deduplicate_with_only_exact_duplicates(model: Encoder) -> None: ] semhash = SemHash.from_records(texts1, model=model) deduplicated = semhash.self_deduplicate() + deduplicated.rethreshold(0.99) assert deduplicated.selected == ["It's dangerous to go alone!"] # Each copy lists only the kept record, so the output grows linearly with the number of copies. - assert [d.duplicates for d in deduplicated.filtered] == [[("It's dangerous to go alone!", 1.0)]] * 2 + assert [(d.exact, d.duplicate_of, d.score) for d in deduplicated.filtered] == [(True, texts1[0], 1.0)] * 2 deduplicated = semhash.deduplicate(texts2) + deduplicated.rethreshold(0.99) assert deduplicated.selected == [] # Records are mapped back to strings, also when every record is an exact duplicate. assert [d.record for d in deduplicated.filtered] == texts2 - assert [d.duplicates for d in deduplicated.filtered] == [[("It's dangerous to go alone!", 1.0)]] * 3 + assert [(d.exact, d.duplicate_of, d.score) for d in deduplicated.filtered] == [(True, texts2[0], 1.0)] * 3 def test_rethreshold_keeps_exact_duplicate_group(model: Encoder) -> None: @@ -378,3 +383,73 @@ def test_deduplicate_edge_cases(model: Encoder) -> None: # Type mismatch: mixed dicts with pytest.raises(ValueError, match="Records must be all dictionaries"): semhash_dict.deduplicate([{"col": "a"}, "b"], threshold=0.95) + + +@pytest.fixture +def angular_model() -> Encoder: + """Encode known angles so similarity thresholds do not depend on a trained model.""" + + class AngularEncoder: + def encode(self, inputs: Sequence[Any] | Any, **kwargs: Any) -> np.ndarray: + angles = np.deg2rad([{"A": 0, "B": 40, "C": 50}[text] for text in inputs]) + return np.column_stack((np.cos(angles), np.sin(angles))).astype(np.float32) + + return AngularEncoder() + + +@pytest.mark.parametrize("backend", ["basic", "usearch"]) +def test_cross_dataset_reports_one_canonical(angular_model: Encoder, backend: str) -> None: + """Report the best reference record, not every near match or exact copy.""" + records = [{"id": i, "text": text} for i, text in enumerate("ABB")] + semhash = SemHash.from_records(records, columns=["text"], model=angular_model, ann_backend=backend) + query = {"id": 3, "text": "C"} + result = semhash.deduplicate([query], threshold=0.6) + assert result.selected == [] + assert result.filtered[0].duplicate_of == records[1] + result.rethreshold(0.995) + assert result.selected == [query] + + +@pytest.mark.parametrize( + "texts,threshold,selected,targets", + [("ABBC", 0.6, [0], [0, 0, 0]), ("ABC", 0.7, [0, 2], [0]), ("ACB", 0.75, [0, 1], [1])], +) +def test_self_deduplication_uses_direct_canonicals( + angular_model: Encoder, texts: str, threshold: float, selected: list[int], targets: list[int] +) -> None: + """Every filtered record points to one selected record it directly matches, also after rethresholding.""" + records = [{"id": i, "text": text, "metadata": [i]} for i, text in enumerate(texts)] + semhash = SemHash.from_records(records, model=angular_model, columns=["text"], ann_backend="basic") + result = semhash.self_deduplicate(threshold) + assert [r["id"] for r in result.selected] == selected + assert [d.duplicate_of["id"] for d in result.filtered] == targets + reconstructed = [r for g in result.selected_with_duplicates for r in [g.record] + [d for d, _ in g.duplicates]] + assert sorted(reconstructed, key=lambda r: r["id"]) == records + for duplicate in result.filtered: + canonical = duplicate.duplicate_of + vectors = angular_model.encode([duplicate.record["text"], canonical["text"]]) + assert canonical in result.selected + assert duplicate.score == pytest.approx(float(vectors[0] @ vectors[1]), abs=1e-6) + assert duplicate.exact is (duplicate.record["text"] == canonical["text"]) + for cutoff in (0.95, 0.99): + result.rethreshold(cutoff) + assert result == semhash.self_deduplicate(cutoff) + + +def test_dense_cluster_beyond_neighbor_limit(angular_model: Encoder) -> None: + """A near-duplicate cluster larger than the ANN neighbor limit keeps a single record, next to a zero vector.""" + rng = np.random.default_rng(0) + cluster = rng.normal(size=16) + rng.normal(scale=0.05, size=(2000, 16)) + embeddings = np.vstack([np.zeros(16), cluster]) + semhash = SemHash.from_embeddings(embeddings, [str(i) for i in range(2001)], model=angular_model) + result = semhash.self_deduplicate(0.9) + assert result.selected == ["0", "1"] + result.rethreshold(0.95) + assert result.selected == ["0", "1"] + + +def test_zero_vectors_are_not_near_duplicates(model: Encoder) -> None: + """Texts that embed to zero vectors are never near duplicates, in self and cross deduplication.""" + semhash = SemHash.from_records(["", " ", "hello world"], model=model) + assert semhash.self_deduplicate().selected == ["", " ", "hello world"] + assert semhash.deduplicate([" ", "hello world"]).selected == [" "] diff --git a/uv.lock b/uv.lock index 6f56a0a..6169816 100644 --- a/uv.lock +++ b/uv.lock @@ -8,9 +8,12 @@ resolution-markers = [ ] [options] -exclude-newer = "2026-04-26T18:41:39.144666Z" +exclude-newer = "0001-01-01T00:00:00Z" # This has no effect and is included for backwards compatibility when using relative exclude-newer values. exclude-newer-span = "P1W" +[options.exclude-newer-package] +vicinity = false + [[package]] name = "asttokens" version = "2.4.1" @@ -1011,7 +1014,7 @@ requires-dist = [ { name = "pytest-coverage", marker = "extra == 'dev'" }, { name = "pyversity" }, { name = "ruff", marker = "extra == 'dev'" }, - { name = "vicinity", extras = ["usearch"], specifier = ">=0.4.3" }, + { name = "vicinity", extras = ["usearch"], specifier = ">=0.4.6" }, ] provides-extras = ["dev"] @@ -1251,16 +1254,16 @@ wheels = [ [[package]] name = "vicinity" -version = "0.4.3" +version = "0.4.6" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "numpy" }, { name = "orjson" }, { name = "tqdm" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/24/c8/dfdc955989f9be22f9ebbe998c7acf47b89b46ed326c9d13583dc7975e42/vicinity-0.4.3.tar.gz", hash = "sha256:4e302b2f9e31416fe578c236bfc61b8931335d2b6db773da5012e945e31dc5e1", size = 831636, upload-time = "2025-10-04T13:23:31.333Z" } +sdist = { url = "https://files.pythonhosted.org/packages/f3/8f/3d1d6b600c6e07cfa61c1c7e2b7c0527416ae540483545ae7add96268143/vicinity-0.4.6.tar.gz", hash = "sha256:ce3e33dab3e6f3f028dcfb8414836d19f600f570d5c0cce720ead088a0343cad", size = 816592, upload-time = "2026-09-26T07:21:46.791Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/bc/22/c20b4393c2dcde671062b481f9ffa0f61dae2dca2cea221d742204dce872/vicinity-0.4.3-py3-none-any.whl", hash = "sha256:05530088e49bf91d79c41cd9db998181fa7c3714fc738384be78e4840e2a3161", size = 28692, upload-time = "2025-10-04T13:23:29.513Z" }, + { url = "https://files.pythonhosted.org/packages/9c/6f/5ed3e84eb7ee82032a8894f54645a6cb5ffcceeb83b78133cb21b40e711f/vicinity-0.4.6-py3-none-any.whl", hash = "sha256:8612a49bf4190bccb58bb3a0c0d34c316653c6fffc01221827032e638da576a6", size = 30019, upload-time = "2026-09-26T07:21:45.419Z" }, ] [package.optional-dependencies]