File size: 5,390 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 | """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
|