File size: 16,214 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
"""Invariants, end-to-end pipeline simulation, realistic data, and snapshots.

Sections:
  1. MinHash invariants β€” properties dedup must always satisfy
  2. Semantic invariants β€” same, for SemanticDeduplicator
  3. End-to-end preprocess-script pipeline simulation
  4. Realistic-data tests using fixtures.REALISTIC_SCIENCE_PROMPTS
  5. Backward-compat snapshots β€” pinned outputs to detect default changes
"""

from __future__ import annotations

import json
import random

import pytest

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

from sdg.preprocessing.dedupe.fixtures import (
    EXACT_DUP_RANGE,
    FORMATTING_RANGE,
    NEAR_DUP_RANGE,
    REALISTIC_SCIENCE_PROMPTS,
    SINGLETON_RANGE,
)
from sdg.preprocessing.dedupe.minhash import MinHashDeduplicator
from sdg.preprocessing.dedupe.semantic import SemanticDeduplicator


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 _make_records(prompts: list[str]) -> list[dict]:
    """Build records that mirror the preprocess script's schema."""
    return [
        {
            "conversations": [
                {"from": "human", "value": p},
                {"from": "gpt", "value": "x" * (50 + i * 10)},
            ],
            "source": "test",
        }
        for i, p in enumerate(prompts)
    ]


# ════════════════════════════════════════════════════════════════════════════
# Section 1: MinHash invariants
# ════════════════════════════════════════════════════════════════════════════

def _sample_minhash_inputs() -> list[str]:
    rng = random.Random(7)
    base_vocab = "alpha beta gamma delta epsilon zeta eta theta iota kappa lambda mu nu xi pi rho sigma tau".split()
    distinct = [" ".join(rng.sample(base_vocab, 12)) + f" id{i}" for i in range(40)]
    duplicates = distinct[:5] * 3  # 15 dup copies of the first 5
    return distinct + duplicates


def test_minhash_keep_indices_are_sorted_unique_in_range():
    inputs = _sample_minhash_inputs()
    n = len(inputs)
    keep = MinHashDeduplicator(threshold=0.8).dedup(inputs, show_progress=False)
    assert keep == sorted(keep)
    assert len(set(keep)) == len(keep)
    assert all(0 <= i < n for i in keep)
    assert len(keep) <= n


def test_minhash_keep_indices_are_python_ints_json_serializable():
    inputs = _sample_minhash_inputs()
    keep = MinHashDeduplicator(threshold=0.8).dedup(inputs, show_progress=False)
    assert all(type(i) is int for i in keep)
    json.dumps(keep)  # must not raise


def test_minhash_idempotent_on_survivors():
    """Running dedup on the survivors of a previous dedup is a no-op."""
    inputs = _sample_minhash_inputs()
    dedup = MinHashDeduplicator(threshold=0.8)
    keep1 = dedup.dedup(inputs, show_progress=False)
    survivors = [inputs[i] for i in keep1]
    keep2 = dedup.dedup(survivors, show_progress=False)
    assert keep2 == list(range(len(survivors)))


def test_minhash_permutation_invariance_keep_count():
    """Shuffling input doesn't change the number of survivors."""
    inputs = _sample_minhash_inputs()
    dedup = MinHashDeduplicator(threshold=0.8)
    keep_a = dedup.dedup(inputs, show_progress=False)

    rng = random.Random(99)
    perm = list(range(len(inputs)))
    rng.shuffle(perm)
    shuffled = [inputs[p] for p in perm]
    keep_b = dedup.dedup(shuffled, show_progress=False)

    assert len(keep_a) == len(keep_b)


def test_minhash_appending_duplicates_does_not_increase_kept_count():
    """Adding more byte-identical copies of existing items can only keep len(keep) the same."""
    inputs = _sample_minhash_inputs()
    dedup = MinHashDeduplicator(threshold=0.8)
    keep_before = dedup.dedup(inputs, show_progress=False)
    kept_texts_before = {inputs[i] for i in keep_before}

    # Append more dups of items that already exist
    augmented = inputs + inputs[:10] + inputs[:10]
    keep_after = dedup.dedup(augmented, show_progress=False)
    assert len(keep_after) == len(keep_before)
    # The set of kept *texts* should be identical (representatives may have shifted by index)
    kept_texts_after = {augmented[i] for i in keep_after}
    assert kept_texts_after == kept_texts_before


# ════════════════════════════════════════════════════════════════════════════
# Section 2: Semantic invariants
# ════════════════════════════════════════════════════════════════════════════

def _sample_embeddings(n: int = 50, dim: int = 64, seed: int = 7) -> np.ndarray:
    rng = np.random.default_rng(seed)
    distinct = rng.standard_normal((n, dim)).astype(np.float32)
    # Plant 5 dup-clusters of 4 each
    seeds = rng.standard_normal((5, dim)).astype(np.float32)
    planted = np.repeat(seeds, 4, axis=0)
    return _normalize(np.concatenate([distinct, planted], axis=0))


def test_semantic_keep_indices_are_sorted_unique_in_range():
    emb = _sample_embeddings()
    n = len(emb)
    keep = SemanticDeduplicator(threshold=0.95).dedup_from_embeddings(emb)
    assert keep == sorted(keep)
    assert len(set(keep)) == len(keep)
    assert all(0 <= i < n for i in keep)
    assert len(keep) <= n


def test_semantic_keep_indices_are_python_ints_json_serializable():
    emb = _sample_embeddings()
    keep = SemanticDeduplicator(threshold=0.95).dedup_from_embeddings(emb)
    assert all(type(i) is int for i in keep)
    json.dumps(keep)


def test_semantic_idempotent_on_survivors():
    emb = _sample_embeddings()
    dedup = SemanticDeduplicator(threshold=0.95, topk=20)
    keep1 = dedup.dedup_from_embeddings(emb)
    survivors = emb[keep1]
    keep2 = dedup.dedup_from_embeddings(survivors)
    assert keep2 == list(range(len(survivors)))


def test_semantic_permutation_invariance_keep_count():
    emb = _sample_embeddings()
    dedup = SemanticDeduplicator(threshold=0.95, topk=20)
    keep_a = dedup.dedup_from_embeddings(emb)

    rng = np.random.default_rng(42)
    perm = rng.permutation(len(emb))
    shuffled = emb[perm]
    keep_b = dedup.dedup_from_embeddings(shuffled)

    assert len(keep_a) == len(keep_b)


def test_semantic_doubling_planted_cluster_does_not_increase_kept_count():
    """Adding more copies of an existing planted cluster can only keep len(keep) the same."""
    emb = _sample_embeddings()
    dedup = SemanticDeduplicator(threshold=0.95, topk=30)
    keep_before = dedup.dedup_from_embeddings(emb)

    # Duplicate the last 4 rows (one of the planted clusters) three more times
    extra = np.tile(emb[-4:], (3, 1))
    augmented = np.concatenate([emb, extra], axis=0)
    keep_after = dedup.dedup_from_embeddings(augmented)
    assert len(keep_after) == len(keep_before)


# ════════════════════════════════════════════════════════════════════════════
# Section 3: End-to-end preprocess-script pipeline simulation
# ════════════════════════════════════════════════════════════════════════════

def test_pipeline_exact_then_minhash_matches_preprocess_script_shape():
    """Exact reservoir dedup + MinHash with longest-response key_fn β€” same shape
    as preprocess_nemotron_cascade_science.py."""
    records_input = _make_records(REALISTIC_SCIENCE_PROMPTS)

    # Stage 1: Exact dedup (mimic the reservoir loop's "first seen wins" net effect)
    seen: dict[str, dict] = {}
    for rec in records_input:
        prompt = rec["conversations"][0]["value"]
        if prompt not in seen:
            seen[prompt] = rec
    records = list(seen.values())
    assert len(records) == 14  # 3 byte-identical -> 1; everything else distinct strings

    # Stage 2: MinHash with longest-response key_fn
    prompts = [r["conversations"][0]["value"] for r in records]
    rep_key = lambda i: -len(records[i]["conversations"][1]["value"])
    keep = MinHashDeduplicator(threshold=0.8).dedup(prompts, key_fn=rep_key, show_progress=False)
    records = [records[i] for i in keep]
    assert len(records) == 10  # formatting cluster -> 1; near-dup cluster -> 1


def test_pipeline_exact_then_minhash_then_semantic_with_orthogonal_embeddings():
    """Three-stage pipeline. Use orthogonal (random high-dim) embeddings so the
    semantic stage shouldn't merge anything β€” verifies the wiring, not the encoder."""
    records_input = _make_records(REALISTIC_SCIENCE_PROMPTS)

    seen: dict[str, dict] = {}
    for rec in records_input:
        seen.setdefault(rec["conversations"][0]["value"], rec)
    records = list(seen.values())

    prompts = [r["conversations"][0]["value"] for r in records]
    rep_key = lambda i: -len(records[i]["conversations"][1]["value"])
    keep = MinHashDeduplicator(threshold=0.8).dedup(prompts, key_fn=rep_key, show_progress=False)
    records = [records[i] for i in keep]
    n_after_minhash = len(records)

    # Stage 3: Semantic with random high-dim orthogonal embeddings
    rng = np.random.default_rng(0)
    mock_emb = _normalize(rng.standard_normal((len(records), 384)).astype(np.float32))
    rep_key = lambda i: -len(records[i]["conversations"][1]["value"])
    keep = SemanticDeduplicator(threshold=0.92).dedup_from_embeddings(mock_emb, key_fn=rep_key)
    records = [records[i] for i in keep]

    assert len(records) == n_after_minhash  # random embeddings don't merge


def test_pipeline_longest_response_representative_actually_wins():
    """Verify the longest-response key_fn picks the right rep within a cluster."""
    # Build 3 byte-identical prompts but with response lengths [10, 1000, 100]
    prompt = "Same prompt text used three times to form a cluster of duplicates."
    records = [
        {"conversations": [{"from": "human", "value": prompt}, {"from": "gpt", "value": "x" * 10}], "source": "t"},
        {"conversations": [{"from": "human", "value": prompt}, {"from": "gpt", "value": "x" * 1000}], "source": "t"},
        {"conversations": [{"from": "human", "value": prompt}, {"from": "gpt", "value": "x" * 100}], "source": "t"},
    ]
    prompts = [r["conversations"][0]["value"] for r in records]
    rep_key = lambda i: -len(records[i]["conversations"][1]["value"])
    keep = MinHashDeduplicator(threshold=0.8).dedup(prompts, key_fn=rep_key, show_progress=False)
    assert keep == [1]  # index 1 has the longest response


# ════════════════════════════════════════════════════════════════════════════
# Section 4: Realistic-data tests
# ════════════════════════════════════════════════════════════════════════════

def test_minhash_on_realistic_prompts_collapses_classes_correctly():
    """When given all 16 realistic prompts directly (no exact-dedup pre-step),
    MinHash should collapse: byte-identical (3->1), formatting (3->1),
    near-dup (3->1), and leave 7 singletons. Total kept = 10."""
    keep = MinHashDeduplicator(threshold=0.8).dedup(
        REALISTIC_SCIENCE_PROMPTS, show_progress=False
    )
    assert len(keep) == 10


def test_minhash_keeps_one_representative_from_each_realistic_cluster():
    """Each cluster (A/B/C) contributes exactly one representative to `keep`."""
    keep = MinHashDeduplicator(threshold=0.8).dedup(
        REALISTIC_SCIENCE_PROMPTS, show_progress=False
    )
    keep_set = set(keep)
    # Class A (byte-identical, indices 0-2): exactly 1 kept
    assert sum(1 for i in EXACT_DUP_RANGE if i in keep_set) == 1
    # Class B (formatting variants, indices 3-5): exactly 1 kept
    assert sum(1 for i in FORMATTING_RANGE if i in keep_set) == 1
    # Class C (near-dup paraphrases, indices 6-8): exactly 1 kept
    assert sum(1 for i in NEAR_DUP_RANGE if i in keep_set) == 1
    # Class D (singletons): all 7 kept
    assert sum(1 for i in SINGLETON_RANGE if i in keep_set) == 7


# ════════════════════════════════════════════════════════════════════════════
# Section 5: Backward-compat snapshots
# ════════════════════════════════════════════════════════════════════════════
#
# These tests pin the EXACT output for known input + seed. If anything in the
# default behavior changes (shingling, normalization, LSH band selection, FAISS
# search defaults, representative-tie-breaking), one of these will break and
# you'll know immediately.
#
# To intentionally update the snapshot, run the test, copy the new output into
# the assertion. Treat this as a deliberate behavioral change requiring review.

def test_minhash_snapshot_default_behavior_seed42():
    """Pinned: defaults (threshold=0.8, num_perm=128, 5-word shingles, normalize=True, seed=42).

    The "fox" text is ~120 tokens so a single-word edit (dawn->dusk) gives
    Jaccard β‰ˆ 0.92, robustly above threshold for LSH detection.
    """
    fox = (
        "the quick brown fox jumps over the lazy dog at sunrise on a cool "
        "morning before dawn arrives in the misty meadow where the rabbits "
        "nest among the tall grasses by the gentle stream that flows down "
        "from the distant hills covered in spring wildflowers and where "
        "songbirds sing as the sun rises over the eastern ridge bringing "
        "warm light to the valley below where farmers begin their daily work "
        "tending fields of corn and wheat that stretch toward the horizon "
        "under a clear blue sky filled with drifting white clouds today"
    )
    different = (
        "completely unrelated content here today tomorrow next week and the "
        "month after that with absolutely no shared vocabulary at all and "
        "no overlapping themes or ideas to be found in this paragraph that "
        "talks about something entirely separate from foxes and meadows or "
        "anything natural touching instead on abstract concepts in pure "
        "mathematics like topology category theory and algebraic geometry "
        "as well as functional programming languages such as Haskell OCaml "
        "and Standard ML which influence modern type systems profoundly today"
    )
    texts = [
        fox,
        fox,  # exact dup
        different,
        fox.replace("dawn", "dusk", 1),  # 1-word edit (dawn->dusk)
    ]
    keep = MinHashDeduplicator(seed=42).dedup(texts, show_progress=False)
    assert keep == [0, 2]


def test_semantic_snapshot_default_behavior():
    """Pinned: defaults (threshold=0.92, topk=10) on handcrafted normalized embeddings."""
    emb = _normalize(
        np.array(
            [
                [1.0, 0.0, 0.0, 0.0],
                [1.0, 0.0, 0.0, 0.0],  # identical to 0
                [0.0, 1.0, 0.0, 0.0],  # orthogonal to 0
                [np.cos(np.radians(15)), np.sin(np.radians(15)), 0.0, 0.0],  # cos=0.966 with 0
            ],
            dtype=np.float32,
        )
    )
    keep = SemanticDeduplicator().dedup_from_embeddings(emb)
    # Cluster {0, 1, 3} via cos>=0.92; singleton {2}
    assert keep == [0, 2]