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
|