Team Ai
Datasetpublic

Angshul/SparseGeometricRAG

SparseGeometricRAG CPU-first sparse geometric retrieval for practical top-10 RAG No transformer inference at retrieval time. No retrieval GPU requirement. No dense document-vector dot products. No external API. SparseGeometricRAG is a retrieval system built around one systems objective: make the retrieval layer cheap enough to run on ordinary multicore CPU hardware without turning the corpus into a dense embedding database. It uses sparse TF-IDF geometry, fuzzy… See the full description on the dataset page: https://huggingface.co/datasets/Angshul/SparseGeometricRAG.

sourceHugging Facefair-noncommercial-research-licenseupdated 2mo agoView on Hugging Face
0likes331downloads
msmarco_lambdaM_uniform_diag.py83 linesDownload Raw Back to msmarco_scale
1import sys,time,json,math2from pathlib import Path3import numpy as np,pandas as pd4sys.path.insert(0,'/mnt/data')5import msmarco_full_search_uniform1m as m6 7ROOT=m.ROOT; WORK=m.WORK; P=m.P; M=m.M8LAMBDAS=[0.0,0.25,0.5,1.0,2.0,4.0,8.0,16.0,64.0]9HGRID=[0,1,2,3]10idx=m.FullIndex(); print('loaded',idx.meta,flush=True)11# exact same deterministic validation split12tr=pd.read_csv(ROOT/'train.tsv',sep='\t',usecols=['query-id']); uq=np.unique(tr['query-id'].to_numpy()); rng=np.random.default_rng(20260815); ids=[str(x) for x in rng.choice(uq,size=1000,replace=False)]; del tr13texts=m.load_query_texts(ids); qrels=m.qrels_from_tsv(ROOT/'train.tsv',ids,positive_only=True)14 15def components(text):16    q=idx.query_vec(text); qd=np.zeros(M,np.float32); qd[q.indices]=q.data; rterms,rd=idx.route(q)17    spans=[(int(j),int(idx.offs[j]),int(idx.offs[j+1])) for j in rterms if idx.offs[j+1]>idx.offs[j]]18    if not spans:return None19    docs=np.concatenate([np.asarray(idx.pd[a:b]) for j,a,b in spans]).astype(np.uint32,copy=False)20    mm=np.concatenate([np.asarray(idx.pm[a:b]) for j,a,b in spans]).astype(np.float32,copy=False)21    rt=np.concatenate([np.asarray(idx.pr[a:b]) for j,a,b in spans]).astype(np.uint16,copy=False)22    sb=np.concatenate([np.asarray(idx.ps[a:b]) for j,a,b in spans]).astype(np.uint16,copy=False)23    nr=len(spans); cent=np.zeros((nr,M),np.float32); rel=np.zeros((nr,M),np.float32); rho=np.empty(nr,np.float32)24    for u,(j,a,b) in enumerate(spans):25        rowt=np.asarray(idx.ct[j]); ok=rowt!=65535; tids=rowt[ok].astype(np.int32,copy=False); cent[u,tids]=np.asarray(idx.cv[j])[ok]26        ra=int(idx.rp[j]); rb=int(idx.rp[j+1]); rel[u,np.asarray(idx.ri[ra:rb],np.int32)]=np.asarray(idx.rv[ra:rb]); rho[u]=rd[j]27    rslot=np.concatenate([np.full(b-a,u,dtype=np.uint8) for u,(j,a,b) in enumerate(spans)])28    hc,tc,cc=m.score_memberships_local(rslot,mm,rt,sb,qd,rho,cent,rel)29    ud,inv=np.unique(docs,return_inverse=True)30    head=np.bincount(inv,weights=hc,minlength=len(ud)).astype(np.float32)31    local=np.bincount(inv,weights=tc,minlength=len(ud)).astype(np.float32)32    cons=np.bincount(inv,weights=cc,minlength=len(ud)).astype(np.float32)33    # lexical vectors independent of lambda34    lexvec=np.zeros(M,np.float32); lexvec[q.indices]=idx.idf[q.indices]; semvec=np.zeros(M,np.float32)35    for t,amp in zip(q.indices,q.data):36        a,b=idx.A.indptr[t],idx.A.indptr[t+1]; nb=idx.A.indices[a:b][:m.SEMK]; sv=idx.A.data[a:b][:m.SEMK]; semvec[nb]+=float(amp)*sv*idx.idf[nb]37    ho=np.argsort(head)[::-1][:max(HGRID)] if max(HGRID)>0 else np.empty(0,np.int64)38    return q,ud,head,local,cons,lexvec,semvec,ho,len(docs)39 40runs={(la,h):{} for la in LAMBDAS for h in HGRID}; pool_hit={la:0 for la in LAMBDAS}; rel_den=0; route_hit=0; times=[]; cand_counts=[]41# warm JIT42_ = components(texts[ids[0]])43for z,qid in enumerate(ids):44    t0=time.perf_counter(); c=components(texts[qid]);45    if c is None:46        for key in runs:runs[key][qid]=[]47        continue48    q,ud,head,local,cons,lexvec,semvec,ho,nmem=c; cand_counts.append(len(ud))49    rels=[int(d) for d,r in qrels[qid].items() if r>0]; rel_den+=len(rels)50    for d in rels:51        k=np.searchsorted(ud,d); route_hit+=int(k<len(ud) and int(ud[k])==d)52    # select each lambda pool; union for one support pass53    pools={}; union=[]54    for la in LAMBDAS:55        tail=local+np.float32(la)*cons; want=min(len(tail),P+max(HGRID)+8)56        if len(tail)>want:57            ci=np.argpartition(tail,-want)[-want:]; oo=ci[np.argsort(tail[ci])[::-1]]58        else: oo=np.argsort(tail)[::-1]59        pools[la]=(ud[oo],tail[oo]); union.append(ud[oo])60    udocs=np.unique(np.concatenate(union)); lx,sm=m.score_support_pool(udocs,idx.sup_ip,idx.sup_ids,lexvec,semvec,idx.dl,idx.avgdl)61    # udocs sorted, so searchsorted maps pool docs62    for la in LAMBDAS:63        docs,ts=pools[la]64        # pool recall before final rerank at h=0 definition65        p0=docs[:P]66        for d in rels: pool_hit[la]+=int(np.any(p0==d))67        pos=np.searchsorted(udocs,docs); lxx=lx[pos]; smm=sm[pos]68        for h in HGRID:69            frozen=ud[ho[:h]] if h else np.empty(0,np.uint32); fs=set(map(int,frozen.tolist()))70            keep=np.asarray([int(d) not in fs for d in docs],bool); dd=docs[keep][:P]; tt=ts[keep][:P]; ll=lxx[keep][:P]; ss=smm[keep][:P]71            fin=m.zscore(tt)+m.LAMBDA_LEX*m.zscore(ll)+m.LAMBDA_SEM*m.zscore(ss); oo=np.argsort(fin)[::-1]; rank=np.concatenate([frozen,dd[oo]])[:100]72            runs[(la,h)][qid]=[int(x) for x in rank]73    times.append((time.perf_counter()-t0)*1000)74    if (z+1)%100==0: print('q',z+1,'median_ms',float(np.median(times)),'route',route_hit/max(1,rel_den),flush=True)75 76rows=[]77for la in LAMBDAS:78    for h in HGRID:79        met=m.eval_run(runs[(la,h)],qrels); met['pool_relevant_recall']=pool_hit[la]/rel_den; rows.append({'lambda_M':la,'h':h,**met}); print('LAM',la,'H',h,met,'pool',pool_hit[la]/rel_den,flush=True)80best=max(rows,key=lambda r:(r['nDCG@10'],r['MRR@10'],r['R@100']))81out={'protocol':'uniform-1M geometry; deterministic 1000 TRAIN validation; joint diagnostic sweep lambda_M and h','route_relevant_recall':route_hit/rel_den,'timing_median_ms':float(np.median(times)),'avg_candidate_docs':float(np.mean(cand_counts)),'rows':rows,'best':best}82json.dump(out,open(WORK/'lambdaM_uniform1m_diag.json','w'),indent=2); print('BEST',best,flush=True)83