-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathhybrid.py
More file actions
360 lines (291 loc) · 14.1 KB
/
Copy pathhybrid.py
File metadata and controls
360 lines (291 loc) · 14.1 KB
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
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
"""Hybrid retrieval: BM25 fused with the dense index by reciprocal rank fusion.
WHY THIS EXISTS, and it comes straight out of this project's own measurements.
The evaluation found that LLM keyword extraction LOWERS precision by 19 to 31%
against simply embedding the whole question, and diagnosed the cause: generic
terms. "adversaries" literally appears in 100% of the documents it retrieves,
so it discriminates nothing while still consuming ten slots in a pool of
twenty five. A sweep of LITERAL_MATCH_BONUS from 0.0 to 1.0 left precision
pinned at 0.24 to 0.25 at every value including zero, so the bonus was never
the fix.
LITERAL_MATCH_BONUS is a hand rolled term weighting scheme with one term and
one weight: it adds a flat 0.15 if the keyword appears anywhere in the
document. It cannot tell "APT29", which appears in 2% of what it retrieves,
from "adversaries", which appears in all of it.
BM25 is that same idea done properly. Inverse document frequency means a rare
token like T1566, LSASS, APT29 or CVE-2021-34527 carries weight in proportion
to how rare it is, and a term appearing everywhere carries almost none. That
is exactly the failure that was measured, and dense embeddings are worst at
precisely this case, because an embedding blurs an exact identifier into the
neighbourhood of similar looking identifiers.
Fusion is by rank rather than by score. BM25 scores and cosine similarities
have no common scale, so any weighted sum of them needs a constant tuned on
one corpus that will not transfer to another. RRF needs no such constant.
NOTHING HERE IS ASSUMED TO HELP. On the sister project, four of five standard
retrieval upgrades made the main task WORSE, including the one predicted to be
the biggest single win. So hybrid retrieval ships behind a flag, evaluate.py
measures it against the plain dense arm on 579 held out queries, and the
numbers decide. Do not put a figure on a resume that has not come out of that
harness.
"""
from __future__ import annotations
import json
import os
import numpy as np
RRF_K = 60 # the constant from the reciprocal rank fusion paper
CORPUS_FILE = os.path.join(os.path.dirname(os.path.abspath(__file__)),
"mitre_corpus.json")
class BM25:
"""BM25 Okapi over a sparse term count matrix.
Written out rather than adding rank_bm25 as a dependency: it is a closed
form, scikit-learn is already installed, and a sparse matrix handles a few
thousand chunks in single digit milliseconds.
"""
def __init__(self, texts, k1=1.5, b=0.75):
from sklearn.feature_extraction.text import CountVectorizer
self.k1, self.b = k1, b
# The default token pattern drops one character tokens and splits on
# punctuation, which would destroy "T1059.001". Kept together instead,
# since a technique id is the single most valuable exact token here.
self.vectorizer = CountVectorizer(
lowercase=True, token_pattern=r"(?u)\b[\w.\-]{2,}\b")
# A corpus with no usable tokens makes CountVectorizer raise
# "empty vocabulary". That is a degenerate corpus, not a programming
# error, and it must not take the app down: the sparse arm simply has
# nothing to say and every query scores zero, which fusion handles.
self.empty = False
try:
counts = self.vectorizer.fit_transform(texts).astype(np.float32)
except ValueError:
self.empty = True
self.n_docs = len(texts)
return
self.doc_len = np.asarray(counts.sum(axis=1)).ravel()
self.avgdl = float(self.doc_len.mean()) if len(self.doc_len) else 0.0
n_docs = counts.shape[0]
df = np.asarray((counts > 0).sum(axis=0)).ravel()
self.idf = np.log(1.0 + (n_docs - df + 0.5) / (df + 0.5)).astype(np.float32)
self.norm = (self.k1 * (1 - self.b + self.b * self.doc_len
/ (self.avgdl or 1.0))).astype(np.float32)
self.counts = counts.tocsc()
def scores(self, query):
if self.empty:
return np.zeros(self.n_docs, dtype=np.float32)
terms = self.vectorizer.build_analyzer()(query)
vocab = self.vectorizer.vocabulary_
indices = [vocab[term] for term in terms if term in vocab]
total = np.zeros(self.counts.shape[0], dtype=np.float32)
if not indices:
return total
for index in indices:
freq = self.counts[:, index].toarray().ravel()
nonzero = freq > 0
if not nonzero.any():
continue
total[nonzero] += (
self.idf[index] * freq[nonzero] * (self.k1 + 1.0)
/ (freq[nonzero] + self.norm[nonzero])
)
return total
def top(self, query, k=50):
scores = self.scores(query)
if not scores.any():
return []
k = min(k, len(scores))
best = np.argpartition(-scores, k - 1)[:k]
return [int(i) for i in best[np.argsort(-scores[best])] if scores[i] > 0]
def rrf_fuse(*rankings, k=RRF_K, top=None):
"""Combine rankings by position. Returns [(item, score), ...] best first."""
fused = {}
for ranking in rankings:
for rank, item in enumerate(ranking, start=1):
fused[item] = fused.get(item, 0.0) + 1.0 / (k + rank)
order = sorted(fused.items(), key=lambda pair: -pair[1])
return order[:top] if top else order
# ------------------------------------------------------------- the corpus
def save_corpus(rows, path=CORPUS_FILE):
"""Write the chunk texts BM25 needs.
Called as a side effect of indexing, so the normal path costs nothing.
Small: a few thousand chunks of MITRE prose is single digit megabytes.
"""
with open(path, "w", encoding="utf-8") as handle:
json.dump(rows, handle)
return len(rows)
def load_corpus(path=CORPUS_FILE, index=None, namespace=""):
"""Chunk texts for the sparse arm, or None if they cannot be had.
Three sources, cheapest first. Returning None rather than raising is
deliberate: no corpus means the app runs dense only, which is exactly what
it did before this module existed, and a missing cache file is not a reason
to refuse to start.
"""
if os.path.exists(path):
try:
with open(path, encoding="utf-8") as handle:
rows = json.load(handle)
if rows:
return rows
except (ValueError, OSError):
pass
if index is None:
return None
rows = pull_corpus_from_index(index, namespace)
if rows:
save_corpus(rows, path)
return rows or None
def pull_corpus_from_index(index, namespace="", batch=100):
"""Read every chunk back out of Pinecone.
The fallback for an index that was populated before this file existed, so
hybrid search can be switched on without a reindex. One pass, then cached.
"""
try:
ids = []
for page in index.list(namespace=namespace):
ids.extend(page if isinstance(page, list) else [page])
except Exception as exc: # noqa: BLE001
print(f"Could not list the index for a corpus pull: {exc}")
return []
rows = []
for start in range(0, len(ids), batch):
chunk_ids = ids[start:start + batch]
try:
response = index.fetch(ids=chunk_ids, namespace=namespace)
except Exception as exc: # noqa: BLE001
print(f"Corpus pull failed at {start}: {exc}")
break
vectors = getattr(response, "vectors", None)
if vectors is None and isinstance(response, dict):
vectors = response.get("vectors", {})
for vector_id, vector in (vectors or {}).items():
metadata = getattr(vector, "metadata", None)
if metadata is None and isinstance(vector, dict):
metadata = vector.get("metadata", {})
metadata = metadata or {}
rows.append({
"id": vector_id,
"technique_id": metadata.get("technique_id", vector_id.split("#")[0]),
"name": metadata.get("name", ""),
"chunk": metadata.get("chunk", ""),
})
return rows
def corpus_rows_from_vectors(vectors):
"""The corpus form of what build_vectors just produced, without the floats.
Keeps indexing and the sparse arm reading the same text: if the two ever
drift, BM25 starts scoring documents the dense index does not contain and
fusion silently degrades to dense only for those ids.
"""
return [{
"id": vector["id"],
"technique_id": vector["metadata"]["technique_id"],
"name": vector["metadata"]["name"],
"chunk": vector["metadata"].get("chunk", ""),
} for vector in vectors]
class SparseArm:
"""BM25 over the indexed chunks, keyed by the same ids Pinecone uses.
Sharing the id space is what makes fusion possible at all: RRF combines
positions of the same item in two rankings, so the two arms have to be
naming the same things.
"""
def __init__(self, rows):
self.rows = rows
self.ids = [row["id"] for row in rows]
self.by_id = {row["id"]: row for row in rows}
# Name prefixed, matching how the vectors were built, so "PowerShell"
# scores on the PowerShell technique rather than only on chunks that
# happen to repeat the word in the body.
#
# The technique id is prefixed too, and it is NOT in the embedded text.
# Measured: searching "T1490" returned T1529 and T1218.014, because the
# only chunks containing that string were the ones CROSS REFERENCING
# T1490, while T1490's own text never names itself. So the one query an
# analyst is most likely to type by hand found everything except the
# right answer.
#
# Safe to differ from the dense text because BM25 only ever reorders
# documents the index already contains; what must not drift is the set
# of ids, not their wording.
self.bm25 = BM25([f"{row['technique_id']} {row['name']}. {row['chunk']}"
for row in rows])
def top_ids(self, query, k=50):
return [self.ids[i] for i in self.bm25.top(query, k)]
def __len__(self):
return len(self.rows)
def as_plain_match(match, metadata=None):
"""A Pinecone match copied into a real dict.
`dict(match)` DOES NOT WORK on what Pinecone returns. Its ScoredVector is an
OpenAPI generated model, and dict() on it reaches for .keys(), which exists
as an unset attribute holding None, so the call fails with the thoroughly
unhelpful "'NoneType' object is not callable".
Only `.get()` is safe on the outer object. Its metadata IS dict convertible,
which is why the older code path never hit this.
Worth remembering generally: a mock that returns plain dicts proves nothing
about a client library that returns models. This crashed on the first real
query while 724 offline checks stayed green.
"""
if metadata is None:
metadata = match.get("metadata") or {}
return {
"id": match.get("id"),
"score": match.get("score", 0.0),
"metadata": dict(metadata),
}
def fuse_matches(dense_matches, sparse_ids, sparse_arm, top_k):
"""Fuse a Pinecone response with a BM25 ranking, keeping match shape.
Sparse-only hits are rebuilt into the same dict a Pinecone match uses so
the caller cannot tell them apart, and they carry the fused score rather
than a cosine. Mixing a cosine and an RRF score in one list would put every
sparse hit at the bottom regardless of rank, which is the quiet way a
hybrid search ends up dense only.
"""
dense_ids = [match.get("id") for match in dense_matches]
by_id = {match.get("id"): match for match in dense_matches}
fused = rrf_fuse(dense_ids, sparse_ids, top=top_k)
sparse_set = set(sparse_ids)
results = []
for vector_id, score in fused:
match = by_id.get(vector_id)
if match is None:
row = sparse_arm.by_id.get(vector_id)
if row is None:
continue
# Only the fields retrieval reads. The association lists live on the
# dense match; a sparse-only hit is enriched by the caller's fetch.
plain = {"id": vector_id, "score": 0.0, "metadata": {
"technique_id": row["technique_id"],
"name": row["name"],
"chunk": row["chunk"],
"description": row["chunk"],
}}
else:
plain = as_plain_match(match)
plain["score"] = float(score)
plain["arm"] = ("both" if vector_id in by_id and vector_id in sparse_set
else ("dense" if vector_id in by_id else "sparse"))
results.append(plain)
return results
def load_reranker(name="cross-encoder/ms-marco-MiniLM-L-6-v2"):
"""Cross-encoder, or None if it will not load.
A bi-encoder embeds query and document separately and never compares them
directly. A cross-encoder reads both together, which is far more accurate
and far too slow to run over a corpus, so it only ever sees a shortlist the
first stage produced. Off by default: unmeasured here, and on the sister
project most such upgrades lost.
"""
try:
from sentence_transformers import CrossEncoder
return CrossEncoder(name, max_length=512)
except Exception as exc: # noqa: BLE001
print(f"Cross-encoder unavailable ({type(exc).__name__}), skipping rerank")
return None
def rerank(reranker, query, matches, keep=None):
"""Reorder matches by a cross-encoder reading the query against each chunk."""
if reranker is None or not matches:
return matches
texts = []
for match in matches:
metadata = match.get("metadata", {})
texts.append(f"{metadata.get('name', '')}. "
f"{metadata.get('chunk') or metadata.get('description', '')}"[:2000])
scores = reranker.predict([(query, text) for text in texts],
show_progress_bar=False)
ordered = [match for match, _ in
sorted(zip(matches, scores), key=lambda pair: -pair[1])]
return ordered[:keep] if keep else ordered