File size: 8,817 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
"""Large-text and large-N correctness tests.

Verifies the deduplicators handle realistic-scale inputs without false
positives (over-merging distinct items) or false negatives (missing planted
duplicates).

Sizes are kept modest enough to run inside pytest in <30s on M-series.
For pure throughput numbers see test_benchmarks.py.
"""

from __future__ import annotations

import random

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


# ────────────────────────────────────────────────────────────────────────────
# Synthetic prompt generators
# ────────────────────────────────────────────────────────────────────────────

_VOCAB = (
    "alpha beta gamma delta epsilon zeta eta theta iota kappa lambda mu nu xi "
    "omicron pi rho sigma tau upsilon phi chi psi omega the quick brown fox "
    "jumps over lazy dog cat house tree mountain river ocean sun moon star "
    "planet galaxy universe science math physics chemistry biology equation "
    "matrix vector function derivative integral hypothesis theorem proof "
    "experiment observation analysis synthesis conclusion problem solution "
    "method approach strategy algorithm structure pattern model framework"
).split()


def _distinct_prompts(n: int, words_per_prompt: int = 25, seed: int = 42) -> list[str]:
    """Generate n distinct random-word prompts (very low cross-similarity)."""
    rng = random.Random(seed)
    return [
        " ".join(rng.sample(_VOCAB, words_per_prompt)) + f" uniqueid{i}"
        for i in range(n)
    ]


def _planted_dup_prompts(
    n_distinct: int, dup_factor: int, words_per_prompt: int = 25, seed: int = 42
) -> tuple[list[str], int]:
    """Generate n_distinct base prompts, each duplicated dup_factor times.

    Returns (texts, expected_unique). Total len(texts) == n_distinct * dup_factor.
    """
    base = _distinct_prompts(n_distinct, words_per_prompt, seed)
    texts: list[str] = []
    for prompt in base:
        for _ in range(dup_factor):
            texts.append(prompt)
    rng = random.Random(seed + 1)
    rng.shuffle(texts)
    return texts, n_distinct


# ────────────────────────────────────────────────────────────────────────────
# MinHash: large-N correctness
# ────────────────────────────────────────────────────────────────────────────

def test_minhash_5k_distinct_no_false_positives():
    texts = _distinct_prompts(5_000)
    keep = MinHashDeduplicator(threshold=0.8).dedup(texts, show_progress=False)
    # Up to a tiny LSH false-positive rate; expect >= 99% kept.
    assert len(keep) >= 4_950, f"Too many false-positive merges: kept {len(keep)}"


def test_minhash_5k_with_planted_clusters_collapses_correctly():
    texts, n_unique = _planted_dup_prompts(n_distinct=1_000, dup_factor=5)
    keep = MinHashDeduplicator(threshold=0.8).dedup(texts, show_progress=False)
    # Allow tiny LSH miss tolerance (rare) -> within +/- 1% of n_unique
    assert n_unique - 10 <= len(keep) <= n_unique + 50


def test_minhash_long_string_100k_chars_does_not_crash():
    """A single >100K-char prompt should be processed without error."""
    rng = random.Random(42)
    # 20K random tokens => ~120K chars, with many distinct 5-grams
    tokens = [rng.choice(_VOCAB) for _ in range(20_000)]
    long_a = " ".join(tokens)
    long_b = long_a + " trailing differentiator words here"
    keep = MinHashDeduplicator(threshold=0.85).dedup([long_a, long_b], show_progress=False)
    # 4 extra tokens out of 20K -> Jaccard ~1.0 -> should collapse
    assert len(keep) == 1


def test_minhash_mixed_short_and_long_prompts():
    short = ["short prompt one alpha", "short prompt two beta", "short prompt three gamma"]
    longs = _distinct_prompts(10, words_per_prompt=40, seed=99)
    texts = short + longs + short  # short prompts duplicated, longs distinct
    keep = MinHashDeduplicator(threshold=0.8).dedup(texts, show_progress=False)
    # 3 unique short + 10 unique long = 13
    assert 12 <= len(keep) <= 14


# ────────────────────────────────────────────────────────────────────────────
# Semantic: large-N correctness (using random embeddings to avoid model load)
# ────────────────────────────────────────────────────────────────────────────

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_5k_random_embeddings_no_false_positives():
    rng = np.random.default_rng(42)
    emb = _normalize(rng.standard_normal((5_000, 384)).astype(np.float32))
    dedup = SemanticDeduplicator(threshold=0.95, hnsw_threshold=1_000)
    keep = dedup.dedup_from_embeddings(emb)
    # 384-dim random gaussians at 0.95 threshold should produce nearly no merges.
    assert len(keep) >= 4_995


def test_semantic_10k_with_planted_clusters_finds_them():
    rng = np.random.default_rng(123)
    distractors = rng.standard_normal((9_000, 384)).astype(np.float32)
    # 100 planted clusters of 10 each
    seeds = rng.standard_normal((100, 384)).astype(np.float32)
    planted = np.repeat(seeds, 10, axis=0)
    raw = np.concatenate([distractors, planted], axis=0)
    emb = _normalize(raw)
    dedup = SemanticDeduplicator(threshold=0.95, hnsw_threshold=1_000, topk=15)
    keep = dedup.dedup_from_embeddings(emb)
    # 9_000 distractors + 100 cluster reps = 9_100; allow small approximation tolerance
    assert 9_080 <= len(keep) <= 9_120


def test_semantic_one_giant_cluster_plus_singletons():
    rng = np.random.default_rng(7)
    distractors = rng.standard_normal((500, 128)).astype(np.float32)
    seed_vec = rng.standard_normal((1, 128)).astype(np.float32)
    giant = np.repeat(seed_vec, 500, axis=0)  # 500 identical
    raw = np.concatenate([distractors, giant], axis=0)
    emb = _normalize(raw)
    dedup = SemanticDeduplicator(threshold=0.95, hnsw_threshold=100, topk=600)
    keep = dedup.dedup_from_embeddings(emb)
    # 500 distractors + 1 representative
    assert len(keep) == 501


# ────────────────────────────────────────────────────────────────────────────
# End-to-end pipeline: MinHash -> Semantic stack on synthetic data
# ────────────────────────────────────────────────────────────────────────────

def test_pipeline_minhash_then_semantic_on_synthetic_records():
    """Two-stage pipeline behaves like SDG preprocess script.

    Build records, run MinHash on prompts, then semantic dedup on the survivors
    (using random embeddings as a stand-in for real encoding). Verify the
    record list shrinks monotonically and final count is plausible.
    """
    texts, _ = _planted_dup_prompts(n_distinct=500, dup_factor=3)
    records = [{"prompt": t, "response": f"resp for {i}"} for i, t in enumerate(texts)]

    n_initial = len(records)
    prompts = [r["prompt"] for r in records]

    keep_mh = MinHashDeduplicator(threshold=0.8).dedup(prompts, show_progress=False)
    records = [records[i] for i in keep_mh]
    assert len(records) < n_initial
    assert 490 <= len(records) <= 530  # ~500 unique

    # Use random embeddings sized to surviving records
    rng = np.random.default_rng(0)
    emb = _normalize(rng.standard_normal((len(records), 384)).astype(np.float32))
    dedup = SemanticDeduplicator(threshold=0.95)
    keep_sem = dedup.dedup_from_embeddings(emb)
    records = [records[i] for i in keep_sem]
    # Random embeddings are nearly orthogonal -> semantic stage shouldn't merge much.
    assert len(records) >= 480