File size: 7,737 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
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
"""Paraphrase dedup via sentence embeddings + FAISS ANN (SemDeDup-style).

Catches semantic near-duplicates that share few n-grams ("Solve x^2-4=0" vs
"Find roots of x squared minus four"). Embeds with sentence-transformers,
indexes with faiss, queries top-k, and clusters pairs with cosine >= threshold.

Defaults: BAAI/bge-small-en-v1.5 (384-dim, MPS-friendly), cosine >= 0.92.
Useful threshold range: 0.90-0.95. Below 0.88 starts dropping legitimately
distinct prompts; above 0.95 only catches near-exact paraphrases.

On Apple Silicon (M-series), `device="auto"` selects MPS for encoding.
FAISS always runs on CPU here (no faiss-gpu on macOS).
"""

from __future__ import annotations

import time
from typing import Any, Callable

from tqdm import tqdm

from .base import cluster_and_pick


class SemanticDeduplicator:
    def __init__(
        self,
        model_name: str = "BAAI/bge-small-en-v1.5",
        threshold: float = 0.92,
        batch_size: int = 128,
        device: str = "auto",  # "auto" | "mps" | "cuda" | "cpu"
        topk: int = 10,
        hnsw_threshold: int = 50_000,  # use IndexFlatIP below this size
    ):
        if not 0.0 < threshold <= 1.0:
            raise ValueError(f"threshold must be in (0, 1], got {threshold}")
        if topk < 1:
            raise ValueError(f"topk must be >= 1, got {topk}")
        self.model_name = model_name
        self.threshold = threshold
        self.batch_size = batch_size
        self.device = device
        self.topk = topk
        self.hnsw_threshold = hnsw_threshold

    def resolve_device(self) -> str:
        if self.device != "auto":
            return self.device
        try:
            import torch

            if torch.backends.mps.is_available():
                return "mps"
            if torch.cuda.is_available():
                return "cuda"
        except ImportError:
            pass
        return "cpu"

    def encode(self, texts: list[str], show_progress: bool = True):
        """Encode texts to L2-normalized float32 embeddings (np.ndarray)."""
        import gc
        import numpy as np
        from sentence_transformers import SentenceTransformer

        device = self.resolve_device()
        if show_progress:
            print(
                f"  Encoding {len(texts):,} texts with {self.model_name} on "
                f"device={device} (batch_size={self.batch_size})"
            )
        t0 = time.perf_counter()
        model = SentenceTransformer(self.model_name, device=device)
        embeddings = model.encode(
            texts,
            batch_size=self.batch_size,
            convert_to_numpy=True,
            normalize_embeddings=True,
            show_progress_bar=show_progress,
        )
        # Take ownership of the buffer so FAISS doesn't segfault on a
        # torch-backed memoryview, then drop the model + flush MPS cache + GC
        # to release allocations before FAISS starts grabbing memory.
        embeddings = np.array(embeddings, dtype=np.float32, copy=True, order="C")
        del model
        try:
            import torch

            if torch.backends.mps.is_available():
                torch.mps.empty_cache()
            elif torch.cuda.is_available():
                torch.cuda.empty_cache()
        except ImportError:
            pass
        gc.collect()
        if show_progress:
            print(f"    encoded in {time.perf_counter() - t0:.1f}s")
        return embeddings

    def dedup_from_embeddings(
        self,
        embeddings,
        key_fn: Callable[[int], Any] | None = None,
        show_progress: bool = True,
        return_clusters: bool = False,
    ) -> "list[int] | tuple[list[int], dict[int, list[int]]]":
        """Cluster pre-computed (L2-normalized) embeddings via FAISS top-k."""
        import faiss
        import gc
        import numpy as np

        # macOS arm64: PyTorch's libomp and FAISS's libomp can clash and
        # segfault under concurrent thread pools. Single-threading FAISS is
        # the standard workaround. (Tiny perf hit at our scale; eliminates
        # the segfault.)
        faiss.omp_set_num_threads(1)
        gc.collect()

        n = len(embeddings)
        if n == 0:
            return ([], {}) if return_clusters else []

        emb = np.ascontiguousarray(embeddings, dtype=np.float32)

        # Defensive: NaN/Inf in input embeddings can segfault FAISS internals.
        # Replace with zero vectors (which have inner-product 0 with everything,
        # well below any reasonable threshold, so they end up as singletons).
        nonfinite_rows = ~np.isfinite(emb).all(axis=1)
        n_nonfinite = int(nonfinite_rows.sum())
        if n_nonfinite > 0:
            if show_progress:
                print(
                    f"  WARNING: {n_nonfinite:,} of {n:,} embeddings contained "
                    f"NaN/Inf — replacing with zero vectors."
                )
            emb = np.nan_to_num(emb, nan=0.0, posinf=0.0, neginf=0.0)

        dim = emb.shape[1]

        if n < self.hnsw_threshold:
            index = faiss.IndexFlatIP(dim)
            index_kind = "Flat"
        else:
            # METRIC_INNER_PRODUCT so search() returns cosine sim (vectors are
            # L2-normalized). Default is L2, which would silently invert the
            # threshold semantics.
            index = faiss.IndexHNSWFlat(dim, 32, faiss.METRIC_INNER_PRODUCT)
            index.hnsw.efConstruction = 200
            index.hnsw.efSearch = 64
            index_kind = "HNSW"

        if show_progress:
            print(f"  Building {index_kind} FAISS index over {n:,} vectors (dim={dim}) ...")
        t0 = time.perf_counter()
        index.add(emb)
        if show_progress:
            print(f"    done in {time.perf_counter() - t0:.2f}s")

        topk = min(self.topk, n)
        if show_progress:
            print(f"  Searching top-{topk} neighbors for {n:,} queries ...")
        t0 = time.perf_counter()
        distances, indices = index.search(emb, topk)
        if show_progress:
            print(f"    done in {time.perf_counter() - t0:.2f}s")

        pairs: list[tuple[int, int]] = []
        t0 = time.perf_counter()
        iter_pairs = (
            tqdm(range(n), desc="  Extracting pairs", unit=" rec", smoothing=0.05, leave=False)
            if show_progress else range(n)
        )
        for i in iter_pairs:
            for d, j in zip(distances[i], indices[i]):
                if j == -1 or j <= i:
                    continue
                if d >= self.threshold:
                    pairs.append((i, int(j)))
        if show_progress:
            print(
                f"  Found {len(pairs):,} candidate pairs (cos>={self.threshold}) "
                f"in {time.perf_counter() - t0:.2f}s"
            )

        t0 = time.perf_counter()
        keep, clusters = cluster_and_pick(n, pairs, key_fn)
        if show_progress:
            print(
                f"  Built {len(keep):,} clusters via union-find "
                f"in {time.perf_counter() - t0:.2f}s"
            )
        if return_clusters:
            return keep, clusters
        return keep

    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]]]":
        n = len(texts)
        if n == 0:
            return ([], {}) if return_clusters else []
        embeddings = self.encode(texts, show_progress=show_progress)
        return self.dedup_from_embeddings(
            embeddings,
            key_fn=key_fn,
            show_progress=show_progress,
            return_clusters=return_clusters,
        )