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