Download sdg/preprocessing/dedupe/test_large.py from fzzhang/svd-code: direct link, hf CLI and curl.
- Browser
- Download file 8.82 kB
-
https://huggingface.co/fzzhang/svd-code/resolve/main/sdg/preprocessing/dedupe/test_large.py
- Command line
-
hf download hf://fzzhang/svd-code/sdg/preprocessing/dedupe/test_large.py
-
curl -L -o test_large.py https://huggingface.co/fzzhang/svd-code/resolve/main/sdg/preprocessing/dedupe/test_large.py
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 | |