File size: 6,323 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
"""Tests for SemanticDeduplicator.

Avoids downloading any embedding model by exercising `dedup_from_embeddings`
directly with handcrafted L2-normalized vectors. Skipped if faiss is not
installed.
"""

from __future__ import annotations

import pytest

pytest.importorskip("faiss")
import numpy as np

from sdg.preprocessing.dedupe.semantic import SemanticDeduplicator


def _normalize(v: np.ndarray) -> np.ndarray:
    norms = np.linalg.norm(v, axis=1, keepdims=True)
    return (v / np.clip(norms, 1e-12, None)).astype(np.float32)


def _vec(*components: float) -> np.ndarray:
    return np.array(components, dtype=np.float32)


# ── Boundary cases ───────────────────────────────────────────────────────────

def test_empty_returns_empty():
    dedup = SemanticDeduplicator()
    assert dedup.dedup_from_embeddings(np.zeros((0, 4), dtype=np.float32)) == []


def test_single_returns_single():
    dedup = SemanticDeduplicator()
    emb = _normalize(np.array([[1.0, 0.0, 0.0, 0.0]]))
    assert dedup.dedup_from_embeddings(emb) == [0]


# ── Identity / orthogonality ─────────────────────────────────────────────────

def test_identical_embeddings_collapse():
    dedup = SemanticDeduplicator(threshold=0.9)
    emb = _normalize(np.stack([_vec(1, 0, 0), _vec(1, 0, 0)]))
    keep = dedup.dedup_from_embeddings(emb)
    assert len(keep) == 1


def test_orthogonal_embeddings_both_kept():
    dedup = SemanticDeduplicator(threshold=0.5)
    emb = _normalize(np.stack([_vec(1, 0, 0), _vec(0, 1, 0)]))
    keep = dedup.dedup_from_embeddings(emb)
    assert keep == [0, 1]


# ── Threshold behavior ───────────────────────────────────────────────────────

def test_above_threshold_collapses():
    # cos(a, b) β‰ˆ 0.95 (small angle)
    a = _vec(1.0, 0.0)
    b = _vec(np.cos(np.radians(18)), np.sin(np.radians(18)))  # ~0.951
    emb = _normalize(np.stack([a, b]))
    dedup = SemanticDeduplicator(threshold=0.92)
    keep = dedup.dedup_from_embeddings(emb)
    assert len(keep) == 1


def test_below_threshold_both_kept():
    # cos(a, b) β‰ˆ 0.866 (30 degrees)
    a = _vec(1.0, 0.0)
    b = _vec(np.cos(np.radians(30)), np.sin(np.radians(30)))
    emb = _normalize(np.stack([a, b]))
    dedup = SemanticDeduplicator(threshold=0.92)
    keep = dedup.dedup_from_embeddings(emb)
    assert keep == [0, 1]


# ── Cluster of three ─────────────────────────────────────────────────────────

def test_cluster_of_three_via_chain_above_threshold():
    a = _vec(1.0, 0.0)
    b = _vec(np.cos(np.radians(10)), np.sin(np.radians(10)))
    c = _vec(np.cos(np.radians(20)), np.sin(np.radians(20)))
    emb = _normalize(np.stack([a, b, c]))
    dedup = SemanticDeduplicator(threshold=0.95)
    keep = dedup.dedup_from_embeddings(emb)
    # All three angles within ~20 degrees -> all pairwise cos > 0.93,
    # so they form one cluster -> one kept.
    assert len(keep) == 1


# ── Representative selection ─────────────────────────────────────────────────

def test_key_fn_selects_representative():
    emb = _normalize(np.stack([_vec(1, 0), _vec(1, 0), _vec(1, 0)]))
    response_lengths = [10, 500, 50]
    key_fn = lambda i: -response_lengths[i]
    dedup = SemanticDeduplicator(threshold=0.99)
    keep = dedup.dedup_from_embeddings(emb, key_fn=key_fn)
    assert keep == [1]


# ── HNSW path ────────────────────────────────────────────────────────────────

def test_hnsw_path_runs_above_threshold_size():
    """Force the HNSW code path with a tiny hnsw_threshold.

    Uses 384-dim random gaussians (matching our default embed model). At that
    dimensionality, random pairs have expected cosine ~0 with std ~1/sqrt(384),
    so a threshold of 0.95 is deep in the tail and nothing should collapse.
    """
    rng = np.random.default_rng(42)
    raw = rng.standard_normal((200, 384)).astype(np.float32)
    emb = _normalize(raw)
    dedup = SemanticDeduplicator(threshold=0.95, hnsw_threshold=10)
    keep = dedup.dedup_from_embeddings(emb)
    assert len(keep) == 200


def test_hnsw_path_collapses_planted_duplicates():
    """HNSW should still find planted near-duplicates among many distractors."""
    rng = np.random.default_rng(123)
    distractors = rng.standard_normal((180, 384)).astype(np.float32)
    # Plant 20 copies of one vector
    seed_vec = rng.standard_normal((1, 384)).astype(np.float32)
    planted = np.repeat(seed_vec, 20, axis=0)
    raw = np.concatenate([distractors, planted], axis=0)
    emb = _normalize(raw)
    dedup = SemanticDeduplicator(threshold=0.95, hnsw_threshold=10, topk=25)
    keep = dedup.dedup_from_embeddings(emb)
    # 180 distractors + 1 representative of the planted cluster = 181
    assert len(keep) == 181


# ── Validation ───────────────────────────────────────────────────────────────

def test_invalid_threshold_raises():
    with pytest.raises(ValueError):
        SemanticDeduplicator(threshold=0.0)


def test_invalid_topk_raises():
    with pytest.raises(ValueError):
        SemanticDeduplicator(topk=0)


# ── Device resolution ────────────────────────────────────────────────────────

def test_device_explicit_passthrough():
    assert SemanticDeduplicator(device="cpu").resolve_device() == "cpu"


def test_device_auto_returns_known_string():
    """auto must resolve to one of mps/cuda/cpu (depending on host)."""
    resolved = SemanticDeduplicator(device="auto").resolve_device()
    assert resolved in {"mps", "cuda", "cpu"}