File size: 2,870 Bytes
1d9bd9b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Lean dense+cross-encoder retrieval harness for the eval loop. Loads ONLY what dense+rerank needs
(no BM25 index -> ~2.5GB less RAM, no 68s/query pure-python scan). Emits run.tsv (qid, rank, doc_id).
Knobs via env: CAND (pool depth), ALPHA (authority prior weight on log1p(cite_indeg))."""
import json, os, time
import numpy as np
from sentence_transformers import SentenceTransformer, CrossEncoder
DATA = os.environ.get("THEMIS_DATA", "/Users/gongura/Code/themis/phase1/data/thor_artifacts")
EVAL = os.environ.get("THEMIS_EVAL", "/Users/gongura/Code/themis/phase1/eval")
DEVICE = os.environ.get("THEMIS_DEVICE", "cpu")       # "cuda" on the GPU box
CAND  = int(os.environ.get("CAND", "40"))
ALPHA = float(os.environ.get("ALPHA", "0"))           # 0 = pure cross-encoder (current system); >0 adds authority prior
BGE_Q = "Represent this sentence for searching relevant passages: "

print("loading chunks/vectors/models ...", flush=True)
chunk_doc = []; texts = []
for l in open(f"{DATA}/escr_chunks.jsonl"):
    c = json.loads(l); texts.append(c["text"]); chunk_doc.append(c["doc_id"])
M = np.load(f"{DATA}/escr_vectors.npy")
cite_indeg = {}
if ALPHA:
    from collections import Counter
    ci = Counter()
    for l in open(f"{DATA}/edges.jsonl"):
        e = json.loads(l)
        if e.get("method") == "cite": ci[e["target"]] += 1
    cite_indeg = ci
st = SentenceTransformer("BAAI/bge-small-en-v1.5", device="cpu")
ce = CrossEncoder("cross-encoder/ms-marco-MiniLM-L-6-v2", device="cpu")
print(f"ready (CAND={CAND} ALPHA={ALPHA})", flush=True)

def dense(q, n=CAND):
    qv = st.encode(BGE_Q + q, normalize_embeddings=True, convert_to_numpy=True).astype(np.float32)
    sim = M @ qv
    cand = np.argpartition(-sim, n)[:n]
    return [int(ci) for ci in cand[np.argsort(-sim[cand])]]

def _sig(x): return 1.0 / (1.0 + np.exp(-x))
def rerank(q, cand, k=20):
    rr = ce.predict([(q, texts[ci]) for ci in cand]); best = {}
    for ci, s in zip(cand, rr):
        d = chunk_doc[ci]
        if d not in best or s > best[d][0]: best[d] = (float(s), ci)
    scored = []
    for d, (s, ci) in best.items():
        sc = _sig(s) + (ALPHA * np.log1p(cite_indeg.get(d, 0)) if ALPHA else 0.0)
        scored.append((sc, d))
    scored.sort(reverse=True)
    return [d for _, d in scored[:k]]

def main():
    qs = [l.rstrip("\n").split("\t", 2) for l in open(f"{EVAL}/queries.tsv")]
    t0 = time.time()
    with open(f"{EVAL}/run.tsv", "w") as f:
        for i, (qid, intent, text) in enumerate(qs):
            for rank, d in enumerate(rerank(text, dense(text)), 1):
                f.write(f"{qid}\t{rank}\t{d}\n")
            f.flush()
            if (i + 1) % 50 == 0:
                print(f"{i+1}/{len(qs)}  {(time.time()-t0)/(i+1):.2f}s/q", flush=True)
    print(f"done {len(qs)} in {time.time()-t0:.0f}s -> run.tsv", flush=True)

if __name__ == "__main__":
    main()