"""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