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.
0331
1from __future__ import annotations2import sys,time,json,math,argparse,gc3from pathlib import Path4import numpy as np5from numba import njit,set_num_threads6set_num_threads(5)7sys.path.insert(0,'/mnt/data/exp')8import eval_scale_rag as e9from opt_two_variants import FastScaleIndex, agg_stamp10 11@njit(cache=True)12def compact_pool_evidence_and_code(docs,rslot,mem,rt,sbits,ev,spanbr,rho,poolmap,poolstamp,token,nmax,qct,qcw,qcbits):13 cap=nmax*e.F14 mp=np.empty(cap,np.int32);mb=np.empty(cap,np.int32);me=np.empty(cap,np.float32);cm=np.zeros(nmax,np.float32);n=015 for z in range(len(docs)):16 d=np.int64(docs[z])17 if poolstamp[d]!=token: continue18 pi=poolmap[d]19 if pi<0 or pi>=nmax: continue20 u=np.int64(rslot[z]); b=spanbr[u]21 mp[n]=pi;mb[n]=b;me[n]=ev[z];n+=122 bits=np.uint16(sbits[z]);s=0.;den=0.23 for a in range(e.S):24 qt=np.int64(qct[u,a])25 if qt==65535: continue26 w=qcw[u,a];den+=w27 for r in range(e.S):28 if np.int64(rt[z,r])==qt:29 qs=1 if ((qcbits[u]>>a)&1)!=0 else -130 ds=1 if ((bits>>r)&1)!=0 else -131 s += w*(1. if qs==ds else -1.)32 break33 if den>0.: cm[pi]+=mem[z]*rho[u]*(s/den)34 return mp[:n],mb[:n],me[:n],cm35 36def qcodes_sparse(idx,spans,q,qd,rel):37 nr=len(spans);qt=np.full((nr,e.S),65535,np.uint16);qw=np.zeros((nr,e.S),np.float32);qb=np.zeros(nr,np.uint16)38 qterms=np.asarray(q.indices,np.int32)39 for u,(j,_,_) in enumerate(spans):40 rowt=np.asarray(idx.ct[j]);ok=rowt!=e.SENT;cids=rowt[ok].astype(np.int32,copy=False)41 cvals=np.asarray(idx.cv[j],np.float32)[ok]42 cand=np.unique(np.concatenate([cids,qterms]))43 # center lookup only for <=80 coordinates44 cvmap={int(t):float(v) for t,v in zip(cids,cvals)}45 diff=np.asarray([float(qd[t])-cvmap.get(int(t),0.0) for t in cand],np.float32);aa=np.abs(diff)46 if len(cand)>e.S:47 ii=np.argpartition(aa,-e.S)[-e.S:];ii=ii[np.argsort(aa[ii])[::-1]]48 else:ii=np.argsort(aa)[::-1]49 for a,k in enumerate(ii[:e.S]):50 t=int(cand[k]);qt[u,a]=t;rr=float(rel[u,t]) if rel[u,t]!=0 else 1.0;qw[u,a]=float(aa[k])*rr51 if diff[k]>=0: qb[u]|=np.uint16(1<<a)52 return qt,qw,qb53 54class FinalOptimizedIndex(FastScaleIndex):55 def search_optimized(self,text,P,gate=None,wcode=.25):56 t0=time.perf_counter();self.token+=1;token=self.token57 q=self.qvec(text)58 if q.nnz==0:return [],{'total_ms':(time.perf_counter()-t0)*1000,'route_docs':0,'gate_docs':0}59 qd=np.zeros(e.M,np.float32);qd[q.indices]=q.data;rterms,rd=self.route(q)60 spans=[(int(j),int(self.offs[j]),int(self.offs[j+1])) for j in rterms if self.offs[j+1]>self.offs[j]]61 if not spans:return [],{'total_ms':(time.perf_counter()-t0)*1000,'route_docs':0,'gate_docs':0}62 docs=np.concatenate([np.asarray(self.pd[a:b]) for j,a,b in spans]).astype(np.uint32,copy=False)63 mem=np.concatenate([np.asarray(self.pm[a:b]) for j,a,b in spans]).astype(np.float32,copy=False)64 rt=np.concatenate([np.asarray(self.pr[a:b]) for j,a,b in spans]).astype(np.uint16,copy=False)65 sb=np.concatenate([np.asarray(self.ps[a:b]) for j,a,b in spans]).astype(np.uint16,copy=False)66 nr=len(spans);cent=np.zeros((nr,e.M),np.float32);rel=np.zeros((nr,e.M),np.float32);rho=np.empty(nr,np.float32);spanbr=np.empty(nr,np.int32);rs=[]67 for u,(j,a,b) in enumerate(spans):68 spanbr[u]=j;rowt=np.asarray(self.ct[j]);ok=rowt!=e.SENT;ids=rowt[ok].astype(np.int32,copy=False);cent[u,ids]=np.asarray(self.cv[j])[ok]69 ra=int(self.rp[j]);rb=int(self.rp[j+1]);rel[u,np.asarray(self.ri[ra:rb],np.int32)]=np.asarray(self.rv[ra:rb]);rho[u]=rd[j];rs.append(np.full(b-a,u,dtype=np.uint8))70 rslot=np.concatenate(rs);ev,cons=e.score_memberships_local(rslot,mem,rt,sb,qd,rho,cent,rel)71 # Optimization 1a: O(K) stamp aggregation instead of np.unique sort.72 n=agg_stamp(docs,ev,cons,self.stamp,self.taila,self.consa,self.touched,token)73 ud=np.asarray(self.touched[:n],np.uint32).copy();tail=(self.taila[ud]+e.LAM_M*self.consa[ud]).astype(np.float32)74 # Optimization 1b: small tail gate before whole-chunk lexical scan.75 if gate is None: gate=min(len(ud),max(10000,40*int(P)))76 if len(ud)>gate:77 gi=e.topk_sorted(tail,gate);cand=ud[gi];ct=tail[gi]78 else:cand=ud;ct=tail79 qlex=np.zeros(e.M,np.float32);qlex[q.indices]=self.idf[q.indices];lex1=e.score_pre(cand,self.ip,self.ids,qlex,self.dl,self.avgdl)80 pre=e.zscore(ct)+e.zscore(lex1);sel=e.topk_sorted(pre,P);pooldocs=cand[sel];pooltail=ct[sel]81 # exact pool map with stamps82 self.poolstamp[pooldocs]=token;self.poolmap[pooldocs]=np.arange(len(pooldocs),dtype=np.int32)83 # Optimization 2: query-conditioned full 16-coordinate residual-code match.84 qct,qcw,qcb=qcodes_sparse(self,spans,q,qd,rel)85 mp,mb,me,cm=compact_pool_evidence_and_code(docs,rslot,mem,rt,sb,ev,spanbr,rho,self.poolmap,self.poolstamp,token,len(pooldocs),qct,qcw,qcb)86 semv=np.zeros(e.M,np.float32)87 for t,amp in zip(q.indices,q.data):88 a,b=self.A.indptr[t],self.A.indptr[t+1];nb=self.A.indices[a:b][:e.SEMK];sv=self.A.data[a:b][:e.SEMK]89 if len(nb):semv[nb]+=float(amp)*sv*self.idf[nb]90 lex2v=np.zeros(e.M,np.float32);lex2v[q.indices]=self.idf[q.indices]**2;qmask=np.zeros(e.M,np.float32);qmask[q.indices]=1.;rarem=np.zeros(e.M,np.float32);rareidx=q.indices[np.argsort(self.idf[q.indices])[::-1]][:e.RAREK];rarem[rareidx]=1.91 p={'q':q,'pooldocs':pooldocs,'pooltail':pooltail,'mem_pool':mp,'mem_branch':mb,'mem_ev':me,'semv':semv,'lex2v':lex2v,'qmask':qmask,'rarem':rarem,'ressem_signed':cm,'route_docs':ud}92 r,rtm=self.rank_variant(p,P,wres=wcode,absres=False)93 return r,{'total_ms':(time.perf_counter()-t0)*1000,'rank_ms':rtm['rank_ms'],'route_docs':len(ud),'gate_docs':len(cand),'pool_size':len(pooldocs)}94 95def full_run(dataset,work,qpath,rpath,P,outfile,gate=None,wcode=.25,start=0,limit=None):96 gc.disable()97 idx=FinalOptimizedIndex(work,dataset);qr=e.load_qrels(rpath);qids=sorted(qr);qs=e.load_queries(qpath,qids);sub=qids[start:] if limit is None else qids[start:start+limit]98 # compile/warm99 if sub:idx.search_optimized(qs[sub[0]],P,gate,wcode)100 run={};times=[];routes=[];gates=[]101 for z,qid in enumerate(sub):102 r,t=idx.search_optimized(qs[qid],P,gate,wcode);run[qid]=r;times.append(t['total_ms']);routes.append(t['route_docs']);gates.append(t['gate_docs'])103 if (z+1)%250==0:104 print(dataset,'done',z+1,'med_ms',float(np.median(times)),'route_med',float(np.median(routes)),flush=True)105 Path(outfile+'.checkpoint').write_text(json.dumps({'done':z+1,'run':run,'times':times,'routes':routes,'gates':gates}))106 qrs={q:qr[q] for q in sub};m=e.metrics_for(run,qrs);a=np.asarray(times)107 m.update({'median_ms':float(np.median(a)),'p95_ms':float(np.percentile(a,95)),'mean_ms':float(np.mean(a)),'qps':float(1000/np.mean(a)),'median_route_docs':float(np.median(routes)),'median_gate_docs':float(np.median(gates)),'P':P,'wcode':wcode,'gate_rule':'min(route,max(10000,40P))' if gate is None else gate})108 out={'dataset':dataset,'metrics':m,'run':run};Path(outfile).write_text(json.dumps(out,indent=2));print(json.dumps(m,indent=2));return out109 110if __name__=='__main__':111 ap=argparse.ArgumentParser();ap.add_argument('dataset',choices=['nq','hotpot']);ap.add_argument('--limit',type=int,default=None);ap.add_argument('--out',required=True);a=ap.parse_args()112 if a.dataset=='nq':full_run('nq','/mnt/data/exp/nq_work','/mnt/data/queries(3).jsonl','/mnt/data/test(3).tsv',100,a.out,limit=a.limit)113 else:full_run('hotpot','/mnt/data/exp/hotpot_work','/mnt/data/queries(2).jsonl','/mnt/data/test(2).tsv',500,a.out,limit=a.limit)114 