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