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,
)
|