File size: 19,417 Bytes
58258b8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
"""Edge cases and limits for MinHash + Semantic deduplicators.

Covers:
  - Degenerate inputs (empty strings, whitespace, single char, unicode, emoji)
  - Threshold boundaries (1.0, near-zero, just-above-cutoff)
  - Shingle/permutation extremes
  - Large clusters, many disjoint clusters, mixed
  - HNSW vs Flat boundary behavior
  - Numerical edge cases for embeddings (zero vector, antipodal, very high dim)
"""

from __future__ import annotations

import pytest

ds = pytest.importorskip("datasketch")
faiss = pytest.importorskip("faiss")
import numpy as np

from sdg.preprocessing.dedupe.minhash import MinHashDeduplicator
from sdg.preprocessing.dedupe.semantic import SemanticDeduplicator


# ────────────────────────────────────────────────────────────────────────────
# MinHash: degenerate inputs
# ────────────────────────────────────────────────────────────────────────────

def test_minhash_all_empty_strings_collapse():
    keep = MinHashDeduplicator().dedup(["", "", ""], show_progress=False)
    assert len(keep) == 1


def test_minhash_one_empty_with_distinct_others():
    keep = MinHashDeduplicator().dedup(
        ["", "alpha beta gamma delta epsilon zeta", "completely different content here"],
        show_progress=False,
    )
    # Empty doesn't share shingles with the others -> kept as singleton
    assert len(keep) == 3


def test_minhash_whitespace_only_strings_collapse_with_normalize():
    keep = MinHashDeduplicator(normalize=True).dedup(
        ["   ", "\t\t", "\n\n", " \t \n "],
        show_progress=False,
    )
    assert len(keep) == 1


def test_minhash_single_character_strings():
    """Each char becomes its own short shingle; identical chars collapse."""
    keep = MinHashDeduplicator().dedup(["a", "a", "b"], show_progress=False)
    assert len(keep) == 2


def test_minhash_string_shorter_than_shingle_size():
    """When text is shorter than shingle size, the whole text is a single shingle."""
    keep = MinHashDeduplicator(shingle_size=10).dedup(
        ["short text", "short text", "very different brief"],
        show_progress=False,
    )
    assert len(keep) == 2


def test_minhash_unicode_cjk():
    text_a = "解决方程 x εΉ³ζ–Ή ε‡εŽ» ε›› η­‰δΊŽ ι›Ά ηš„ 详细 ζ­₯ιͺ€ θ§£ι‡Š"
    text_b = text_a  # identical
    text_c = "ε…‰εˆδ½œη”¨ 是 怍物 εˆ©η”¨ 光能 εˆΆι€  ζœ‰ζœΊη‰© ηš„ 过程 θ―¦θ§£"
    keep = MinHashDeduplicator().dedup([text_a, text_b, text_c], show_progress=False)
    assert len(keep) == 2


def test_minhash_emoji_inputs():
    a = "Solve this 🎯 equation x squared minus four equals zero step by step"
    b = "Solve this 🎯 equation x squared minus four equals zero step by step"
    c = "Cook the 🍝 pasta in salted boiling water for nine to ten minutes"
    keep = MinHashDeduplicator().dedup([a, b, c], show_progress=False)
    assert len(keep) == 2


def test_minhash_repeated_token_pattern():
    """Many copies of the same token -> all shingles are identical."""
    a = ("yes " * 50).strip()
    b = ("yes " * 100).strip()
    c = ("no " * 50).strip()
    keep = MinHashDeduplicator(threshold=0.5).dedup([a, b, c], show_progress=False)
    # a and b share their only shingle "yes yes yes yes yes" -> Jaccard 1.0
    # c is disjoint from them
    assert len(keep) == 2


# ────────────────────────────────────────────────────────────────────────────
# MinHash: threshold limits
# ────────────────────────────────────────────────────────────────────────────

def test_minhash_threshold_1_rejected_by_validator():
    """datasketch's LSH cannot find valid bands at threshold=1.0; we reject early."""
    with pytest.raises(ValueError, match="threshold"):
        MinHashDeduplicator(threshold=1.0)


def test_minhash_high_threshold_keeps_near_dups_apart():
    """At a strict threshold (0.95), a clearly-below pair (Jaccard ~0.4) does not collapse."""
    base = " ".join(f"word{i}" for i in range(20))  # 20 tokens, 16 5-grams
    a = base
    # Append 20 unrelated tokens -> b has 36 5-grams, 16 shared with a, Jaccard 16/36 β‰ˆ 0.44
    b = base + " " + " ".join(f"extra{i}" for i in range(20))
    keep = MinHashDeduplicator(threshold=0.95).dedup([a, b], show_progress=False)
    assert len(keep) == 2


def test_minhash_high_threshold_collapses_exact_duplicates():
    """At a strict threshold (0.95), byte-identical strings still collapse."""
    a = " ".join(f"word{i}" for i in range(20))
    keep = MinHashDeduplicator(threshold=0.95).dedup([a, a, a], show_progress=False)
    assert len(keep) == 1


def test_minhash_low_threshold_aggressive_collapse():
    """At very low threshold, weak similarity is enough to merge."""
    texts = [
        "alpha beta gamma delta epsilon zeta eta theta iota kappa lambda mu",
        "alpha xxxxx yyyyy zzzzz wwwww vvvvv uuuuu ttttt sssss rrrrr qqqqq pppp",  # shares only 'alpha' as 1-gram
    ]
    # threshold 0.05 + 1-word shingles -> very low bar; "alpha" alone may suffice
    keep = MinHashDeduplicator(threshold=0.05, shingle_size=1).dedup(texts, show_progress=False)
    # Tolerant: result is 1 (collapsed) or 2 (LSH missed). Either is acceptable
    # at extreme threshold; just don't crash.
    assert len(keep) in {1, 2}


def test_minhash_invalid_threshold_zero_raises():
    with pytest.raises(ValueError):
        MinHashDeduplicator(threshold=0.0)


def test_minhash_invalid_threshold_negative_raises():
    with pytest.raises(ValueError):
        MinHashDeduplicator(threshold=-0.5)


# ────────────────────────────────────────────────────────────────────────────
# MinHash: shingle and permutation extremes
# ────────────────────────────────────────────────────────────────────────────

def test_minhash_shingle_size_one():
    """1-word shingles == bag-of-words Jaccard."""
    a = "the quick brown fox"
    b = "fox the brown quick"  # same words, different order
    keep = MinHashDeduplicator(shingle_size=1, threshold=0.9).dedup([a, b], show_progress=False)
    # Same shingle set -> Jaccard 1.0 -> collapse
    assert len(keep) == 1


def test_minhash_shingle_size_one_word_order_matters_with_5gram():
    a = "the quick brown fox jumps over the lazy dog quickly today"
    b = "today quickly dog lazy the over jumps fox brown quick the"
    keep = MinHashDeduplicator(shingle_size=5, threshold=0.7).dedup([a, b], show_progress=False)
    # Word order changes -> different 5-grams -> Jaccard near 0 -> kept apart
    assert len(keep) == 2


def test_minhash_invalid_shingle_size_zero_raises():
    with pytest.raises(ValueError):
        MinHashDeduplicator(shingle_size=0)


def test_minhash_low_num_perm_still_works():
    """Low num_perm == noisier similarity estimates but still functional."""
    a = "alpha beta gamma delta epsilon zeta eta theta iota kappa lambda"
    keep = MinHashDeduplicator(num_perm=16, threshold=0.7).dedup([a, a], show_progress=False)
    assert len(keep) == 1


def test_minhash_high_num_perm_still_works():
    a = "alpha beta gamma delta epsilon zeta eta theta iota kappa lambda mu"
    keep = MinHashDeduplicator(num_perm=512, threshold=0.7).dedup([a, a], show_progress=False)
    assert len(keep) == 1


# ────────────────────────────────────────────────────────────────────────────
# MinHash: cluster topology
# ────────────────────────────────────────────────────────────────────────────

def test_minhash_large_identical_cluster_collapses_to_one():
    a = "Solve x squared minus four equals zero showing all your work step by step in detail"
    keep = MinHashDeduplicator().dedup([a] * 200, show_progress=False)
    assert len(keep) == 1


def test_minhash_many_disjoint_pairs_all_collapse():
    """50 pairs of identical strings -> 50 kept."""
    pairs = []
    for i in range(50):
        # each pair shares its own unique vocabulary, no cross-pair overlap
        prompt = f"unique prompt vocabulary words alpha{i} beta{i} gamma{i} delta{i} epsilon{i} zeta{i}"
        pairs.append(prompt)
        pairs.append(prompt)
    keep = MinHashDeduplicator().dedup(pairs, show_progress=False)
    assert len(keep) == 50


def test_minhash_chain_of_similar_strings_clusters():
    """Three near-identical strings (each one-word edit from the base) form one cluster.

    Pure transitivity (A-B above, B-C above, A-C below) is hard to construct
    reliably under MinHash's probabilistic LSH at moderate num_perm. Union-find
    behavior on transitivity is independently tested in test_base.py; here we
    just verify cluster formation works for chained near-dups.
    """
    base = " ".join(f"word{i}" for i in range(40))  # 40 tokens
    a = base
    b = a.replace("word0 ", "REPLACED_B ", 1)
    c = a.replace("word39", "REPLACED_C", 1)
    # Pairwise Jaccard ~0.95+, well above threshold.
    keep = MinHashDeduplicator(threshold=0.85).dedup([a, b, c], show_progress=False)
    assert len(keep) == 1


# ────────────────────────────────────────────────────────────────────────────
# MinHash: normalization corner cases
# ────────────────────────────────────────────────────────────────────────────

def test_minhash_normalize_handles_extreme_whitespace():
    a = "alpha beta gamma delta epsilon zeta eta theta iota kappa"
    b = "ALPHA\tBETA  GAMMA\n\nDELTA\rEPSILON  ZETA   ETA\tTHETA IOTA KAPPA"
    keep = MinHashDeduplicator(normalize=True, threshold=0.9).dedup([a, b], show_progress=False)
    assert len(keep) == 1


def test_minhash_normalize_off_keeps_case_variants_apart():
    a = "ALPHA BETA GAMMA DELTA EPSILON ZETA ETA THETA IOTA KAPPA"
    b = "alpha beta gamma delta epsilon zeta eta theta iota kappa"
    keep = MinHashDeduplicator(normalize=False, threshold=0.9).dedup([a, b], show_progress=False)
    # Different shingles entirely (case-sensitive) -> Jaccard 0 -> both kept
    assert len(keep) == 2


# ────────────────────────────────────────────────────────────────────────────
# Semantic: degenerate vectors
# ────────────────────────────────────────────────────────────────────────────

def _normalize(v: np.ndarray) -> np.ndarray:
    norms = np.linalg.norm(v, axis=1, keepdims=True)
    return (v / np.clip(norms, 1e-12, None)).astype(np.float32)


def test_semantic_zero_vector_does_not_crash():
    """A zero vector has undefined direction; clipping keeps it numerically safe."""
    emb = np.zeros((3, 8), dtype=np.float32)
    # All-zero inner product is 0 -> below any reasonable threshold -> all kept
    keep = SemanticDeduplicator(threshold=0.5).dedup_from_embeddings(emb)
    assert len(keep) == 3


def test_semantic_antipodal_vectors_not_clustered():
    """Cosine = -1 must not be treated as similar."""
    a = np.array([[1.0, 0.0, 0.0]], dtype=np.float32)
    b = np.array([[-1.0, 0.0, 0.0]], dtype=np.float32)
    emb = np.vstack([a, b])
    keep = SemanticDeduplicator(threshold=0.5).dedup_from_embeddings(emb)
    assert keep == [0, 1]


def test_semantic_one_dim_embeddings():
    emb = np.array([[1.0], [1.0], [-1.0]], dtype=np.float32)
    keep = SemanticDeduplicator(threshold=0.99).dedup_from_embeddings(emb)
    # Two +1 collapse, -1 stays separate
    assert len(keep) == 2


def test_semantic_high_dim_1024():
    rng = np.random.default_rng(0)
    emb = _normalize(rng.standard_normal((100, 1024)).astype(np.float32))
    keep = SemanticDeduplicator(threshold=0.95).dedup_from_embeddings(emb)
    assert len(keep) == 100  # nearly orthogonal in 1024D


# ────────────────────────────────────────────────────────────────────────────
# Semantic: topk extremes
# ────────────────────────────────────────────────────────────────────────────

def test_semantic_topk_one_misses_cross_matches():
    """topk=1 returns each vector's nearest neighbor = self. j>i is never
    satisfied, so even highly similar vectors are not merged.

    Uses slightly distinct (cosβ‰ˆ0.996) vectors to avoid the tied-IP nondeterminism
    you'd get from two byte-identical embeddings.
    """
    emb = _normalize(
        np.array(
            [[1.0, 0.0], [np.cos(np.radians(5)), np.sin(np.radians(5))]],
            dtype=np.float32,
        )
    )
    keep = SemanticDeduplicator(threshold=0.95, topk=1).dedup_from_embeddings(emb)
    assert len(keep) == 2


def test_semantic_topk_equal_to_n_full_pairwise():
    rng = np.random.default_rng(1)
    base = rng.standard_normal((1, 16)).astype(np.float32)
    emb = _normalize(np.vstack([base, base, base]))  # all identical
    keep = SemanticDeduplicator(threshold=0.99, topk=3).dedup_from_embeddings(emb)
    assert len(keep) == 1


def test_semantic_topk_larger_than_n_clamped_safely():
    emb = _normalize(np.array([[1.0, 0.0], [1.0, 0.0]], dtype=np.float32))
    # topk=100 > n=2 -> internal min() prevents overrun
    keep = SemanticDeduplicator(threshold=0.99, topk=100).dedup_from_embeddings(emb)
    assert len(keep) == 1


# ────────────────────────────────────────────────────────────────────────────
# Semantic: threshold limits
# ────────────────────────────────────────────────────────────────────────────

def test_semantic_threshold_1_only_exactly_identical():
    a = _normalize(np.array([[1.0, 0.0]], dtype=np.float32))
    b = _normalize(np.array([[np.cos(np.radians(2)), np.sin(np.radians(2))]], dtype=np.float32))
    emb = np.vstack([a, a, b])
    keep = SemanticDeduplicator(threshold=1.0).dedup_from_embeddings(emb)
    # a == a collapses (cos == 1), b stays apart (cos β‰ˆ 0.999 < 1.0)
    assert len(keep) == 2


def test_semantic_low_threshold_collapses_many():
    """At cos>=0.5, mildly related vectors collapse together."""
    angles = [0, 10, 20, 30, 40, 50]  # all within 60 degrees
    emb = _normalize(
        np.array(
            [[np.cos(np.radians(a)), np.sin(np.radians(a))] for a in angles],
            dtype=np.float32,
        )
    )
    keep = SemanticDeduplicator(threshold=0.5, topk=6).dedup_from_embeddings(emb)
    # All pairwise cos > 0.5 -> all merge transitively
    assert len(keep) == 1


# ────────────────────────────────────────────────────────────────────────────
# Semantic: HNSW vs Flat boundary
# ────────────────────────────────────────────────────────────────────────────

def test_semantic_just_below_hnsw_threshold_uses_flat():
    """Boundary check: n < hnsw_threshold -> Flat path; just verify no crash + correctness."""
    rng = np.random.default_rng(7)
    emb = _normalize(rng.standard_normal((49, 64)).astype(np.float32))
    dedup = SemanticDeduplicator(threshold=0.99, hnsw_threshold=50)
    keep = dedup.dedup_from_embeddings(emb)
    assert len(keep) == 49


def test_semantic_just_above_hnsw_threshold_uses_hnsw():
    rng = np.random.default_rng(7)
    emb = _normalize(rng.standard_normal((51, 64)).astype(np.float32))
    dedup = SemanticDeduplicator(threshold=0.99, hnsw_threshold=50)
    keep = dedup.dedup_from_embeddings(emb)
    assert len(keep) == 51


def test_semantic_hnsw_path_finds_planted_cluster_among_distractors():
    """HNSW must still detect a planted duplicate cluster despite approximation."""
    rng = np.random.default_rng(99)
    distractors = rng.standard_normal((300, 128)).astype(np.float32)
    seed_vec = rng.standard_normal((1, 128)).astype(np.float32)
    planted = np.repeat(seed_vec, 30, axis=0)
    raw = np.concatenate([distractors, planted], axis=0)
    emb = _normalize(raw)
    dedup = SemanticDeduplicator(threshold=0.95, hnsw_threshold=50, topk=40)
    keep = dedup.dedup_from_embeddings(emb)
    # 300 distractors + 1 cluster representative
    assert len(keep) == 301


# ────────────────────────────────────────────────────────────────────────────
# Semantic: representative selection across many clusters
# ────────────────────────────────────────────────────────────────────────────

def test_semantic_key_fn_picks_max_across_many_clusters():
    """5 clusters of 3 identical vectors each; key_fn picks the longest in each."""
    rng = np.random.default_rng(11)
    seeds = rng.standard_normal((5, 32)).astype(np.float32)
    raw = np.repeat(seeds, 3, axis=0)  # 15 vectors, 5 distinct directions x 3
    emb = _normalize(raw)

    # response_lengths designed so that within each cluster of 3 (idx triples
    # 0-2, 3-5, ...), the LAST index has the largest length.
    response_lengths = [10, 20, 30, 40, 50, 60, 70, 80, 90, 100, 110, 120, 130, 140, 150]
    key_fn = lambda i: -response_lengths[i]
    dedup = SemanticDeduplicator(threshold=0.99, topk=5)
    keep = dedup.dedup_from_embeddings(emb, key_fn=key_fn)
    assert keep == [2, 5, 8, 11, 14]