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