svd-code / sdg /preprocessing /dedupe /minhash.py
fzzhang's picture
Upload folder using huggingface_hub
58258b8 verified
Raw History Blame Contribute Delete
5.39 kB
"""Lexical near-duplicate dedup via MinHash + LSH (datasketch).
Catches near-duplicates that share most of their token n-grams: copy-paste
with edits, formatting changes, slight rewordings. Misses true paraphrases
(use SemanticDeduplicator for those).
Defaults follow the recipe used in FineWeb / Dolma / DataComp-LM:
threshold=0.8 (Jaccard), num_perm=128, 5-word shingles.
"""
from __future__ import annotations
from typing import Any, Callable
from tqdm import tqdm
from .base import cluster_and_pick, normalize_text
class MinHashDeduplicator:
def __init__(
self,
threshold: float = 0.8,
num_perm: int = 128,
shingle_size: int = 5,
shingle_kind: str = "word", # "word" or "char"
normalize: bool = True,
seed: int = 42,
):
if shingle_kind not in {"word", "char"}:
raise ValueError(f"shingle_kind must be 'word' or 'char', got {shingle_kind!r}")
if not 0.0 < threshold < 1.0:
raise ValueError(
f"threshold must be in (0, 1) (datasketch's LSH cannot pick "
f"valid bands at threshold=1.0); got {threshold}"
)
if shingle_size < 1:
raise ValueError(f"shingle_size must be >= 1, got {shingle_size}")
self.threshold = threshold
self.num_perm = num_perm
self.shingle_size = shingle_size
self.shingle_kind = shingle_kind
self.normalize = normalize
self.seed = seed
def shingles(self, text: str) -> set[bytes]:
if self.shingle_kind == "word":
tokens = text.split()
if len(tokens) < self.shingle_size:
return {" ".join(tokens).encode("utf-8")}
return {
" ".join(tokens[i : i + self.shingle_size]).encode("utf-8")
for i in range(len(tokens) - self.shingle_size + 1)
}
if len(text) < self.shingle_size:
return {text.encode("utf-8")}
return {
text[i : i + self.shingle_size].encode("utf-8")
for i in range(len(text) - self.shingle_size + 1)
}
def dedup(
self,
texts: list[str],
key_fn: Callable[[int], Any] | None = None,
show_progress: bool = True,
return_clusters: bool = False,
) -> "list[int] | tuple[list[int], dict[int, list[int]]]":
from datasketch import MinHash, MinHashLSH
n = len(texts)
if n == 0:
return ([], {}) if return_clusters else []
iter_sigs = (
tqdm(texts, desc="MinHash signatures", unit=" rec", smoothing=0.05)
if show_progress
else texts
)
minhashes: list[MinHash] = []
for text in iter_sigs:
norm = normalize_text(text) if self.normalize else text
shingles = self.shingles(norm)
m = MinHash(num_perm=self.num_perm, seed=self.seed)
for sh in shingles:
m.update(sh)
minhashes.append(m)
lsh = MinHashLSH(threshold=self.threshold, num_perm=self.num_perm)
iter_insert = (
tqdm(enumerate(minhashes), total=n, desc="MinHash LSH insert", unit=" rec", smoothing=0.05)
if show_progress
else enumerate(minhashes)
)
for i, m in iter_insert:
lsh.insert(str(i), m)
iter_query = (
tqdm(enumerate(minhashes), total=n, desc="MinHash query", unit=" rec", smoothing=0.05)
if show_progress
else enumerate(minhashes)
)
# LSH banding returns *candidate* pairs from its probabilistic S-curve,
# which includes sub-threshold false positives. Without verification, a
# dense near-threshold blob (distinct items sharing a long boilerplate
# block, each pair ~0.65 < threshold) gets transitively collapsed by
# union-find. We verify each candidate with EXACT Jaccard before it
# becomes a merge edge. (The signature-estimated Jaccard is too noisy
# near the boundary — ~±0.1 at num_perm=128 — to use here.)
#
# Shingles are recomputed lazily only for candidate-involved indices and
# cached. Pairs are always (i, j) with j > i and queries run in
# increasing i, so index i is never referenced again after its own
# query — we evict it immediately to keep the cache bounded.
norm_fn = normalize_text if self.normalize else (lambda t: t)
shingle_cache: dict[int, set[bytes]] = {}
def shingles_for(idx: int) -> set[bytes]:
cached = shingle_cache.get(idx)
if cached is None:
cached = self.shingles(norm_fn(texts[idx]))
shingle_cache[idx] = cached
return cached
pairs: list[tuple[int, int]] = []
for i, m in iter_query:
candidates = [j for j in (int(c) for c in lsh.query(m)) if j > i]
if candidates:
si = shingles_for(i)
for j in candidates:
sj = shingles_for(j)
if len(si & sj) / len(si | sj) >= self.threshold:
pairs.append((i, j))
shingle_cache.pop(i, None)
keep, clusters = cluster_and_pick(n, pairs, key_fn)
if return_clusters:
return keep, clusters
return keep