83 lines
2.6 KiB
Python
83 lines
2.6 KiB
Python
import xxhash
|
|
from typing import Set, Tuple, List
|
|
from datasketch import MinHash, MinHashLSH
|
|
import numpy as np
|
|
|
|
class WinNowing:
|
|
def __init__(self, k: int = 5, window_size: int = 4):
|
|
self.k = k
|
|
self.window_size = window_size
|
|
|
|
def get_kgrams(self, text: str) -> List[str]:
|
|
words = text.split()
|
|
kgrams = []
|
|
for i in range(len(words) - self.k + 1):
|
|
kgram = ' '.join(words[i:i+self.k])
|
|
kgrams.append(kgram)
|
|
return kgrams
|
|
|
|
def fingerprint(self, text: str) -> Set[int]:
|
|
kgrams = self.get_kgrams(text)
|
|
hashes = []
|
|
for kgram in kgrams:
|
|
h = int(xxhash.xxh64(kgram).hexdigest(), 16)
|
|
hashes.append(h)
|
|
|
|
if not hashes:
|
|
return set()
|
|
|
|
fingerprints = set()
|
|
for i in range(len(hashes) - self.window_size + 1):
|
|
window = hashes[i:i+self.window_size]
|
|
min_hash = min(window)
|
|
fingerprints.add(min_hash)
|
|
|
|
return fingerprints
|
|
|
|
def compare(self, fp1: Set[int], fp2: Set[int]) -> float:
|
|
if not fp1 or not fp2:
|
|
return 0.0
|
|
intersection = len(fp1 & fp2)
|
|
union = len(fp1 | fp2)
|
|
return intersection / union if union > 0 else 0.0
|
|
|
|
class MinHashLSHIndex:
|
|
def __init__(self, num_perm: int = 128, threshold: float = 0.5):
|
|
self.num_perm = num_perm
|
|
self.threshold = threshold
|
|
self.lsh = MinHashLSH(threshold=threshold, num_perm=num_perm)
|
|
self.documents = {}
|
|
|
|
def add_document(self, doc_id: str, text: str):
|
|
kgrams = self._get_kgrams(text)
|
|
m = MinHash(num_perm=self.num_perm)
|
|
for kgram in kgrams:
|
|
m.update(kgram.encode())
|
|
|
|
self.lsh.insert(doc_id, m)
|
|
self.documents[doc_id] = m
|
|
|
|
def query(self, text: str, top_k: int = 10) -> List[Tuple[str, float]]:
|
|
kgrams = self._get_kgrams(text)
|
|
query_m = MinHash(num_perm=self.num_perm)
|
|
for kgram in kgrams:
|
|
query_m.update(kgram.encode())
|
|
|
|
candidates = self.lsh.query(query_m)
|
|
|
|
results = []
|
|
for doc_id in candidates:
|
|
similarity = query_m.jaccard(self.documents[doc_id])
|
|
results.append((doc_id, similarity))
|
|
|
|
results.sort(key=lambda x: x[1], reverse=True)
|
|
return results[:top_k]
|
|
|
|
def _get_kgrams(self, text: str, k: int = 5) -> List[str]:
|
|
words = text.split()
|
|
kgrams = []
|
|
for i in range(max(1, len(words) - k + 1)):
|
|
kgram = ' '.join(words[i:i+k])
|
|
kgrams.append(kgram)
|
|
return kgrams
|