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_extract_structural_features.py85 linesDownload Raw Back to msmarco_scale
1from __future__ import annotations2import sys,time,json3from pathlib import Path4import numpy as np,pandas as pd5from numba import njit,prange,set_num_threads6sys.path.insert(0,'/mnt/data')7import msmarco_best_tail_core as b8import msmarco_full_search_uniform1m as m9ROOT=m.ROOT; WORK=m.WORK; idx=b.idx; M=m.M; P=200010OUT=WORK/'structural_fusion'; OUT.mkdir(exist_ok=True)11set_num_threads(5)12 13@njit(parallel=True,cache=False)14def support_features(cand_docs,ip,ids,lexvec,semvec,dl,avgdl):15 n=len(cand_docs); lex=np.zeros(n,np.float32); rawlex=np.zeros(n,np.float32); sem=np.zeros(n,np.float32); qcount=np.zeros(n,np.uint8); lenfac=np.zeros(n,np.float32)16 for z in prange(n):17  d=int(cand_docs[z]); a=int(ip[d]); bb=int(ip[d+1]); raw=0.0; sm=0.0; cnt=018  for k in range(a,bb):19   t=int(ids[k]); v=lexvec[t]20   if v>0:21    raw += v; cnt += 122   sm += semvec[t]23  denom=(1.0-m.LENGTH_B)+m.LENGTH_B*(float(dl[d])/avgdl)24  rawlex[z]=raw; lex[z]=raw/(denom if denom>0 else 1.0); sem[z]=sm; qcount[z]=cnt; lenfac[z]=denom25 return lex,rawlex,sem,qcount,lenfac26 27def topk(score,k):28 n=len(score); k=min(k,n)29 if n<=k:return np.argsort(score)[::-1]30 ii=np.argpartition(score,-k)[-k:]; return ii[np.argsort(score[ii])[::-1]]31 32# exact validation IDs33tr=pd.read_csv(ROOT/'train.tsv',sep='\t',usecols=['query-id']); uq=np.unique(tr['query-id'].to_numpy()); del tr34rng=np.random.default_rng(20260815); qids=[str(x) for x in rng.choice(uq,size=1000,replace=False)]35texts=m.load_query_texts(qids)36# allocate fixed P arrays37shape=(len(qids),P)38docsO=np.zeros(shape,np.uint32); valid=np.zeros(len(qids),np.int32)39features={k:np.zeros(shape,dtype) for k,dtype in [40 ('geom',np.float32),('cons',np.float32),('tail',np.float32),('lex',np.float32),('rawlex',np.float32),('sem',np.float32),('qcount',np.float32),('lenfac',np.float32),('branch_count',np.float32),('max_cons',np.float32),('max_geom',np.float32)]}41qterms=np.zeros(len(qids),np.int16)42prep=[]43# warmup44_=support_features(np.array([0],np.uint32),idx.sup_ip,idx.sup_ids,np.zeros(M,np.float32),np.zeros(M,np.float32),idx.dl,idx.avgdl)45for qi,qid in enumerate(qids):46 t0=time.perf_counter(); text=texts[qid]47 q=idx.query_vec(text); qd=np.zeros(M,np.float32); qd[q.indices]=q.data; qterms[qi]=len(q.indices)48 rterms,rd=idx.route(q)49 spans=[(int(j),int(idx.offs[j]),int(idx.offs[j+1])) for j in rterms if idx.offs[j+1]>idx.offs[j]]50 if not spans: continue51 dmem=np.concatenate([np.asarray(idx.pd[a:bb]) for j,a,bb in spans]).astype(np.uint32,copy=False)52 mm=np.concatenate([np.asarray(idx.pm[a:bb]) for j,a,bb in spans]).astype(np.float32,copy=False)53 rt=np.concatenate([np.asarray(idx.pr[a:bb]) for j,a,bb in spans]).astype(np.uint16,copy=False)54 sb=np.concatenate([np.asarray(idx.ps[a:bb]) for j,a,bb in spans]).astype(np.uint16,copy=False)55 nr=len(spans); cent=np.zeros((nr,M),np.float32); rel=np.zeros((nr,M),np.float32); rho=np.empty(nr,np.float32)56 for u,(j,a,bb) in enumerate(spans):57  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]58  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]59 rslot=np.concatenate([np.full(bb-a,u,dtype=np.uint8) for u,(j,a,bb) in enumerate(spans)])60 base,sig,consmem=b.score_components(rslot,mm,rt,sb,qd,rho,cent,rel)61 ud,inv=np.unique(dmem,return_inverse=True)62 geom=np.bincount(inv,weights=base*np.power(sig,b.GAMMA,dtype=np.float32),minlength=len(ud)).astype(np.float32)63 cons=np.bincount(inv,weights=consmem,minlength=len(ud)).astype(np.float32)64 tail=geom+b.LAM*cons65 bc=np.bincount(inv,minlength=len(ud)).astype(np.float32)66 # max per-document membership-level contributions, preserved as structural evidence67 mxcon=np.full(len(ud),-np.inf,np.float32); np.maximum.at(mxcon,inv,consmem); mxcon[~np.isfinite(mxcon)]=068 memgeom=base*np.power(sig,b.GAMMA,dtype=np.float32); mxg=np.full(len(ud),-np.inf,np.float32); np.maximum.at(mxg,inv,memgeom); mxg[~np.isfinite(mxg)]=069 # support features for all routed docs, needed for eta=1 selection70 lexvec=np.zeros(M,np.float32); lexvec[q.indices]=idx.idf[q.indices]71 semvec=np.zeros(M,np.float32)72 for tt,amp in zip(q.indices,q.data):73  a,bb=idx.A.indptr[tt],idx.A.indptr[tt+1]; nb=idx.A.indices[a:bb][:m.SEMK]; sv=idx.A.data[a:bb][:m.SEMK]; semvec[nb]+=float(amp)*sv*idx.idf[nb]74 lx,rlx,sm,qc,lf=support_features(ud,idx.sup_ip,idx.sup_ids,lexvec,semvec,idx.dl,idx.avgdl)75 sel=topk(m.zscore(tail)+m.zscore(lx),P) # locked eta=176 k=len(sel); valid[qi]=k; docsO[qi,:k]=ud[sel]77 for name,arr in [('geom',geom),('cons',cons),('tail',tail),('lex',lx),('rawlex',rlx),('sem',sm),('qcount',qc.astype(np.float32)),('lenfac',lf),('branch_count',bc),('max_cons',mxcon),('max_geom',mxg)]: features[name][qi,:k]=arr[sel]78 prep.append((time.perf_counter()-t0)*1000)79 if (qi+1)%100==0: print('q',qi+1,'median',float(np.median(prep)),flush=True)80 81save={'qids':np.asarray(qids),'valid':valid,'docs':docsO,'qterms':qterms}|features82np.savez_compressed(OUT/'fixed_eta1_structural_features.npz',**save)83meta={'protocol':'same deterministic 1000 TRAIN validation and locked eta=1 P=2000 pools; features use only current index','features':list(features),'median_prepare_ms':float(np.median(prep)),'p95_prepare_ms':float(np.percentile(prep,95))}84json.dump(meta,open(OUT/'structural_feature_meta.json','w'),indent=2); print('DONE',meta,flush=True)85