svd-code / sdg /preprocessing /dedupe /semantic.py
fzzhang's picture
Upload folder using huggingface_hub
58258b8 verified
Raw History Blame Contribute Delete
7.74 kB
"""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,
)