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

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@ classifiers = [

dependencies = [
"model2vec>=0.3.4",
"vicinity[usearch]>=0.4.3",
"vicinity[usearch]>=0.4.6",
"frozendict",
"pyversity",
]
Expand Down Expand Up @@ -106,3 +106,4 @@ version = {attr = "semhash.version.__version__"}

[tool.uv]
exclude-newer = "1 week"
exclude-newer-package = { vicinity = false }
99 changes: 69 additions & 30 deletions semhash/datamodels.py
Original file line number Diff line number Diff line change
@@ -1,37 +1,47 @@
from __future__ import annotations

import json
import warnings
from collections import defaultdict
from collections.abc import Hashable, Sequence
from dataclasses import dataclass, field
from functools import cached_property
from typing import Generic

import numpy as np
from frozendict import frozendict

from semhash.utils import DuplicateList, Record, to_frozendict
from semhash.utils import DuplicateList, Neighbors, Record, select_canonicals, to_frozendict


@dataclass
class DuplicateRecord(Generic[Record]):
"""
A single record with its duplicates.
A record that was filtered as a duplicate of another record.

Attributes
----------
record: The original record being deduplicated.
exact: Whether the record was identified as an exact match.
duplicates: List of tuples consisting of duplicate records and their associated scores.
duplicate_of: The record that this record is a duplicate of.
score: The similarity score between record and duplicate_of.
duplicates: Deprecated, use duplicate_of and score instead.

"""

record: Record
exact: bool
duplicates: DuplicateList = field(default_factory=list)
duplicate_of: Record
score: float

def _rethreshold(self, threshold: float) -> None:
"""Rethreshold the duplicates."""
self.duplicates = [(d, score) for d, score in self.duplicates if score >= threshold]
@property
def duplicates(self) -> DuplicateList:
"""Deprecated, use duplicate_of and score instead."""
warnings.warn(
"'duplicates' is deprecated and will be removed in a future release. Use 'duplicate_of' and 'score' instead.",
DeprecationWarning,
stacklevel=2,
)
return [(self.duplicate_of, self.score)]


@dataclass
Expand Down Expand Up @@ -69,6 +79,36 @@ class DeduplicationResult(Generic[Record]):
threshold: float = field(default=0.9)
columns: Sequence[str] | None = field(default=None)

# Inputs of self_deduplicate, so rethreshold can rerun the selection at a higher threshold.
_selection_inputs: tuple[list[list[Record]], list[Neighbors], np.ndarray] | None = field(
default=None, init=False, repr=False, compare=False
)

@classmethod
def _from_groups(
cls,
groups: list[list[Record]],
neighbors: list[Neighbors],
vectors: np.ndarray,
threshold: float,
columns: Sequence[str] | None,
) -> DeduplicationResult[Record]:
"""Build a self-deduplication result where every filtered record points to a directly matching selected record."""
result = cls(threshold=threshold, columns=columns)
canonicals = select_canonicals(vectors=vectors, neighbors=neighbors, threshold=threshold)
for i, (group, (canonical, score)) in enumerate(zip(groups, canonicals)):
is_selected = canonical == i
if is_selected:
result.selected.append(group[0])
# The rest of a selected group are exact copies of its first record.
filtered_records = group[1:] if is_selected else group
result.filtered.extend(
DuplicateRecord(record=record, exact=is_selected, duplicate_of=groups[canonical][0], score=score)
for record in filtered_records
)
result._selection_inputs = (groups, neighbors, vectors)
return result

@property
def duplicate_ratio(self) -> float:
"""Return the percentage of records dropped."""
Expand All @@ -90,7 +130,7 @@ def get_least_similar_from_duplicates(self, n: int = 1) -> list[tuple[Record, Re
:param n: The number of least similar pairs to return.
:return: A list of tuples consisting of (original_record, duplicate_record, score).
"""
all_pairs = [(dup.record, d, score) for dup in self.filtered for d, score in dup.duplicates]
all_pairs = [(dup.record, dup.duplicate_of, dup.score) for dup in self.filtered]
sorted_pairs = sorted(all_pairs, key=lambda x: x[2]) # Sort by score
return sorted_pairs[:n]

Expand All @@ -100,12 +140,21 @@ def rethreshold(self, threshold: float) -> None:
raise ValueError("Threshold is smaller than the given value.")
# Invalidate cached property before modifying data
self.__dict__.pop("selected_with_duplicates", None)
# Rethreshold duplicates and move records without duplicates to selected
for dup in list(self.filtered):
dup._rethreshold(threshold)
if not dup.duplicates:
self.filtered.remove(dup)
self.selected.append(dup.record)
if (inputs := self._selection_inputs) is not None:
# Rerun the selection, since a record that is no longer filtered can become the canonical of later records.
groups, neighbors, vectors = inputs
result = self._from_groups(
groups=groups, neighbors=neighbors, vectors=vectors, threshold=threshold, columns=self.columns
)
self.selected, self.filtered = result.selected, result.filtered
else:
filtered = []
for dup in self.filtered:
if dup.score >= threshold:
filtered.append(dup)
else:
self.selected.append(dup.record)
self.filtered = filtered
self.threshold = threshold

@cached_property
Expand All @@ -127,24 +176,14 @@ def _to_hashable(record: Record) -> frozendict[str, str] | str:
# Build a mapping from original-record to [(duplicate, score), …]
buckets: defaultdict[Hashable, DuplicateList] = defaultdict(list)
for duplicate_record in self.filtered:
for original_record, score in duplicate_record.duplicates:
buckets[_to_hashable(original_record)].append((duplicate_record.record, float(score)))
buckets[_to_hashable(duplicate_record.duplicate_of)].append(
(duplicate_record.record, duplicate_record.score)
)

result: list[SelectedWithDuplicates[Record]] = []
for selected in self.selected:
# Get the list of duplicates for the selected record
raw_list = buckets.get(_to_hashable(selected), [])
# Ensure we don't have duplicates in the list
# Use full-record canonical JSON for dicts so that unhashable values are handled correctly
deduped = {
(
json.dumps(rec, sort_keys=True, separators=(",", ":"), ensure_ascii=False)
if isinstance(rec, dict)
else rec
): (rec, score)
for rec, score in raw_list
}
result.append(SelectedWithDuplicates(record=selected, duplicates=list(deduped.values())))
# Preserve occurrences, even when multiple input records have identical values.
result.append(SelectedWithDuplicates(record=selected, duplicates=buckets.get(_to_hashable(selected), [])))

return result

Expand Down
31 changes: 16 additions & 15 deletions semhash/index.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,8 +7,8 @@
from vicinity.backends import AbstractBackend, get_backend_class
from vicinity.datatypes import SingleQueryResult

DocScore = tuple[dict[str, str], float]
DocScores = list[DocScore]
from semhash.utils import MAX_NEIGHBORS, NEIGHBORS_PER_QUERY, Neighbors

DictItem = list[dict[str, str]]


Expand Down Expand Up @@ -47,25 +47,26 @@ def from_vectors_and_items(

return cls(vectors, items, backend)

def query_threshold(self, vectors: np.ndarray, threshold: float) -> list[DocScores]:
def query_threshold(self, vectors: np.ndarray, threshold: float) -> list[Neighbors]:
"""
Query the index with a threshold.

:param vectors: The vectors to query.
:param threshold: The similarity threshold.
:return: The query results.
:return: Arrays of group indices and cosine similarity scores for each query.
"""
out: list[DocScores] = []
for result in self.backend.threshold(vectors, threshold=1 - threshold, max_k=100):
intermediate = []
for index, distance in zip(*result):
# Every item in the index contains one or more records that are exact duplicates of each other.
# Only the first is returned, since listing every copy grows with the size of the group.
# The score is the cosine similarity. The backend returns distances, so we need to convert.
intermediate.append((self.items[index][0], 1 - distance))
out.append(intermediate)

return out
neighbors = [
(indices, 1 - distances)
for indices, distances in self.backend.threshold(
vectors, threshold=1 - threshold, max_k=NEIGHBORS_PER_QUERY
)
]
# Query rows that hit the limit again with a larger one, so dense clusters rarely need a brute-force fallback.
if capped := [i for i, (indices, _) in enumerate(neighbors) if len(indices) >= NEIGHBORS_PER_QUERY]:
results = self.backend.threshold(vectors[capped], threshold=1 - threshold, max_k=MAX_NEIGHBORS)
for i, (indices, distances) in zip(capped, results):
neighbors[i] = (indices, 1 - distances)
return neighbors

def query_top_k(self, vectors: np.ndarray, k: int, vectors_are_in_index: bool) -> list[SingleQueryResult]:
"""
Expand Down
9 changes: 2 additions & 7 deletions semhash/records.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
80 changes: 21 additions & 59 deletions semhash/semhash.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@
from semhash.datamodels import DeduplicationResult, DuplicateRecord, FilterResult
from semhash.index import Index
from semhash.records import (
add_scores_to_records,
dict_to_string,
group_records_by_key,
map_deduplication_result_to_strings,
prepare_records,
Expand All @@ -25,6 +25,7 @@
coerce_value,
compute_candidate_limit,
featurize,
normalize,
to_frozendict,
)

Expand Down Expand Up @@ -173,26 +174,27 @@ def deduplicate(
)
duplicate_records = []
for record, duplicates in exact_duplicates:
duplicated_with_score = add_scores_to_records(duplicates)
duplicate_record = DuplicateRecord(record=record, duplicates=duplicated_with_score, exact=True)
duplicate_record = DuplicateRecord(record=record, duplicate_of=duplicates[0], score=1.0, exact=True)
duplicate_records.append(duplicate_record)

# Only embed and query the records that are left after removing exact duplicates
results = []
deduplicated_records = []
if dict_records:
embeddings = featurize(records=dict_records, columns=self.columns, model=self.model)
results = self.index.query_threshold(embeddings, threshold=threshold)

deduplicated_records = []
for record, similar_items in zip(dict_records, results):
if not similar_items:
# No duplicates found, keep this record
deduplicated_records.append(record)
else:
for record, embedding, (indices, _) in zip(dict_records, embeddings, results):
# Rescore the neighbors with exact cosine similarity, like self_deduplicate does.
scores = normalize(self.index.vectors[indices]) @ normalize(embedding)
if not len(indices) or scores.max() < threshold:
# No duplicates found, keep this record
deduplicated_records.append(record)
continue
best = int(np.argmax(scores))
duplicate_records.append(
DuplicateRecord(
record=record,
duplicates=[(item, score) for item, score in similar_items],
duplicate_of=self.index.items[indices[best]][0],
score=float(scores[best]),
exact=False,
)
)
Expand All @@ -217,54 +219,14 @@ def self_deduplicate(
:param threshold: Similarity threshold for deduplication.
:return: A deduplicated list of records.
"""
# Query the fitted index
results = self.index.query_threshold(self.index.vectors, threshold=threshold)
column_set = set(self.columns)

duplicate_records = []

deduplicated_records = []
seen_items: set[frozendict[str, str]] = set()
for item, similar_items in zip(self.index.items, results):
# Items is a list of items which are exact duplicates of each other.
# The first record is kept, and every other copy is an exact duplicate of it. Each copy only lists
# the kept record, since listing every other copy grows quadratically with the size of the group.
record, *duplicates = item
for curr_record in duplicates:
duplicate_records.append(DuplicateRecord(record=curr_record, duplicates=[(record, 1.0)], exact=True))

# If we don't see any similar_items, we know the record is not a duplicate.
# In rare cases, the item itself might not be returned by the index.
if not similar_items: # pragma: no cover
deduplicated_records.append(record)
continue
items, _ = zip(*similar_items)
frozen_items = [to_frozendict(item, column_set) for item in items]
# similar_items includes 'record' itself
# If we've seen any of these items before, this is a duplicate cluster.
if any(item in seen_items for item in frozen_items):
duplicate_records.append(
DuplicateRecord(
record=record,
duplicates=[(item, score) for item, score in similar_items if item != record],
exact=False,
)
)
continue
# This is the first time we see this cluster of similar items
deduplicated_records.append(record)
# Mark all items in this cluster as seen
seen_items.update(frozen_items)

result = DeduplicationResult(
selected=deduplicated_records, filtered=duplicate_records, threshold=threshold, columns=self.columns
)

neighbors = self.index.query_threshold(self.index.vectors, threshold=threshold)
groups: list[list[Any]] = self.index.items
if self._was_string:
# Convert records back to strings if the records were originally strings
return map_deduplication_result_to_strings(result, columns=self.columns)

return result
# Convert before selection, so the result holds strings when rethreshold reruns it.
groups = [[dict_to_string(record, self.columns) for record in group] for group in groups]
return DeduplicationResult._from_groups(
groups=groups, neighbors=neighbors, vectors=self.index.vectors, threshold=threshold, columns=self.columns
)

def _validate_if_strings(self, records: Sequence[dict[str, Any] | str]) -> list[dict[str, Any]]:
"""
Expand Down
Loading
Loading