Download sdg/preprocessing/dedupe/semantic.py from fzzhang/svd-code: direct link, hf CLI and curl.
- Browser
- Download file 7.74 kB
-
https://huggingface.co/fzzhang/svd-code/resolve/main/sdg/preprocessing/dedupe/semantic.py
- Command line
-
hf download hf://fzzhang/svd-code/sdg/preprocessing/dedupe/semantic.py
-
curl -L -o semantic.py https://huggingface.co/fzzhang/svd-code/resolve/main/sdg/preprocessing/dedupe/semantic.py
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, | |
| ) | |