File size: 17,759 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
396
397
398
399
400
401
402
"""Speed benchmarks for MinHash + Semantic deduplicators.

Prints timing tables to stdout (uses `capsys.disabled()` so output appears
even without `pytest -s`). All benchmarks are bounded so the suite stays
under ~60s on an M-series machine.

To include the (slower) sentence-transformers encoding benchmark, set
SDG_BENCH_ENCODE=1 in the environment. That one downloads the model on
first run.

Run:
    pytest sdg/preprocessing/dedupe/test_benchmarks.py -v
    SDG_BENCH_ENCODE=1 pytest sdg/preprocessing/dedupe/test_benchmarks.py -v
"""

from __future__ import annotations

import gc
import os
import platform
import random
import resource
import time

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


# ────────────────────────────────────────────────────────────────────────────
# Helpers
# ────────────────────────────────────────────────────────────────────────────

_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"
).split()


def _distinct_prompts(n: int, words: int = 25, seed: int = 42) -> list[str]:
    rng = random.Random(seed)
    return [" ".join(rng.sample(_VOCAB, words)) + f" id{i}" for i in range(n)]


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 _print_machine_header(out, label: str) -> None:
    out(f"\n{'=' * 78}")
    out(f"  {label}")
    out(f"  host={platform.machine()} python={platform.python_version()} "
        f"darwin={platform.release()}")
    out(f"{'=' * 78}")


def _peak_rss_mb() -> float:
    """Peak resident memory in MB (macOS reports bytes; Linux reports KB)."""
    raw = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss
    if platform.system() == "Darwin":
        return raw / 1e6
    return raw / 1e3  # Linux: KB -> MB


# ────────────────────────────────────────────────────────────────────────────
# MinHash benchmarks
# ────────────────────────────────────────────────────────────────────────────

def test_bench_minhash_full_dedup_increasing_n(capsys):
    """End-to-end MinHash dedup throughput at increasing N (distinct prompts)."""
    sizes = [100, 500, 2_000, 10_000]
    rows = []
    for n in sizes:
        texts = _distinct_prompts(n)
        t0 = time.perf_counter()
        keep = MinHashDeduplicator(threshold=0.8).dedup(texts, show_progress=False)
        elapsed = time.perf_counter() - t0
        rps = n / elapsed if elapsed > 0 else float("inf")
        rows.append((n, elapsed, rps, len(keep)))

    with capsys.disabled():
        out = print
        _print_machine_header(out, "MinHash full dedup (signatures + LSH + cluster)")
        out(f"  {'N':>10s}  {'time':>10s}  {'rec/sec':>14s}  {'kept':>10s}")
        for n, elapsed, rps, kept in rows:
            out(f"  {n:>10,}  {elapsed:>9.3f}s  {rps:>12,.0f}  {kept:>10,}")
        out("=" * 78)


def test_bench_minhash_signature_only(capsys):
    """Signature creation alone (no LSH/query) β€” isolates shingling+hashing cost."""
    from datasketch import MinHash

    sizes = [100, 1_000, 10_000]
    rows = []
    deduper = MinHashDeduplicator()
    for n in sizes:
        texts = _distinct_prompts(n)
        t0 = time.perf_counter()
        for text in texts:
            shingles = deduper.shingles(text.lower())
            m = MinHash(num_perm=128, seed=42)
            for sh in shingles:
                m.update(sh)
        elapsed = time.perf_counter() - t0
        rps = n / elapsed if elapsed > 0 else float("inf")
        rows.append((n, elapsed, rps))

    with capsys.disabled():
        out = print
        _print_machine_header(out, "MinHash signature creation only (no LSH)")
        out(f"  {'N':>10s}  {'time':>10s}  {'rec/sec':>14s}")
        for n, elapsed, rps in rows:
            out(f"  {n:>10,}  {elapsed:>9.3f}s  {rps:>12,.0f}")
        out("=" * 78)


def test_bench_minhash_threshold_sweep(capsys):
    """Same N, vary threshold β€” high threshold should be faster (fewer LSH bands)."""
    n = 5_000
    texts = _distinct_prompts(n)
    rows = []
    for threshold in [0.5, 0.7, 0.8, 0.9, 0.95]:
        t0 = time.perf_counter()
        keep = MinHashDeduplicator(threshold=threshold).dedup(texts, show_progress=False)
        elapsed = time.perf_counter() - t0
        rows.append((threshold, elapsed, len(keep)))

    with capsys.disabled():
        out = print
        _print_machine_header(out, f"MinHash threshold sweep (N={n:,} distinct prompts)")
        out(f"  {'threshold':>10s}  {'time':>10s}  {'kept':>10s}")
        for threshold, elapsed, kept in rows:
            out(f"  {threshold:>10.2f}  {elapsed:>9.3f}s  {kept:>10,}")
        out("=" * 78)


def test_bench_minhash_num_perm_sweep(capsys):
    """Higher num_perm = better accuracy but slower."""
    n = 2_000
    texts = _distinct_prompts(n)
    rows = []
    for num_perm in [32, 64, 128, 256, 512]:
        t0 = time.perf_counter()
        MinHashDeduplicator(num_perm=num_perm).dedup(texts, show_progress=False)
        elapsed = time.perf_counter() - t0
        rows.append((num_perm, elapsed))

    with capsys.disabled():
        out = print
        _print_machine_header(out, f"MinHash num_perm sweep (N={n:,})")
        out(f"  {'num_perm':>10s}  {'time':>10s}")
        for num_perm, elapsed in rows:
            out(f"  {num_perm:>10d}  {elapsed:>9.3f}s")
        out("=" * 78)


# ────────────────────────────────────────────────────────────────────────────
# Semantic benchmarks (skip the encoder; bench dedup_from_embeddings)
# ────────────────────────────────────────────────────────────────────────────

def test_bench_semantic_dedup_from_embeddings_increasing_n(capsys):
    """Pure FAISS+cluster throughput at increasing N (random 384-dim vectors).

    Capped at 25K to stay clear of a faiss-cpu 1.8.0 macOS-arm64 memory-pressure
    regime that can segfault when the suite has accumulated state from prior
    tests. (50K runs fine in isolation; not when chained.)
    """
    sizes = [500, 2_000, 10_000, 25_000]
    rows = []
    for n in sizes:
        rng = np.random.default_rng(42)
        emb = _normalize(rng.standard_normal((n, 384)).astype(np.float32))
        dedup = SemanticDeduplicator(threshold=0.9, hnsw_threshold=5_000)
        kind = "Flat" if n < dedup.hnsw_threshold else "HNSW"
        t0 = time.perf_counter()
        keep = dedup.dedup_from_embeddings(emb, show_progress=False)
        elapsed = time.perf_counter() - t0
        rps = n / elapsed if elapsed > 0 else float("inf")
        rows.append((n, kind, elapsed, rps, len(keep)))

    with capsys.disabled():
        out = print
        _print_machine_header(out, "Semantic dedup_from_embeddings (FAISS index + cluster)")
        out(f"  {'N':>10s}  {'index':>8s}  {'time':>10s}  {'rec/sec':>14s}  {'kept':>10s}")
        for n, kind, elapsed, rps, kept in rows:
            out(f"  {n:>10,}  {kind:>8s}  {elapsed:>9.3f}s  {rps:>12,.0f}  {kept:>10,}")
        out("=" * 78)


def test_bench_semantic_threshold_sweep(capsys):
    """Lower threshold = more clusters to materialize but FAISS time is the same."""
    n = 5_000
    rng = np.random.default_rng(42)
    emb = _normalize(rng.standard_normal((n, 384)).astype(np.float32))
    rows = []
    for threshold in [0.7, 0.8, 0.9, 0.95, 0.99]:
        dedup = SemanticDeduplicator(threshold=threshold, hnsw_threshold=1_000)
        t0 = time.perf_counter()
        keep = dedup.dedup_from_embeddings(emb, show_progress=False)
        elapsed = time.perf_counter() - t0
        rows.append((threshold, elapsed, len(keep)))

    with capsys.disabled():
        out = print
        _print_machine_header(out, f"Semantic threshold sweep (N={n:,}, 384-dim)")
        out(f"  {'threshold':>10s}  {'time':>10s}  {'kept':>10s}")
        for threshold, elapsed, kept in rows:
            out(f"  {threshold:>10.2f}  {elapsed:>9.3f}s  {kept:>10,}")
        out("=" * 78)


def test_bench_semantic_dim_sweep(capsys):
    """Embedding dimensionality vs throughput at fixed N."""
    n = 5_000
    rows = []
    for dim in [128, 256, 384, 768, 1024]:
        rng = np.random.default_rng(42)
        emb = _normalize(rng.standard_normal((n, dim)).astype(np.float32))
        dedup = SemanticDeduplicator(threshold=0.95, hnsw_threshold=1_000)
        t0 = time.perf_counter()
        dedup.dedup_from_embeddings(emb, show_progress=False)
        elapsed = time.perf_counter() - t0
        rows.append((dim, elapsed))

    with capsys.disabled():
        out = print
        _print_machine_header(out, f"Semantic embedding-dim sweep (N={n:,}, HNSW)")
        out(f"  {'dim':>10s}  {'time':>10s}")
        for dim, elapsed in rows:
            out(f"  {dim:>10d}  {elapsed:>9.3f}s")
        out("=" * 78)


# ────────────────────────────────────────────────────────────────────────────
# Optional: real-encoder benchmark (gated by env var)
# ────────────────────────────────────────────────────────────────────────────

@pytest.mark.skipif(
    os.getenv("SDG_BENCH_ENCODE") != "1",
    reason="set SDG_BENCH_ENCODE=1 to run the real-encoder benchmark "
           "(downloads BAAI/bge-small-en-v1.5 on first run)",
)
def test_bench_semantic_real_encoder_small(capsys):
    pytest.importorskip("sentence_transformers")
    from sentence_transformers import SentenceTransformer

    sizes = [128, 512, 1024]
    dedup = SemanticDeduplicator()
    device = dedup.resolve_device()
    model = SentenceTransformer(dedup.model_name, device=device)

    rows = []
    for n in sizes:
        texts = _distinct_prompts(n)
        t0 = time.perf_counter()
        emb = model.encode(
            texts,
            batch_size=128,
            convert_to_numpy=True,
            normalize_embeddings=True,
            show_progress_bar=False,
        )
        elapsed = time.perf_counter() - t0
        rps = n / elapsed if elapsed > 0 else float("inf")
        rows.append((n, elapsed, rps))
        assert emb.shape == (n, 384)

    with capsys.disabled():
        out = print
        _print_machine_header(
            out, f"Semantic encoding (BAAI/bge-small-en-v1.5, device={device}, batch=128)"
        )
        out(f"  {'N':>10s}  {'time':>10s}  {'rec/sec':>14s}")
        for n, elapsed, rps in rows:
            out(f"  {n:>10,}  {elapsed:>9.3f}s  {rps:>12,.0f}")
        out("=" * 78)


# ────────────────────────────────────────────────────────────────────────────
# Large-scale projection: how long would 3M prompts take?
# ────────────────────────────────────────────────────────────────────────────
#
# The Nemotron-Cascade-2 science split has ~3M unique prompts after exact
# dedup. We benchmark MinHash and semantic-FAISS at 50K/100K/250K and
# extrapolate. Encoder cost (~30 min for 3M on MPS) is a separate stage β€”
# enable with SDG_BENCH_ENCODE=1 to measure.
#
# Memory: 250K signatures + LSH ~500-700 MB; 250K x 384 float32 = 384 MB.
# Both fit easily on a laptop. We `gc.collect()` between sizes defensively
# because faiss-cpu 1.8.0 has shown memory-pressure segfaults under pytest.

def test_bench_3m_projection_minhash(capsys):
    """Project MinHash dedup time to 3M using N=50K/100K/250K data points."""
    sizes = [50_000, 100_000, 250_000]
    rows = []
    for n in sizes:
        gc.collect()
        texts = _distinct_prompts(n)
        t0 = time.perf_counter()
        keep = MinHashDeduplicator(threshold=0.8).dedup(texts, show_progress=False)
        elapsed = time.perf_counter() - t0
        rps = n / elapsed if elapsed > 0 else float("inf")
        rss = _peak_rss_mb()
        rows.append((n, elapsed, rps, len(keep), rss))
        del texts, keep
        gc.collect()

    n_ref, _, rps_ref, _, _ = rows[-1]
    proj_sec = 3_000_000 / rps_ref

    with capsys.disabled():
        out = print
        _print_machine_header(out, "MinHash large-scale projection -> 3M target")
        out(f"  {'N':>10s}  {'time':>10s}  {'rec/sec':>12s}  {'kept':>10s}  {'peak RSS MB':>14s}")
        for n, elapsed, rps, kept, rss in rows:
            out(f"  {n:>10,}  {elapsed:>9.2f}s  {rps:>10,.0f}  {kept:>10,}  {rss:>13,.0f}")
        out(f"  {'PROJ 3M':>10s}  {proj_sec:>9.0f}s  ({proj_sec/60:.1f} min, "
            f"linear extrap from N={n_ref:,})")
        out("=" * 78)


def test_bench_3m_projection_semantic_faiss(capsys):
    """Project semantic dedup_from_embeddings time to 3M (FAISS only, no encoding)."""
    sizes = [50_000, 100_000, 250_000]
    rows = []
    for n in sizes:
        gc.collect()
        rng = np.random.default_rng(42)
        emb = _normalize(rng.standard_normal((n, 384)).astype(np.float32))
        dedup = SemanticDeduplicator(threshold=0.92, hnsw_threshold=10_000)
        t0 = time.perf_counter()
        keep = dedup.dedup_from_embeddings(emb, show_progress=False)
        elapsed = time.perf_counter() - t0
        rps = n / elapsed if elapsed > 0 else float("inf")
        rss = _peak_rss_mb()
        rows.append((n, elapsed, rps, len(keep), rss))
        del emb, keep, dedup
        gc.collect()

    # HNSW per-record cost grows with N (~log N). Use ratio from largest point
    # but warn that this is optimistic.
    n_ref, t_ref, rps_ref, _, _ = rows[-1]
    proj_sec_linear = 3_000_000 / rps_ref
    # NlogN scaling estimate: t_3M β‰ˆ t_ref * (3M/N_ref) * log(3M)/log(N_ref)
    import math
    proj_sec_nlogn = t_ref * (3_000_000 / n_ref) * (math.log(3_000_000) / math.log(n_ref))

    with capsys.disabled():
        out = print
        _print_machine_header(out, "Semantic FAISS large-scale projection -> 3M target")
        out(f"  {'N':>10s}  {'time':>10s}  {'rec/sec':>12s}  {'peak RSS MB':>14s}")
        for n, elapsed, rps, kept, rss in rows:
            out(f"  {n:>10,}  {elapsed:>9.2f}s  {rps:>10,.0f}  {rss:>13,.0f}")
        out(f"  {'PROJ 3M (linear)':>20s}  {proj_sec_linear:>9.0f}s  ({proj_sec_linear/60:.1f} min)")
        out(f"  {'PROJ 3M (N log N)':>20s}  {proj_sec_nlogn:>9.0f}s  ({proj_sec_nlogn/60:.1f} min)")
        out("=" * 78)


# ────────────────────────────────────────────────────────────────────────────
# End-to-end pipeline benchmark
# ────────────────────────────────────────────────────────────────────────────

def test_bench_full_pipeline_minhash_then_semantic(capsys):
    """End-to-end MinHash + semantic dedup on synthetic 5K-prompt dataset.

    Uses random embeddings for the semantic stage to avoid model download.
    """
    n = 5_000
    texts = _distinct_prompts(n)

    t0 = time.perf_counter()
    keep_mh = MinHashDeduplicator(threshold=0.8).dedup(texts, show_progress=False)
    t_mh = time.perf_counter() - t0

    rng = np.random.default_rng(0)
    emb = _normalize(rng.standard_normal((len(keep_mh), 384)).astype(np.float32))

    dedup = SemanticDeduplicator(threshold=0.92, hnsw_threshold=1_000)
    t0 = time.perf_counter()
    keep_sem = dedup.dedup_from_embeddings(emb, show_progress=False)
    t_sem = time.perf_counter() - t0

    with capsys.disabled():
        out = print
        _print_machine_header(out, f"End-to-end pipeline (N={n:,} prompts)")
        out(f"  {'stage':<32s} {'time':>10s} {'kept':>10s}")
        out(f"  {'MinHash (Jaccard 0.8)':<32s} {t_mh:>9.3f}s {len(keep_mh):>10,}")
        out(f"  {'Semantic (cosine 0.92, HNSW)':<32s} {t_sem:>9.3f}s {len(keep_sem):>10,}")
        out(f"  {'TOTAL':<32s} {t_mh + t_sem:>9.3f}s {len(keep_sem):>10,}")
        out("=" * 78)