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