svd-code / sdg /preprocessing /dedupe /test_large.py
fzzhang's picture
Upload folder using huggingface_hub
58258b8 verified
Raw History Blame Contribute Delete
8.82 kB
"""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