From 8d19c4eb74bbb308dcae573e8bbc0975800dc4cc Mon Sep 17 00:00:00 2001 From: Pringled Date: Fri, 25 Sep 2026 15:17:25 +0200 Subject: [PATCH 01/12] fix: use consistent canonical duplicate references --- README.md | 12 ++- semhash/datamodels.py | 77 +++++++++++++------ semhash/semhash.py | 54 ++----------- tests/test_datamodels.py | 8 +- tests/test_duplicate_reporting.py | 123 ++++++++++++++++++++++++++++++ tests/test_semhash.py | 2 +- 6 files changed, 200 insertions(+), 76 deletions(-) create mode 100644 tests/test_duplicate_reporting.py diff --git a/README.md b/README.md index abb230f..1a69b97 100644 --- a/README.md +++ b/README.md @@ -177,7 +177,7 @@ print(f"Exact duplicate ratio: {result.exact_duplicate_ratio}") # Find edge cases to tune your threshold least_similar = result.get_least_similar_from_duplicates(n=5) -# Adjust threshold without re-deduplicating +# Adjust threshold using cached matches, without re-embedding result.rethreshold(0.95) # View each kept record with its duplicate cluster @@ -186,6 +186,16 @@ for item in result.selected_with_duplicates: print(f"Duplicates: {item.duplicates}") # List of (duplicate_text, similarity_score) ``` +Each filtered record's `.duplicates` is a one-element list containing its canonical record and their similarity. +For self-deduplication, canonicals are kept in input order; a group is filtered only if it directly matches +an already-kept canonical above the threshold. If several canonicals match, the highest-scoring one is used. +Cross-dataset deduplication similarly reports one reference canonical, not every matching reference record. + +`selected_with_duplicates` reconstructs the self-deduplication groups, preserving every occurrence, ID and metadata. +`exact` describes the match to the final canonical: if an exact group is removed as a near duplicate of another group, +all its records point to that group's canonical with the near-match score and `exact=False`. +`rethreshold()` reruns self-deduplication using the original cached group matches, not edits to the result lists. + ## Main Features - **Fast**: SemHash uses [model2vec](https://github.com/MinishLab/model2vec) to embed texts and [vicinity](https://github.com/MinishLab/vicinity) to perform similarity search, making it extremely fast. diff --git a/semhash/datamodels.py b/semhash/datamodels.py index 9f98c4b..6da4c3c 100644 --- a/semhash/datamodels.py +++ b/semhash/datamodels.py @@ -1,6 +1,5 @@ from __future__ import annotations -import json from collections import defaultdict from collections.abc import Hashable, Sequence from dataclasses import dataclass, field @@ -20,8 +19,8 @@ class DuplicateRecord(Generic[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. + exact: Whether the record matches its canonical exactly on the deduplication columns. + duplicates: The canonical record and its similarity score, stored as a one-element list. """ @@ -69,6 +68,43 @@ class DeduplicationResult(Generic[Record]): threshold: float = field(default=0.9) columns: Sequence[str] | None = field(default=None) + def __post_init__(self) -> None: + """Retain compact self-deduplication context without changing dataclass construction or serialization.""" + self._self_deduplication: tuple[list[list[Record]], list[list[tuple[int, float]]]] | None = None + + @classmethod + def _from_groups( + cls, + groups: list[list[Record]], + neighbors: list[list[tuple[int, float]]], + threshold: float, + columns: Sequence[str] | None, + ) -> DeduplicationResult[Record]: + """Assign each record to a directly matching kept canonical, preserving input group order.""" + result = cls(threshold=threshold, columns=columns) + kept: set[int] = set() + for i, group in enumerate(groups): + owner = max( + ((j, score) for j, score in neighbors[i] if j in kept and score >= threshold), + key=lambda match: match[1], + default=None, + ) + if owner is None: + kept.add(i) + result.selected.append(group[0]) + canonical, score = group[0], 1.0 + removed = group[1:] + else: + owner_index, score = owner + canonical = groups[owner_index][0] + removed = group + result.filtered.extend( + DuplicateRecord(record=record, exact=owner is None, duplicates=[(canonical, score)]) + for record in removed + ) + result._self_deduplication = (groups, neighbors) + return result + @property def duplicate_ratio(self) -> float: """Return the percentage of records dropped.""" @@ -100,12 +136,20 @@ 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 (state := getattr(self, "_self_deduplication", None)) is not None: + # Replay selection over cached group matches; filtered records must not keep each other filtered. + groups, neighbors = state + result = self._from_groups(groups, neighbors, threshold, self.columns) + self.selected, self.filtered = result.selected, result.filtered + else: + filtered = [] + for dup in self.filtered: + dup._rethreshold(threshold) + if not dup.duplicates: + self.selected.append(dup.record) + else: + filtered.append(dup) + self.filtered = filtered self.threshold = threshold @cached_property @@ -132,19 +176,8 @@ def _to_hashable(record: Record) -> frozendict[str, str] | str: 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/semhash.py b/semhash/semhash.py index 27e883d..84ee787 100644 --- a/semhash/semhash.py +++ b/semhash/semhash.py @@ -14,6 +14,7 @@ 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, @@ -192,7 +193,7 @@ def deduplicate( duplicate_records.append( DuplicateRecord( record=record, - duplicates=[(item, score) for item, score in similar_items], + duplicates=[max(similar_items, key=lambda match: match[1])], exact=False, ) ) @@ -217,54 +218,13 @@ 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 - ) - + indices = {id(item[0]): i for i, item in enumerate(self.index.items)} + neighbors = [[(indices[id(record)], score) for record, score in matches] for matches in results] + 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 + groups = [[dict_to_string(record, self.columns) for record in group] for group in groups] + return DeduplicationResult._from_groups(groups, neighbors, threshold, self.columns) def _validate_if_strings(self, records: Sequence[dict[str, Any] | str]) -> list[dict[str, Any]]: """ diff --git a/tests/test_datamodels.py b/tests/test_datamodels.py index 307eeeb..2d638e9 100644 --- a/tests/test_datamodels.py +++ b/tests/test_datamodels.py @@ -162,8 +162,8 @@ 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"} @@ -184,9 +184,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: diff --git a/tests/test_duplicate_reporting.py b/tests/test_duplicate_reporting.py new file mode 100644 index 0000000..563a985 --- /dev/null +++ b/tests/test_duplicate_reporting.py @@ -0,0 +1,123 @@ +import tracemalloc +from collections.abc import Sequence +from typing import Any + +import numpy as np +import pytest + +from semhash import SemHash +from semhash.utils import Encoder + + +@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"]) +@pytest.mark.parametrize("texts,query_text,canonical_id", [("AAA", "A", 0), ("AAA", "B", 0), ("AAB", "C", 2)]) +@pytest.mark.parametrize("columns", [["text"], ["text", "context"]]) +def test_cross_dataset_reports_one_canonical( + angular_model: Encoder, backend: str, texts: str, query_text: str, canonical_id: int, columns: list[str] +) -> None: + """Exact and near matches report just the best reference canonical, including its metadata.""" + records = [{"id": i, "text": text, "context": "C", "metadata": [i]} for i, text in enumerate(texts)] + queries = [{"id": i, "text": query_text, "context": "C"} for i in (3, 4)] + semhash = SemHash.from_records(records, columns=columns, model=angular_model, ann_backend=backend) + result = semhash.deduplicate(queries, threshold=0.6) + assert result.selected == [] + for duplicate, query in zip(result.filtered, queries): + assert duplicate.record == query + assert duplicate.exact is (query_text == "A") + assert len(duplicate.duplicates) == 1 + canonical, score = duplicate.duplicates[0] + assert canonical == records[canonical_id] + assert score >= 0.6 + result.rethreshold(0.995) + assert result.selected == ([] if query_text == "A" else queries) + + +@pytest.mark.parametrize("backend", ["basic", "usearch"]) +@pytest.mark.parametrize("from_embeddings", [False, True]) +def test_self_canonicals_preserve_records_and_rethreshold( + angular_model: Encoder, backend: str, from_embeddings: bool +) -> None: + """Every occurrence points directly to a kept canonical, including after semantic groups split.""" + records = [{"id": i, "text": text, "metadata": [i]} for i, text in enumerate(["A", "B", "B", "C"])] + if from_embeddings: + semhash = SemHash.from_embeddings( + angular_model.encode([record["text"] for record in records]), + records, + model=angular_model, + columns=["text"], + ann_backend=backend, + ) + else: + semhash = SemHash.from_records(records, model=angular_model, columns=["text"], ann_backend=backend) + result = semhash.self_deduplicate(0.6) + assert result.selected == records[:1] + for threshold in (0.6, 0.95, 0.99): + result.rethreshold(threshold) + assert result == semhash.self_deduplicate(threshold) + assert sorted(result.selected + [d.record for d in result.filtered], key=lambda r: r["id"]) == records + 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: + assert len(duplicate.duplicates) == 1 + canonical, score = duplicate.duplicates[0] + assert canonical in result.selected + assert score >= threshold + vectors = angular_model.encode([duplicate.record["text"], canonical["text"]]) + # USEARCH may quantize vectors (e.g. BF16), unlike BASIC's full-precision cosine. + tolerance = 0.01 if backend == "usearch" else 1e-6 + assert score == pytest.approx(float(vectors[0] @ vectors[1]), abs=tolerance) + assert duplicate.exact is (duplicate.record["text"] == canonical["text"]) + assert [r["id"] for r in result.selected] == [0, 1, 3] + assert result.selected_with_duplicates[1].duplicates == [(records[2], 1.0)] + + +@pytest.mark.parametrize( + "texts,threshold,selected,links", + [ + ("ABC", 0.7, ["A", "C"], [("B", "A")]), + ("ACB", 0.75, ["A", "C"], [("B", "C")]), + ("ABC", 0.6, ["A"], [("B", "A"), ("C", "A")]), + ], +) +def test_near_canonical_requires_a_direct_best_match( + angular_model: Encoder, texts: str, threshold: float, selected: list[str], links: list[tuple[str, str]] +) -> None: + """Near matches use the best kept canonical, never transitive links through filtered records.""" + semhash = SemHash.from_records(list(texts), model=angular_model, ann_backend="basic") + result = semhash.self_deduplicate(threshold) + assert result.selected == selected + assert all(len(d.duplicates) == 1 for d in result.filtered) + assert [(d.record, d.duplicates[0][0]) for d in result.filtered] == links + + +def test_large_exact_groups_have_linear_reporting(angular_model: Encoder) -> None: + """Self/cross results, grouping, edge inspection and rethresholding avoid all-pairs allocation.""" + n = 1000 + semhash = SemHash.from_records(["A"] * n, model=angular_model, ann_backend="basic") + tracemalloc.start() + try: + result = semhash.self_deduplicate() + cross = semhash.deduplicate(["A"] * n) + assert len(result.filtered) == n - 1 + assert len(cross.filtered) == n + assert all(d.duplicates == [("A", 1.0)] for d in result.filtered + cross.filtered) + assert len(result.selected_with_duplicates[0].duplicates) == n - 1 + assert result.get_least_similar_from_duplicates(3) == [("A", "A", 1.0)] * 3 + result.rethreshold(0.99) + cross.rethreshold(0.99) + _, peak = tracemalloc.get_traced_memory() + assert peak < 8 * 1024 * 1024 + finally: + tracemalloc.stop() diff --git a/tests/test_semhash.py b/tests/test_semhash.py index bf3ffb2..9bef786 100644 --- a/tests/test_semhash.py +++ b/tests/test_semhash.py @@ -144,7 +144,7 @@ def test_deduplicate_with_only_exact_duplicates(model: Encoder) -> None: def test_rethreshold_keeps_exact_duplicate_group(model: Encoder) -> None: - """A near-duplicate with exact copies is not listed as a duplicate of its own copies, so rethresholding keeps it.""" + """Exact copies must not prevent their representative from being restored by rethresholding.""" records = [ {"text": "It's dangerous to go alone!", "id": 1}, {"text": "It's dangerous to go alone! Take this.", "id": 2}, From d89f0f225978a85f8538d2772ab6e4593c6099cb Mon Sep 17 00:00:00 2001 From: Pringled Date: Fri, 25 Sep 2026 15:40:48 +0200 Subject: [PATCH 02/12] refactor: trim canonical reporting tests and documentation --- README.md | 12 +--- semhash/datamodels.py | 2 +- tests/test_duplicate_reporting.py | 95 +++++++++---------------------- 3 files changed, 29 insertions(+), 80 deletions(-) diff --git a/README.md b/README.md index 1a69b97..abb230f 100644 --- a/README.md +++ b/README.md @@ -177,7 +177,7 @@ print(f"Exact duplicate ratio: {result.exact_duplicate_ratio}") # Find edge cases to tune your threshold least_similar = result.get_least_similar_from_duplicates(n=5) -# Adjust threshold using cached matches, without re-embedding +# Adjust threshold without re-deduplicating result.rethreshold(0.95) # View each kept record with its duplicate cluster @@ -186,16 +186,6 @@ for item in result.selected_with_duplicates: print(f"Duplicates: {item.duplicates}") # List of (duplicate_text, similarity_score) ``` -Each filtered record's `.duplicates` is a one-element list containing its canonical record and their similarity. -For self-deduplication, canonicals are kept in input order; a group is filtered only if it directly matches -an already-kept canonical above the threshold. If several canonicals match, the highest-scoring one is used. -Cross-dataset deduplication similarly reports one reference canonical, not every matching reference record. - -`selected_with_duplicates` reconstructs the self-deduplication groups, preserving every occurrence, ID and metadata. -`exact` describes the match to the final canonical: if an exact group is removed as a near duplicate of another group, -all its records point to that group's canonical with the near-match score and `exact=False`. -`rethreshold()` reruns self-deduplication using the original cached group matches, not edits to the result lists. - ## Main Features - **Fast**: SemHash uses [model2vec](https://github.com/MinishLab/model2vec) to embed texts and [vicinity](https://github.com/MinishLab/vicinity) to perform similarity search, making it extremely fast. diff --git a/semhash/datamodels.py b/semhash/datamodels.py index 6da4c3c..cc8e9e1 100644 --- a/semhash/datamodels.py +++ b/semhash/datamodels.py @@ -69,7 +69,7 @@ class DeduplicationResult(Generic[Record]): columns: Sequence[str] | None = field(default=None) def __post_init__(self) -> None: - """Retain compact self-deduplication context without changing dataclass construction or serialization.""" + """Initialize the cache used for rethresholding.""" self._self_deduplication: tuple[list[list[Record]], list[list[tuple[int, float]]]] | None = None @classmethod diff --git a/tests/test_duplicate_reporting.py b/tests/test_duplicate_reporting.py index 563a985..ea787a7 100644 --- a/tests/test_duplicate_reporting.py +++ b/tests/test_duplicate_reporting.py @@ -22,88 +22,48 @@ def encode(self, inputs: Sequence[Any] | Any, **kwargs: Any) -> np.ndarray: @pytest.mark.parametrize("backend", ["basic", "usearch"]) -@pytest.mark.parametrize("texts,query_text,canonical_id", [("AAA", "A", 0), ("AAA", "B", 0), ("AAB", "C", 2)]) -@pytest.mark.parametrize("columns", [["text"], ["text", "context"]]) -def test_cross_dataset_reports_one_canonical( - angular_model: Encoder, backend: str, texts: str, query_text: str, canonical_id: int, columns: list[str] -) -> None: - """Exact and near matches report just the best reference canonical, including its metadata.""" - records = [{"id": i, "text": text, "context": "C", "metadata": [i]} for i, text in enumerate(texts)] - queries = [{"id": i, "text": query_text, "context": "C"} for i in (3, 4)] - semhash = SemHash.from_records(records, columns=columns, model=angular_model, ann_backend=backend) - result = semhash.deduplicate(queries, threshold=0.6) +def test_cross_dataset_reports_one_canonical(angular_model: Encoder, backend: str) -> None: + """Report the best reference canonical, 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 == [] - for duplicate, query in zip(result.filtered, queries): - assert duplicate.record == query - assert duplicate.exact is (query_text == "A") - assert len(duplicate.duplicates) == 1 - canonical, score = duplicate.duplicates[0] - assert canonical == records[canonical_id] - assert score >= 0.6 + assert len(result.filtered[0].duplicates) == 1 + assert result.filtered[0].duplicates[0][0] == records[1] result.rethreshold(0.995) - assert result.selected == ([] if query_text == "A" else queries) + assert result.selected == [query] -@pytest.mark.parametrize("backend", ["basic", "usearch"]) -@pytest.mark.parametrize("from_embeddings", [False, True]) -def test_self_canonicals_preserve_records_and_rethreshold( - angular_model: Encoder, backend: str, from_embeddings: bool +@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_canonicals_and_rethreshold( + angular_model: Encoder, texts: str, threshold: float, selected: list[int], targets: list[int] ) -> None: - """Every occurrence points directly to a kept canonical, including after semantic groups split.""" - records = [{"id": i, "text": text, "metadata": [i]} for i, text in enumerate(["A", "B", "B", "C"])] - if from_embeddings: - semhash = SemHash.from_embeddings( - angular_model.encode([record["text"] for record in records]), - records, - model=angular_model, - columns=["text"], - ann_backend=backend, - ) - else: - semhash = SemHash.from_records(records, model=angular_model, columns=["text"], ann_backend=backend) - result = semhash.self_deduplicate(0.6) - assert result.selected == records[:1] - for threshold in (0.6, 0.95, 0.99): - result.rethreshold(threshold) - assert result == semhash.self_deduplicate(threshold) - assert sorted(result.selected + [d.record for d in result.filtered], key=lambda r: r["id"]) == records + """Preserve complete groups with direct canonical links, including when higher thresholds split them.""" + 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.duplicates[0][0]["id"] for d in result.filtered] == targets + for cutoff in (threshold, 0.95, 0.99): + result.rethreshold(cutoff) + assert result == semhash.self_deduplicate(cutoff) 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: assert len(duplicate.duplicates) == 1 canonical, score = duplicate.duplicates[0] - assert canonical in result.selected - assert score >= threshold + assert canonical in result.selected and score >= cutoff vectors = angular_model.encode([duplicate.record["text"], canonical["text"]]) - # USEARCH may quantize vectors (e.g. BF16), unlike BASIC's full-precision cosine. - tolerance = 0.01 if backend == "usearch" else 1e-6 - assert score == pytest.approx(float(vectors[0] @ vectors[1]), abs=tolerance) + assert score == pytest.approx(float(vectors[0] @ vectors[1]), abs=1e-6) assert duplicate.exact is (duplicate.record["text"] == canonical["text"]) - assert [r["id"] for r in result.selected] == [0, 1, 3] - assert result.selected_with_duplicates[1].duplicates == [(records[2], 1.0)] - - -@pytest.mark.parametrize( - "texts,threshold,selected,links", - [ - ("ABC", 0.7, ["A", "C"], [("B", "A")]), - ("ACB", 0.75, ["A", "C"], [("B", "C")]), - ("ABC", 0.6, ["A"], [("B", "A"), ("C", "A")]), - ], -) -def test_near_canonical_requires_a_direct_best_match( - angular_model: Encoder, texts: str, threshold: float, selected: list[str], links: list[tuple[str, str]] -) -> None: - """Near matches use the best kept canonical, never transitive links through filtered records.""" - semhash = SemHash.from_records(list(texts), model=angular_model, ann_backend="basic") - result = semhash.self_deduplicate(threshold) - assert result.selected == selected - assert all(len(d.duplicates) == 1 for d in result.filtered) - assert [(d.record, d.duplicates[0][0]) for d in result.filtered] == links def test_large_exact_groups_have_linear_reporting(angular_model: Encoder) -> None: - """Self/cross results, grouping, edge inspection and rethresholding avoid all-pairs allocation.""" + """Self/cross results, grouping and rethresholding avoid all-pairs allocation.""" n = 1000 semhash = SemHash.from_records(["A"] * n, model=angular_model, ann_backend="basic") tracemalloc.start() @@ -114,7 +74,6 @@ def test_large_exact_groups_have_linear_reporting(angular_model: Encoder) -> Non assert len(cross.filtered) == n assert all(d.duplicates == [("A", 1.0)] for d in result.filtered + cross.filtered) assert len(result.selected_with_duplicates[0].duplicates) == n - 1 - assert result.get_least_similar_from_duplicates(3) == [("A", "A", 1.0)] * 3 result.rethreshold(0.99) cross.rethreshold(0.99) _, peak = tracemalloc.get_traced_memory() From 1ab13869d8e769c0ee812cd244bdd37e8d17c2f7 Mon Sep 17 00:00:00 2001 From: Pringled Date: Fri, 25 Sep 2026 15:54:55 +0200 Subject: [PATCH 03/12] refactor: simplify threshold results and align naming --- semhash/datamodels.py | 32 ++++++++++++++++---------------- semhash/index.py | 21 ++++++--------------- semhash/semhash.py | 7 +++---- 3 files changed, 25 insertions(+), 35 deletions(-) diff --git a/semhash/datamodels.py b/semhash/datamodels.py index cc8e9e1..eb37992 100644 --- a/semhash/datamodels.py +++ b/semhash/datamodels.py @@ -76,33 +76,33 @@ def __post_init__(self) -> None: def _from_groups( cls, groups: list[list[Record]], - neighbors: list[list[tuple[int, float]]], + results: list[list[tuple[int, float]]], threshold: float, columns: Sequence[str] | None, ) -> DeduplicationResult[Record]: """Assign each record to a directly matching kept canonical, preserving input group order.""" result = cls(threshold=threshold, columns=columns) - kept: set[int] = set() + selected_indices: set[int] = set() for i, group in enumerate(groups): - owner = max( - ((j, score) for j, score in neighbors[i] if j in kept and score >= threshold), + best_match = max( + ((j, score) for j, score in results[i] if j in selected_indices and score >= threshold), key=lambda match: match[1], default=None, ) - if owner is None: - kept.add(i) + if best_match is None: + selected_indices.add(i) result.selected.append(group[0]) - canonical, score = group[0], 1.0 - removed = group[1:] + canonical_record, score = group[0], 1.0 + filtered_records = group[1:] else: - owner_index, score = owner - canonical = groups[owner_index][0] - removed = group + index, score = best_match + canonical_record = groups[index][0] + filtered_records = group result.filtered.extend( - DuplicateRecord(record=record, exact=owner is None, duplicates=[(canonical, score)]) - for record in removed + DuplicateRecord(record=record, exact=best_match is None, duplicates=[(canonical_record, score)]) + for record in filtered_records ) - result._self_deduplication = (groups, neighbors) + result._self_deduplication = (groups, results) return result @property @@ -138,8 +138,8 @@ def rethreshold(self, threshold: float) -> None: self.__dict__.pop("selected_with_duplicates", None) if (state := getattr(self, "_self_deduplication", None)) is not None: # Replay selection over cached group matches; filtered records must not keep each other filtered. - groups, neighbors = state - result = self._from_groups(groups, neighbors, threshold, self.columns) + groups, results = state + result = self._from_groups(groups, results, threshold, self.columns) self.selected, self.filtered = result.selected, result.filtered else: filtered = [] diff --git a/semhash/index.py b/semhash/index.py index 207a860..e35812f 100644 --- a/semhash/index.py +++ b/semhash/index.py @@ -7,8 +7,6 @@ from vicinity.backends import AbstractBackend, get_backend_class from vicinity.datatypes import SingleQueryResult -DocScore = tuple[dict[str, str], float] -DocScores = list[DocScore] DictItem = list[dict[str, str]] @@ -47,25 +45,18 @@ 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[list[tuple[int, float]]]: """ Query the index with a threshold. :param vectors: The vectors to query. :param threshold: The similarity threshold. - :return: The query results. + :return: 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 + return [ + [(int(index), 1 - distance) for index, distance in zip(*result)] + for result in self.backend.threshold(vectors, threshold=1 - threshold, max_k=100) + ] def query_top_k(self, vectors: np.ndarray, k: int, vectors_are_in_index: bool) -> list[SingleQueryResult]: """ diff --git a/semhash/semhash.py b/semhash/semhash.py index 84ee787..d74fa8f 100644 --- a/semhash/semhash.py +++ b/semhash/semhash.py @@ -190,10 +190,11 @@ def deduplicate( # No duplicates found, keep this record deduplicated_records.append(record) else: + index, score = max(similar_items, key=lambda match: match[1]) duplicate_records.append( DuplicateRecord( record=record, - duplicates=[max(similar_items, key=lambda match: match[1])], + duplicates=[(self.index.items[index][0], score)], exact=False, ) ) @@ -219,12 +220,10 @@ def self_deduplicate( :return: A deduplicated list of records. """ results = self.index.query_threshold(self.index.vectors, threshold=threshold) - indices = {id(item[0]): i for i, item in enumerate(self.index.items)} - neighbors = [[(indices[id(record)], score) for record, score in matches] for matches in results] groups: list[list[Any]] = self.index.items if self._was_string: groups = [[dict_to_string(record, self.columns) for record in group] for group in groups] - return DeduplicationResult._from_groups(groups, neighbors, threshold, self.columns) + return DeduplicationResult._from_groups(groups, results, threshold, self.columns) def _validate_if_strings(self, records: Sequence[dict[str, Any] | str]) -> list[dict[str, Any]]: """ From f831500d69c31b3a62c43bf8612fe9375d7ec2b6 Mon Sep 17 00:00:00 2001 From: Pringled Date: Fri, 25 Sep 2026 16:42:21 +0200 Subject: [PATCH 04/12] fix: keep dense near-duplicate clusters deduplicated beyond the ANN neighbor limit Neighbors propose their canonical (in both directions) and each candidate is verified by direct cosine similarity, so matches stay non-transitive. When a record's neighbor list is truncated at MAX_NEIGHBORS without a match, it is compared against all selected records directly. --- semhash/datamodels.py | 52 +++++++++++++++++++++++-------- semhash/index.py | 3 +- semhash/semhash.py | 2 +- tests/test_duplicate_reporting.py | 11 +++++++ 4 files changed, 53 insertions(+), 15 deletions(-) diff --git a/semhash/datamodels.py b/semhash/datamodels.py index eb37992..398e7e7 100644 --- a/semhash/datamodels.py +++ b/semhash/datamodels.py @@ -6,8 +6,10 @@ from functools import cached_property from typing import Generic +import numpy as np from frozendict import frozendict +from semhash.index import MAX_NEIGHBORS from semhash.utils import DuplicateList, Record, to_frozendict @@ -70,39 +72,63 @@ class DeduplicationResult(Generic[Record]): def __post_init__(self) -> None: """Initialize the cache used for rethresholding.""" - self._self_deduplication: tuple[list[list[Record]], list[list[tuple[int, float]]]] | None = None + self._self_deduplication: tuple[list[list[Record]], list[list[tuple[int, float]]], np.ndarray] | None = None @classmethod def _from_groups( cls, groups: list[list[Record]], results: list[list[tuple[int, float]]], + vectors: np.ndarray, threshold: float, columns: Sequence[str] | None, ) -> DeduplicationResult[Record]: """Assign each record to a directly matching kept canonical, preserving input group order.""" result = cls(threshold=threshold, columns=columns) - selected_indices: set[int] = set() + norms = np.linalg.norm(vectors, axis=1) + # Normalized vectors of selected records, in selection order, for comparing against all of them at once. + selected_vectors = np.empty(vectors.shape, dtype=np.float32) + selected_indices: list[int] = [] + canonical_indices: dict[int, int] = {} + # Canonicals proposed by earlier neighbors, so a match missed by one query is still found through the other. + proposals: defaultdict[int, set[int]] = defaultdict(set) + + def closest(i: int, candidates: list[int], candidate_vectors: np.ndarray) -> tuple[int, float] | None: + if not candidates: + return None + scores = candidate_vectors @ (vectors[i] / norms[i]) + best = int(np.argmax(scores)) + return (candidates[best], float(scores[best])) if scores[best] >= threshold else None + for i, group in enumerate(groups): - best_match = max( - ((j, score) for j, score in results[i] if j in selected_indices and score >= threshold), - key=lambda match: match[1], - default=None, + matches = [j for j, score in results[i] if score >= threshold] + # Neighbors propose their canonical, whose similarity is checked directly to avoid transitive matches. + candidates = list( + proposals.pop(i, set()).union(canonical_indices[j] for j in matches if j in canonical_indices) ) + best_match = closest(i, candidates, vectors[candidates] / norms[candidates, None]) + if best_match is None and len(matches) >= MAX_NEIGHBORS: + # The neighbors may be truncated, so compare against every selected record directly. + best_match = closest(i, selected_indices, selected_vectors[: len(selected_indices)]) if best_match is None: - selected_indices.add(i) + canonical_index, score = i, 1.0 + selected_vectors[len(selected_indices)] = vectors[i] / norms[i] + selected_indices.append(i) result.selected.append(group[0]) - canonical_record, score = group[0], 1.0 filtered_records = group[1:] else: - index, score = best_match - canonical_record = groups[index][0] + canonical_index, score = best_match filtered_records = group + canonical_indices[i] = canonical_index + canonical_record = groups[canonical_index][0] + for j in matches: + if j > i: + proposals[j].add(canonical_index) result.filtered.extend( DuplicateRecord(record=record, exact=best_match is None, duplicates=[(canonical_record, score)]) for record in filtered_records ) - result._self_deduplication = (groups, results) + result._self_deduplication = (groups, results, vectors) return result @property @@ -138,8 +164,8 @@ def rethreshold(self, threshold: float) -> None: self.__dict__.pop("selected_with_duplicates", None) if (state := getattr(self, "_self_deduplication", None)) is not None: # Replay selection over cached group matches; filtered records must not keep each other filtered. - groups, results = state - result = self._from_groups(groups, results, threshold, self.columns) + groups, results, vectors = state + result = self._from_groups(groups, results, vectors, threshold, self.columns) self.selected, self.filtered = result.selected, result.filtered else: filtered = [] diff --git a/semhash/index.py b/semhash/index.py index e35812f..f079f4b 100644 --- a/semhash/index.py +++ b/semhash/index.py @@ -8,6 +8,7 @@ from vicinity.datatypes import SingleQueryResult DictItem = list[dict[str, str]] +MAX_NEIGHBORS = 100 class Index: @@ -55,7 +56,7 @@ def query_threshold(self, vectors: np.ndarray, threshold: float) -> list[list[tu """ return [ [(int(index), 1 - distance) for index, distance in zip(*result)] - for result in self.backend.threshold(vectors, threshold=1 - threshold, max_k=100) + for result in self.backend.threshold(vectors, threshold=1 - threshold, max_k=MAX_NEIGHBORS) ] def query_top_k(self, vectors: np.ndarray, k: int, vectors_are_in_index: bool) -> list[SingleQueryResult]: diff --git a/semhash/semhash.py b/semhash/semhash.py index d74fa8f..8bf1ab0 100644 --- a/semhash/semhash.py +++ b/semhash/semhash.py @@ -223,7 +223,7 @@ def self_deduplicate( groups: list[list[Any]] = self.index.items if self._was_string: groups = [[dict_to_string(record, self.columns) for record in group] for group in groups] - return DeduplicationResult._from_groups(groups, results, threshold, self.columns) + return DeduplicationResult._from_groups(groups, results, self.index.vectors, threshold, self.columns) def _validate_if_strings(self, records: Sequence[dict[str, Any] | str]) -> list[dict[str, Any]]: """ diff --git a/tests/test_duplicate_reporting.py b/tests/test_duplicate_reporting.py index ea787a7..be54f28 100644 --- a/tests/test_duplicate_reporting.py +++ b/tests/test_duplicate_reporting.py @@ -62,6 +62,17 @@ def test_self_canonicals_and_rethreshold( assert duplicate.exact is (duplicate.record["text"] == canonical["text"]) +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.""" + rng = np.random.default_rng(0) + embeddings = rng.normal(size=16) + rng.normal(scale=0.05, size=(2000, 16)) + semhash = SemHash.from_embeddings(embeddings, [str(i) for i in range(2000)], model=angular_model) + result = semhash.self_deduplicate(0.9) + assert result.selected == ["0"] + result.rethreshold(0.95) + assert result.selected == ["0"] + + def test_large_exact_groups_have_linear_reporting(angular_model: Encoder) -> None: """Self/cross results, grouping and rethresholding avoid all-pairs allocation.""" n = 1000 From ade60f8b8d42964c31023d25f94b836cec75e30a Mon Sep 17 00:00:00 2001 From: Pringled Date: Fri, 25 Sep 2026 17:02:27 +0200 Subject: [PATCH 05/12] fix: treat zero vectors as matching nothing in canonical selection Zero embeddings (e.g. empty text) produced NaN normalized vectors, which could make the all-selected fallback pick NaN and keep an extra canonical. --- semhash/datamodels.py | 2 ++ tests/test_duplicate_reporting.py | 11 ++++++----- 2 files changed, 8 insertions(+), 5 deletions(-) diff --git a/semhash/datamodels.py b/semhash/datamodels.py index 398e7e7..46f96d8 100644 --- a/semhash/datamodels.py +++ b/semhash/datamodels.py @@ -86,6 +86,8 @@ def _from_groups( """Assign each record to a directly matching kept canonical, preserving input group order.""" result = cls(threshold=threshold, columns=columns) norms = np.linalg.norm(vectors, axis=1) + # Zero vectors (e.g. empty text) are similar to nothing, instead of producing NaN scores. + norms[norms == 0] = 1.0 # Normalized vectors of selected records, in selection order, for comparing against all of them at once. selected_vectors = np.empty(vectors.shape, dtype=np.float32) selected_indices: list[int] = [] diff --git a/tests/test_duplicate_reporting.py b/tests/test_duplicate_reporting.py index be54f28..6cf24a6 100644 --- a/tests/test_duplicate_reporting.py +++ b/tests/test_duplicate_reporting.py @@ -63,14 +63,15 @@ def test_self_canonicals_and_rethreshold( 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.""" + """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) - embeddings = rng.normal(size=16) + rng.normal(scale=0.05, size=(2000, 16)) - semhash = SemHash.from_embeddings(embeddings, [str(i) for i in range(2000)], model=angular_model) + 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"] + assert result.selected == ["0", "1"] result.rethreshold(0.95) - assert result.selected == ["0"] + assert result.selected == ["0", "1"] def test_large_exact_groups_have_linear_reporting(angular_model: Encoder) -> None: From 17a2b32e062acc229ac452b0e62ec0d1b808aab6 Mon Sep 17 00:00:00 2001 From: Pringled Date: Fri, 25 Sep 2026 17:21:33 +0200 Subject: [PATCH 06/12] fix: rescore cross-dataset matches with exact cosine similarity Cross-dataset deduplication now verifies ANN neighbors with exact cosine similarity, like self-deduplication, so scores are exact and zero vectors (e.g. empty text) are never reported as near duplicates. --- semhash/semhash.py | 14 ++++++++++---- tests/test_duplicate_reporting.py | 7 +++++++ 2 files changed, 17 insertions(+), 4 deletions(-) diff --git a/semhash/semhash.py b/semhash/semhash.py index 8bf1ab0..95a7926 100644 --- a/semhash/semhash.py +++ b/semhash/semhash.py @@ -180,21 +180,27 @@ def deduplicate( # Only embed and query the records that are left after removing exact duplicates results = [] + embeddings = np.empty((0, self.index.vectors.shape[1])) 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: + for record, embedding, similar_items in zip(dict_records, embeddings, results): + # Rescore the neighbors with exact cosine similarity, like self_deduplicate does. + indices = [index for index, _ in similar_items] + candidates = self.index.vectors[indices] + norms = np.linalg.norm(candidates, axis=1) * np.linalg.norm(embedding) + scores = candidates @ embedding / np.where(norms == 0, 1.0, norms) + best = int(np.argmax(scores)) if indices else 0 + if not indices or scores[best] < threshold: # No duplicates found, keep this record deduplicated_records.append(record) else: - index, score = max(similar_items, key=lambda match: match[1]) duplicate_records.append( DuplicateRecord( record=record, - duplicates=[(self.index.items[index][0], score)], + duplicates=[(self.index.items[indices[best]][0], float(scores[best]))], exact=False, ) ) diff --git a/tests/test_duplicate_reporting.py b/tests/test_duplicate_reporting.py index 6cf24a6..0c09a54 100644 --- a/tests/test_duplicate_reporting.py +++ b/tests/test_duplicate_reporting.py @@ -92,3 +92,10 @@ def test_large_exact_groups_have_linear_reporting(angular_model: Encoder) -> Non assert peak < 8 * 1024 * 1024 finally: tracemalloc.stop() + + +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 == [" "] From 279610fe38540dd645c1e29fcbfed6d21da28e06 Mon Sep 17 00:00:00 2001 From: Pringled Date: Sat, 26 Sep 2026 09:24:47 +0200 Subject: [PATCH 07/12] chore: require vicinity>=0.4.6 for correct backend distances --- pyproject.toml | 3 ++- uv.lock | 13 ++++++++----- 2 files changed, 10 insertions(+), 6 deletions(-) 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/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] From bc2221779694c63f8cce8d070d6c496fe829a8d3 Mon Sep 17 00:00:00 2001 From: Pringled Date: Sat, 26 Sep 2026 09:35:59 +0200 Subject: [PATCH 08/12] refactor: drop neighbor proposals and revert unneeded docstring changes --- semhash/datamodels.py | 17 +++++------------ tests/test_semhash.py | 2 +- 2 files changed, 6 insertions(+), 13 deletions(-) diff --git a/semhash/datamodels.py b/semhash/datamodels.py index 46f96d8..8c2dc4a 100644 --- a/semhash/datamodels.py +++ b/semhash/datamodels.py @@ -21,8 +21,8 @@ class DuplicateRecord(Generic[Record]): Attributes ---------- record: The original record being deduplicated. - exact: Whether the record matches its canonical exactly on the deduplication columns. - duplicates: The canonical record and its similarity score, stored as a one-element list. + exact: Whether the record was identified as an exact match. + duplicates: The canonical record and its similarity score. """ @@ -92,8 +92,6 @@ def _from_groups( selected_vectors = np.empty(vectors.shape, dtype=np.float32) selected_indices: list[int] = [] canonical_indices: dict[int, int] = {} - # Canonicals proposed by earlier neighbors, so a match missed by one query is still found through the other. - proposals: defaultdict[int, set[int]] = defaultdict(set) def closest(i: int, candidates: list[int], candidate_vectors: np.ndarray) -> tuple[int, float] | None: if not candidates: @@ -104,10 +102,8 @@ def closest(i: int, candidates: list[int], candidate_vectors: np.ndarray) -> tup for i, group in enumerate(groups): matches = [j for j, score in results[i] if score >= threshold] - # Neighbors propose their canonical, whose similarity is checked directly to avoid transitive matches. - candidates = list( - proposals.pop(i, set()).union(canonical_indices[j] for j in matches if j in canonical_indices) - ) + # Check the canonicals of matching neighbors directly, to avoid transitive matches. + candidates = list({canonical_indices[j] for j in matches if j in canonical_indices}) best_match = closest(i, candidates, vectors[candidates] / norms[candidates, None]) if best_match is None and len(matches) >= MAX_NEIGHBORS: # The neighbors may be truncated, so compare against every selected record directly. @@ -123,9 +119,6 @@ def closest(i: int, candidates: list[int], candidate_vectors: np.ndarray) -> tup filtered_records = group canonical_indices[i] = canonical_index canonical_record = groups[canonical_index][0] - for j in matches: - if j > i: - proposals[j].add(canonical_index) result.filtered.extend( DuplicateRecord(record=record, exact=best_match is None, duplicates=[(canonical_record, score)]) for record in filtered_records @@ -164,7 +157,7 @@ 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) - if (state := getattr(self, "_self_deduplication", None)) is not None: + if (state := self._self_deduplication) is not None: # Replay selection over cached group matches; filtered records must not keep each other filtered. groups, results, vectors = state result = self._from_groups(groups, results, vectors, threshold, self.columns) diff --git a/tests/test_semhash.py b/tests/test_semhash.py index 9bef786..bf3ffb2 100644 --- a/tests/test_semhash.py +++ b/tests/test_semhash.py @@ -144,7 +144,7 @@ def test_deduplicate_with_only_exact_duplicates(model: Encoder) -> None: def test_rethreshold_keeps_exact_duplicate_group(model: Encoder) -> None: - """Exact copies must not prevent their representative from being restored by rethresholding.""" + """A near-duplicate with exact copies is not listed as a duplicate of its own copies, so rethresholding keeps it.""" records = [ {"text": "It's dangerous to go alone!", "id": 1}, {"text": "It's dangerous to go alone! Take this.", "id": 2}, From bbe7633635a7616042c22945d73745cd26ffb5b1 Mon Sep 17 00:00:00 2001 From: Pringled Date: Sat, 26 Sep 2026 13:08:53 +0200 Subject: [PATCH 09/12] refactor: move canonical selection to utils and match semhash conventions --- semhash/datamodels.py | 69 +++++++------------- semhash/index.py | 5 +- semhash/semhash.py | 33 +++++----- semhash/utils.py | 47 ++++++++++++++ tests/test_duplicate_reporting.py | 101 ------------------------------ tests/test_semhash.py | 96 ++++++++++++++++++++++++++++ 6 files changed, 184 insertions(+), 167 deletions(-) delete mode 100644 tests/test_duplicate_reporting.py diff --git a/semhash/datamodels.py b/semhash/datamodels.py index 8c2dc4a..37f793b 100644 --- a/semhash/datamodels.py +++ b/semhash/datamodels.py @@ -9,8 +9,7 @@ import numpy as np from frozendict import frozendict -from semhash.index import MAX_NEIGHBORS -from semhash.utils import DuplicateList, Record, to_frozendict +from semhash.utils import DuplicateList, Neighbors, Record, select_canonicals, to_frozendict @dataclass @@ -22,7 +21,7 @@ class DuplicateRecord(Generic[Record]): ---------- record: The original record being deduplicated. exact: Whether the record was identified as an exact match. - duplicates: The canonical record and its similarity score. + duplicates: The record it duplicates and their similarity score. """ @@ -70,60 +69,34 @@ class DeduplicationResult(Generic[Record]): threshold: float = field(default=0.9) columns: Sequence[str] | None = field(default=None) - def __post_init__(self) -> None: - """Initialize the cache used for rethresholding.""" - self._self_deduplication: tuple[list[list[Record]], list[list[tuple[int, float]]], np.ndarray] | None = None + # Self-deduplication inputs, kept so rethreshold can replay the selection. + _replay_state: 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]], - results: list[list[tuple[int, float]]], + neighbors: list[Neighbors], vectors: np.ndarray, threshold: float, columns: Sequence[str] | None, ) -> DeduplicationResult[Record]: - """Assign each record to a directly matching kept canonical, preserving input group order.""" + """Build a self-deduplication result where every filtered record points to a directly matching selected record.""" result = cls(threshold=threshold, columns=columns) - norms = np.linalg.norm(vectors, axis=1) - # Zero vectors (e.g. empty text) are similar to nothing, instead of producing NaN scores. - norms[norms == 0] = 1.0 - # Normalized vectors of selected records, in selection order, for comparing against all of them at once. - selected_vectors = np.empty(vectors.shape, dtype=np.float32) - selected_indices: list[int] = [] - canonical_indices: dict[int, int] = {} - - def closest(i: int, candidates: list[int], candidate_vectors: np.ndarray) -> tuple[int, float] | None: - if not candidates: - return None - scores = candidate_vectors @ (vectors[i] / norms[i]) - best = int(np.argmax(scores)) - return (candidates[best], float(scores[best])) if scores[best] >= threshold else None - - for i, group in enumerate(groups): - matches = [j for j, score in results[i] if score >= threshold] - # Check the canonicals of matching neighbors directly, to avoid transitive matches. - candidates = list({canonical_indices[j] for j in matches if j in canonical_indices}) - best_match = closest(i, candidates, vectors[candidates] / norms[candidates, None]) - if best_match is None and len(matches) >= MAX_NEIGHBORS: - # The neighbors may be truncated, so compare against every selected record directly. - best_match = closest(i, selected_indices, selected_vectors[: len(selected_indices)]) - if best_match is None: - canonical_index, score = i, 1.0 - selected_vectors[len(selected_indices)] = vectors[i] / norms[i] - selected_indices.append(i) + 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]) - filtered_records = group[1:] - else: - canonical_index, score = best_match - filtered_records = group - canonical_indices[i] = canonical_index - canonical_record = groups[canonical_index][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=best_match is None, duplicates=[(canonical_record, score)]) + DuplicateRecord(record=record, exact=is_selected, duplicates=[(groups[canonical][0], score)]) for record in filtered_records ) - result._self_deduplication = (groups, results, vectors) + result._replay_state = (groups, neighbors, vectors) return result @property @@ -157,10 +130,12 @@ 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) - if (state := self._self_deduplication) is not None: - # Replay selection over cached group matches; filtered records must not keep each other filtered. - groups, results, vectors = state - result = self._from_groups(groups, results, vectors, threshold, self.columns) + if (state := self._replay_state) is not None: + # Replay the selection, since a record that is no longer filtered can become the canonical of later records. + groups, neighbors, vectors = state + 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 = [] diff --git a/semhash/index.py b/semhash/index.py index f079f4b..a92092b 100644 --- a/semhash/index.py +++ b/semhash/index.py @@ -7,8 +7,9 @@ from vicinity.backends import AbstractBackend, get_backend_class from vicinity.datatypes import SingleQueryResult +from semhash.utils import MAX_NEIGHBORS, Neighbors + DictItem = list[dict[str, str]] -MAX_NEIGHBORS = 100 class Index: @@ -46,7 +47,7 @@ def from_vectors_and_items( return cls(vectors, items, backend) - def query_threshold(self, vectors: np.ndarray, threshold: float) -> list[list[tuple[int, float]]]: + def query_threshold(self, vectors: np.ndarray, threshold: float) -> list[Neighbors]: """ Query the index with a threshold. diff --git a/semhash/semhash.py b/semhash/semhash.py index 95a7926..f169def 100644 --- a/semhash/semhash.py +++ b/semhash/semhash.py @@ -26,6 +26,7 @@ coerce_value, compute_candidate_limit, featurize, + normalize, to_frozendict, ) @@ -179,24 +180,19 @@ def deduplicate( duplicate_records.append(duplicate_record) # Only embed and query the records that are left after removing exact duplicates - results = [] - embeddings = np.empty((0, self.index.vectors.shape[1])) + 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, embedding, similar_items in zip(dict_records, embeddings, results): - # Rescore the neighbors with exact cosine similarity, like self_deduplicate does. - indices = [index for index, _ in similar_items] - candidates = self.index.vectors[indices] - norms = np.linalg.norm(candidates, axis=1) * np.linalg.norm(embedding) - scores = candidates @ embedding / np.where(norms == 0, 1.0, norms) - best = int(np.argmax(scores)) if indices else 0 - if not indices or scores[best] < threshold: - # No duplicates found, keep this record - deduplicated_records.append(record) - else: + for record, embedding, neighbors in zip(dict_records, embeddings, results): + # Rescore the neighbors with exact cosine similarity, like self_deduplicate does. + indices = [index for index, _ in neighbors] + scores = normalize(self.index.vectors[indices]) @ normalize(embedding) + if not 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, @@ -225,11 +221,14 @@ def self_deduplicate( :param threshold: Similarity threshold for deduplication. :return: A deduplicated list of records. """ - results = self.index.query_threshold(self.index.vectors, threshold=threshold) + neighbors = self.index.query_threshold(self.index.vectors, threshold=threshold) groups: list[list[Any]] = self.index.items if self._was_string: + # Convert before selection, so the result holds strings when rethreshold replays it. groups = [[dict_to_string(record, self.columns) for record in group] for group in groups] - return DeduplicationResult._from_groups(groups, results, self.index.vectors, threshold, self.columns) + 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..c00ad61 100644 --- a/semhash/utils.py +++ b/semhash/utils.py @@ -8,6 +8,9 @@ # Type definitions Record = TypeVar("Record", str, dict[str, Any]) DuplicateList: TypeAlias = list[tuple[Record, float]] +Neighbors: TypeAlias = list[tuple[int, float]] + +MAX_NEIGHBORS = 100 class Encoder(Protocol): @@ -147,3 +150,47 @@ 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 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. + """ + # Normalized vectors of selected groups, in selection order, for comparing against all of them at once. + selected_vectors = np.empty(vectors.shape, dtype=np.float32) + selected_indices: list[int] = [] + canonicals: list[tuple[int, float]] = [] + + def _closest(i: int, candidates: list[int], candidate_vectors: np.ndarray) -> tuple[int, float] | None: + """Return the candidate most similar to group i, if it reaches the threshold.""" + if not candidates: + return None + scores = candidate_vectors @ normalize(vectors[i]) + best = int(np.argmax(scores)) + return (candidates[best], float(scores[best])) if scores[best] >= threshold else None + + for i, group_neighbors in enumerate(neighbors): + matches = [j for j, score in group_neighbors if score >= threshold] + # Check the canonicals of earlier matching neighbors directly, to avoid transitive matches. + candidates = list({canonicals[j][0] for j in matches if j < i}) + best_match = _closest(i, candidates, normalize(vectors[candidates])) + 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, selected_indices, selected_vectors[: len(selected_indices)]) + if best_match is None: + selected_vectors[len(selected_indices)] = normalize(vectors[i]) + selected_indices.append(i) + best_match = (i, 1.0) + canonicals.append(best_match) + return canonicals diff --git a/tests/test_duplicate_reporting.py b/tests/test_duplicate_reporting.py deleted file mode 100644 index 0c09a54..0000000 --- a/tests/test_duplicate_reporting.py +++ /dev/null @@ -1,101 +0,0 @@ -import tracemalloc -from collections.abc import Sequence -from typing import Any - -import numpy as np -import pytest - -from semhash import SemHash -from semhash.utils import Encoder - - -@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 canonical, 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 len(result.filtered[0].duplicates) == 1 - assert result.filtered[0].duplicates[0][0] == 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_canonicals_and_rethreshold( - angular_model: Encoder, texts: str, threshold: float, selected: list[int], targets: list[int] -) -> None: - """Preserve complete groups with direct canonical links, including when higher thresholds split them.""" - 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.duplicates[0][0]["id"] for d in result.filtered] == targets - for cutoff in (threshold, 0.95, 0.99): - result.rethreshold(cutoff) - assert result == semhash.self_deduplicate(cutoff) - 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: - assert len(duplicate.duplicates) == 1 - canonical, score = duplicate.duplicates[0] - assert canonical in result.selected and score >= cutoff - vectors = angular_model.encode([duplicate.record["text"], canonical["text"]]) - assert score == pytest.approx(float(vectors[0] @ vectors[1]), abs=1e-6) - assert duplicate.exact is (duplicate.record["text"] == canonical["text"]) - - -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_large_exact_groups_have_linear_reporting(angular_model: Encoder) -> None: - """Self/cross results, grouping and rethresholding avoid all-pairs allocation.""" - n = 1000 - semhash = SemHash.from_records(["A"] * n, model=angular_model, ann_backend="basic") - tracemalloc.start() - try: - result = semhash.self_deduplicate() - cross = semhash.deduplicate(["A"] * n) - assert len(result.filtered) == n - 1 - assert len(cross.filtered) == n - assert all(d.duplicates == [("A", 1.0)] for d in result.filtered + cross.filtered) - assert len(result.selected_with_duplicates[0].duplicates) == n - 1 - result.rethreshold(0.99) - cross.rethreshold(0.99) - _, peak = tracemalloc.get_traced_memory() - assert peak < 8 * 1024 * 1024 - finally: - tracemalloc.stop() - - -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/tests/test_semhash.py b/tests/test_semhash.py index bf3ffb2..09c023d 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 @@ -378,3 +381,96 @@ 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 len(result.filtered[0].duplicates) == 1 + assert result.filtered[0].duplicates[0][0] == 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, and the groups are complete.""" + 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.duplicates[0][0]["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, score)] = duplicate.duplicates + vectors = angular_model.encode([duplicate.record["text"], canonical["text"]]) + assert canonical in result.selected + assert score == pytest.approx(float(vectors[0] @ vectors[1]), abs=1e-6) + assert duplicate.exact is (duplicate.record["text"] == canonical["text"]) + + +@pytest.mark.parametrize("texts,threshold", [("ABBC", 0.6), ("ABC", 0.7), ("ACB", 0.75)]) +def test_self_rethreshold_matches_fresh_run(angular_model: Encoder, texts: str, threshold: float) -> None: + """Rethresholding a self-deduplication gives the same result as deduplicating at the new threshold.""" + 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) + 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_exact_copies_report_one_canonical(angular_model: Encoder) -> None: + """Every exact copy points to one selected record in self and cross results, so reporting stays linear.""" + n = 1000 + semhash = SemHash.from_records(["A"] * n, model=angular_model, ann_backend="basic") + result = semhash.self_deduplicate() + cross = semhash.deduplicate(["A"] * n) + result.rethreshold(0.99) + cross.rethreshold(0.99) + assert len(result.filtered) == n - 1 + assert len(cross.filtered) == n + assert all(d.exact and d.duplicates == [("A", 1.0)] for d in result.filtered + cross.filtered) + assert len(result.selected_with_duplicates[0].duplicates) == n - 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 == [" "] From e9ebbfcd66641d720f8a727da18ba0ce603ed6ce Mon Sep 17 00:00:00 2001 From: Pringled Date: Sat, 26 Sep 2026 13:13:56 +0200 Subject: [PATCH 10/12] refactor: consolidate duplicate reporting tests and rename selection inputs --- semhash/datamodels.py | 12 ++++++------ semhash/semhash.py | 2 +- tests/test_semhash.py | 30 +++++------------------------- 3 files changed, 12 insertions(+), 32 deletions(-) diff --git a/semhash/datamodels.py b/semhash/datamodels.py index 37f793b..767e970 100644 --- a/semhash/datamodels.py +++ b/semhash/datamodels.py @@ -69,8 +69,8 @@ class DeduplicationResult(Generic[Record]): threshold: float = field(default=0.9) columns: Sequence[str] | None = field(default=None) - # Self-deduplication inputs, kept so rethreshold can replay the selection. - _replay_state: tuple[list[list[Record]], list[Neighbors], np.ndarray] | None = field( + # 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 ) @@ -96,7 +96,7 @@ def _from_groups( DuplicateRecord(record=record, exact=is_selected, duplicates=[(groups[canonical][0], score)]) for record in filtered_records ) - result._replay_state = (groups, neighbors, vectors) + result._selection_inputs = (groups, neighbors, vectors) return result @property @@ -130,9 +130,9 @@ 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) - if (state := self._replay_state) is not None: - # Replay the selection, since a record that is no longer filtered can become the canonical of later records. - groups, neighbors, vectors = state + 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 ) diff --git a/semhash/semhash.py b/semhash/semhash.py index f169def..6bd6a47 100644 --- a/semhash/semhash.py +++ b/semhash/semhash.py @@ -224,7 +224,7 @@ def self_deduplicate( neighbors = self.index.query_threshold(self.index.vectors, threshold=threshold) groups: list[list[Any]] = self.index.items if self._was_string: - # Convert before selection, so the result holds strings when rethreshold replays it. + # 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 diff --git a/tests/test_semhash.py b/tests/test_semhash.py index 09c023d..0ed2b16 100644 --- a/tests/test_semhash.py +++ b/tests/test_semhash.py @@ -135,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.duplicates) 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.duplicates) for d in deduplicated.filtered] == [(True, [(texts2[0], 1.0)])] * 3 def test_rethreshold_keeps_exact_duplicate_group(model: Encoder) -> None: @@ -416,7 +418,7 @@ def test_cross_dataset_reports_one_canonical(angular_model: Encoder, backend: st 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, and the groups are complete.""" + """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) @@ -430,14 +432,6 @@ def test_self_deduplication_uses_direct_canonicals( assert canonical in result.selected assert score == pytest.approx(float(vectors[0] @ vectors[1]), abs=1e-6) assert duplicate.exact is (duplicate.record["text"] == canonical["text"]) - - -@pytest.mark.parametrize("texts,threshold", [("ABBC", 0.6), ("ABC", 0.7), ("ACB", 0.75)]) -def test_self_rethreshold_matches_fresh_run(angular_model: Encoder, texts: str, threshold: float) -> None: - """Rethresholding a self-deduplication gives the same result as deduplicating at the new threshold.""" - 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) for cutoff in (0.95, 0.99): result.rethreshold(cutoff) assert result == semhash.self_deduplicate(cutoff) @@ -455,20 +449,6 @@ def test_dense_cluster_beyond_neighbor_limit(angular_model: Encoder) -> None: assert result.selected == ["0", "1"] -def test_exact_copies_report_one_canonical(angular_model: Encoder) -> None: - """Every exact copy points to one selected record in self and cross results, so reporting stays linear.""" - n = 1000 - semhash = SemHash.from_records(["A"] * n, model=angular_model, ann_backend="basic") - result = semhash.self_deduplicate() - cross = semhash.deduplicate(["A"] * n) - result.rethreshold(0.99) - cross.rethreshold(0.99) - assert len(result.filtered) == n - 1 - assert len(cross.filtered) == n - assert all(d.exact and d.duplicates == [("A", 1.0)] for d in result.filtered + cross.filtered) - assert len(result.selected_with_duplicates[0].duplicates) == n - 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) From ba41f42c061bafed717eca8599b630850fe08786 Mon Sep 17 00:00:00 2001 From: Pringled Date: Sun, 27 Sep 2026 07:32:36 +0200 Subject: [PATCH 11/12] feat: replace DuplicateRecord.duplicates with duplicate_of and score, deprecating duplicates --- semhash/datamodels.py | 38 ++++++++++++++++---------- semhash/records.py | 9 ++----- semhash/semhash.py | 7 +++-- tests/test_datamodels.py | 58 ++++++++++++++++++---------------------- tests/test_semhash.py | 13 +++++---- 5 files changed, 61 insertions(+), 64 deletions(-) diff --git a/semhash/datamodels.py b/semhash/datamodels.py index 767e970..26480cf 100644 --- a/semhash/datamodels.py +++ b/semhash/datamodels.py @@ -1,5 +1,6 @@ from __future__ import annotations +import warnings from collections import defaultdict from collections.abc import Hashable, Sequence from dataclasses import dataclass, field @@ -15,23 +16,32 @@ @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: The record it duplicates and their similarity score. + 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 @@ -93,7 +103,7 @@ def _from_groups( # 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, duplicates=[(groups[canonical][0], score)]) + DuplicateRecord(record=record, exact=is_selected, duplicate_of=groups[canonical][0], score=score) for record in filtered_records ) result._selection_inputs = (groups, neighbors, vectors) @@ -120,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] @@ -140,11 +150,10 @@ def rethreshold(self, threshold: float) -> None: else: filtered = [] for dup in self.filtered: - dup._rethreshold(threshold) - if not dup.duplicates: - self.selected.append(dup.record) - else: + if dup.score >= threshold: filtered.append(dup) + else: + self.selected.append(dup.record) self.filtered = filtered self.threshold = threshold @@ -167,8 +176,9 @@ 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: 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 6bd6a47..b7e75c2 100644 --- a/semhash/semhash.py +++ b/semhash/semhash.py @@ -13,7 +13,6 @@ 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, @@ -175,8 +174,7 @@ 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 @@ -196,7 +194,8 @@ def deduplicate( duplicate_records.append( DuplicateRecord( record=record, - duplicates=[(self.index.items[indices[best]][0], float(scores[best]))], + duplicate_of=self.index.items[indices[best]][0], + score=float(scores[best]), exact=False, ) ) diff --git a/tests/test_datamodels.py b/tests/test_datamodels.py index 2d638e9..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"], ) @@ -170,8 +164,8 @@ def test_selected_with_duplicates_preserves_occurrences() -> None: 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"], @@ -192,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, ) @@ -210,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 0ed2b16..416c071 100644 --- a/tests/test_semhash.py +++ b/tests/test_semhash.py @@ -138,14 +138,14 @@ def test_deduplicate_with_only_exact_duplicates(model: Encoder) -> None: 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.exact, d.duplicates) for d in deduplicated.filtered] == [(True, [(texts1[0], 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.exact, d.duplicates) for d in deduplicated.filtered] == [(True, [(texts2[0], 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: @@ -405,8 +405,7 @@ def test_cross_dataset_reports_one_canonical(angular_model: Encoder, backend: st query = {"id": 3, "text": "C"} result = semhash.deduplicate([query], threshold=0.6) assert result.selected == [] - assert len(result.filtered[0].duplicates) == 1 - assert result.filtered[0].duplicates[0][0] == records[1] + assert result.filtered[0].duplicate_of == records[1] result.rethreshold(0.995) assert result.selected == [query] @@ -423,14 +422,14 @@ def test_self_deduplication_uses_direct_canonicals( 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.duplicates[0][0]["id"] for d in result.filtered] == targets + 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, score)] = duplicate.duplicates + canonical = duplicate.duplicate_of vectors = angular_model.encode([duplicate.record["text"], canonical["text"]]) assert canonical in result.selected - assert score == pytest.approx(float(vectors[0] @ vectors[1]), abs=1e-6) + 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) From cfbf35bd2d186e75bdbfabd90a1a98800ed31ff8 Mon Sep 17 00:00:00 2001 From: Pringled Date: Mon, 28 Sep 2026 08:49:52 +0200 Subject: [PATCH 12/12] perf: keep neighbours as arrays and widen capped queries before brute force --- semhash/index.py | 18 +++++++++++++----- semhash/semhash.py | 5 ++--- semhash/utils.py | 30 +++++++++++++++--------------- 3 files changed, 30 insertions(+), 23 deletions(-) diff --git a/semhash/index.py b/semhash/index.py index a92092b..331ac6d 100644 --- a/semhash/index.py +++ b/semhash/index.py @@ -7,7 +7,7 @@ from vicinity.backends import AbstractBackend, get_backend_class from vicinity.datatypes import SingleQueryResult -from semhash.utils import MAX_NEIGHBORS, Neighbors +from semhash.utils import MAX_NEIGHBORS, NEIGHBORS_PER_QUERY, Neighbors DictItem = list[dict[str, str]] @@ -53,12 +53,20 @@ def query_threshold(self, vectors: np.ndarray, threshold: float) -> list[Neighbo :param vectors: The vectors to query. :param threshold: The similarity threshold. - :return: Group indices and cosine similarity scores for each query. + :return: Arrays of group indices and cosine similarity scores for each query. """ - return [ - [(int(index), 1 - distance) for index, distance in zip(*result)] - for result in self.backend.threshold(vectors, threshold=1 - threshold, max_k=MAX_NEIGHBORS) + 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/semhash.py b/semhash/semhash.py index b7e75c2..a34088a 100644 --- a/semhash/semhash.py +++ b/semhash/semhash.py @@ -182,11 +182,10 @@ def deduplicate( if dict_records: embeddings = featurize(records=dict_records, columns=self.columns, model=self.model) results = self.index.query_threshold(embeddings, threshold=threshold) - for record, embedding, neighbors in zip(dict_records, embeddings, results): + for record, embedding, (indices, _) in zip(dict_records, embeddings, results): # Rescore the neighbors with exact cosine similarity, like self_deduplicate does. - indices = [index for index, _ in neighbors] scores = normalize(self.index.vectors[indices]) @ normalize(embedding) - if not indices or scores.max() < threshold: + if not len(indices) or scores.max() < threshold: # No duplicates found, keep this record deduplicated_records.append(record) continue diff --git a/semhash/utils.py b/semhash/utils.py index c00ad61..72ddc77 100644 --- a/semhash/utils.py +++ b/semhash/utils.py @@ -8,9 +8,10 @@ # Type definitions Record = TypeVar("Record", str, dict[str, Any]) DuplicateList: TypeAlias = list[tuple[Record, float]] -Neighbors: TypeAlias = list[tuple[int, float]] +Neighbors: TypeAlias = tuple[np.ndarray, np.ndarray] -MAX_NEIGHBORS = 100 +NEIGHBORS_PER_QUERY = 100 +MAX_NEIGHBORS = 1000 class Encoder(Protocol): @@ -163,34 +164,33 @@ def select_canonicals(vectors: np.ndarray, neighbors: list[Neighbors], threshold 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 group indices and similarity scores. + :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. """ - # Normalized vectors of selected groups, in selection order, for comparing against all of them at once. - selected_vectors = np.empty(vectors.shape, dtype=np.float32) + 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: list[int], candidate_vectors: np.ndarray) -> tuple[int, float] | None: + 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 candidates: + if not len(candidates): return None - scores = candidate_vectors @ normalize(vectors[i]) + scores = unit_vectors[candidates] @ unit_vectors[i] best = int(np.argmax(scores)) - return (candidates[best], float(scores[best])) if scores[best] >= threshold else None + return (int(candidates[best]), float(scores[best])) if scores[best] >= threshold else None - for i, group_neighbors in enumerate(neighbors): - matches = [j for j, score in group_neighbors if score >= threshold] + for i, (indices, scores) in enumerate(neighbors): + matches = indices[scores >= threshold] # Check the canonicals of earlier matching neighbors directly, to avoid transitive matches. - candidates = list({canonicals[j][0] for j in matches if j < i}) - best_match = _closest(i, candidates, normalize(vectors[candidates])) + 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, selected_indices, selected_vectors[: len(selected_indices)]) + best_match = _closest(i, np.asarray(selected_indices)) if best_match is None: - selected_vectors[len(selected_indices)] = normalize(vectors[i]) selected_indices.append(i) best_match = (i, 1.0) + canonical_indices[i] = best_match[0] canonicals.append(best_match) return canonicals