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_encode_full.py112 linesDownload Raw Back to msmarco_scale
1from __future__ import annotations2import gzip,json,pickle,time,re,gc,os3from pathlib import Path4import numpy as np5from scipy import sparse6from sklearn.feature_extraction.text import CountVectorizer7from sklearn.preprocessing import normalize8from numba import njit, prange, set_num_threads9 10ROOT=Path('/mnt/data'); WORK=ROOT/'msmarco_scale_work'; GEOM=WORK/'geometry_1m'; IDX=WORK/'full_index'; IDX.mkdir(parents=True,exist_ok=True)11N=8_841_823; M=50_000; F=4; S=16; SENT=np.uint16(65535)12set_num_threads(5)13TOKEN_RE=re.compile(r'(?u)\b\w\w+\b')14 15def shard_path(i):16 hits=list(ROOT.glob(f'corpus_{i:04d}.jsonl*.gz')); assert len(hits)==1,(i,hits); return hits[0]17 18def load_vocab():19 with gzip.open(WORK/'final_vocab_50k.pkl.gz','rb') as g: z=pickle.load(g)20 terms=z['terms'].tolist(); idf=np.asarray(z['idf'],np.float32); return terms,idf,{t:i for i,t in enumerate(terms)}21 22@njit(cache=False)23def lookup_center(ct,cv,t):24 lo=0; hi=ct.size25 while lo<hi:26  mid=(lo+hi)//2; x=ct[mid]27  if x==65535 or x>=t: hi=mid28  else: lo=mid+129 if lo<ct.size and ct[lo]==t: return cv[lo]30 return 0.031 32@njit(parallel=True,cache=False)33def encode_kernel(indptr,indices,data,center_terms,center_values):34 n=indptr.size-135 branches=np.full((n,F),SENT,np.uint16); mem=np.zeros((n,F),np.float32)36 rt=np.full((n,F,S),SENT,np.uint16); signbits=np.zeros((n,F),np.uint16)37 for d in prange(n):38  a=indptr[d]; b=indptr[d+1]39  # top4 document coordinates40  tv=np.zeros(F,np.float32); tt=np.full(F,SENT,np.uint16)41  for p in range(a,b):42   v=data[p]; t=np.uint16(indices[p]); pos=F43   for r in range(F):44    if v>tv[r]: pos=r; break45   if pos<F:46    for r in range(F-1,pos,-1): tv[r]=tv[r-1]; tt[r]=tt[r-1]47    tv[pos]=v; tt[pos]=t48  den=0.049  for s in range(F): den+=tv[s]50  if den<=0: continue51  for s in range(F): branches[d,s]=tt[s]; mem[d,s]=tv[s]/den52  # residual codes per fuzzy branch53  for sl in range(F):54   j=int(tt[sl])55   if j==65535: continue56   best=np.zeros(S,np.float32); bt=np.full(S,SENT,np.uint16); bp=np.zeros(S,np.uint8)57   for p in range(a,b):58    t=np.uint16(indices[p]); r=data[p]-lookup_center(center_terms[j],center_values[j],t); ar=abs(r)59    mi=0; mv=best[0]60    for q in range(1,S):61     if best[q]<mv: mi=q; mv=best[q]62    if ar>mv:63     best[mi]=ar; bt[mi]=t; bp[mi]=1 if r>=0 else 064   # sort descending to stabilize65   for x in range(S):66    mx=x67    for y in range(x+1,S):68     if best[y]>best[mx]: mx=y69    if mx!=x:70     z=best[x]; best[x]=best[mx]; best[mx]=z71     zt=bt[x]; bt[x]=bt[mx]; bt[mx]=zt72     zp=bp[x]; bp[x]=bp[mx]; bp[mx]=zp73   bits=np.uint16(0)74   for q in range(S):75    rt[d,sl,q]=bt[q]76    if bt[q]!=SENT and bp[q]: bits |= np.uint16(1<<q)77   signbits[d,sl]=bits78 return branches,mem,rt,signbits79 80if __name__=='__main__':81 terms,idf,vocab=load_vocab(); center_terms=np.load(GEOM/'center_terms.npy',mmap_mode='r'); center_values=np.load(GEOM/'center_values.npy',mmap_mode='r')82 # Global disk-backed arrays83 branches=np.memmap(IDX/'branches.u16',dtype=np.uint16,mode='w+',shape=(N,F)); branches[:]=SENT84 memberships=np.memmap(IDX/'memberships.f32',dtype=np.float32,mode='w+',shape=(N,F)); memberships[:]=085 res_terms=np.memmap(IDX/'res_terms.u16',dtype=np.uint16,mode='w+',shape=(N,F,S)); res_terms[:]=SENT86 signbits=np.memmap(IDX/'signbits.u16',dtype=np.uint16,mode='w+',shape=(N,F)); signbits[:]=087 doc_lengths=np.memmap(IDX/'doc_lengths.u16',dtype=np.uint16,mode='w+',shape=(N,)); doc_lengths[:]=088 cv=CountVectorizer(vocabulary=vocab,lowercase=True,token_pattern=r'(?u)\b\w\w+\b',dtype=np.int32)89 total_len=0; t_all=time.time(); offset=090 for sid in range(36):91  t=time.time(); texts=[]; lens=[]92  with gzip.open(shard_path(sid),'rt',encoding='utf-8') as f:93   for line in f:94    o=json.loads(line); tx=((o.get('title') or '')+' '+(o.get('text') or '')).strip(); texts.append(tx); lens.append(min(65535,len(TOKEN_RE.findall(tx.lower()))))95  n=len(texts); X=cv.transform(texts).tocsr().astype(np.float32); X.data *= idf[X.indices]; normalize(X,norm='l2',axis=1,copy=False); X.sort_indices()96  br,mm,rt,sb=encode_kernel(X.indptr.astype(np.int64),X.indices.astype(np.int32),X.data.astype(np.float32),center_terms,center_values)97  sl=slice(offset,offset+n); branches[sl]=br; memberships[sl]=mm; res_terms[sl]=rt; signbits[sl]=sb; doc_lengths[sl]=np.asarray(lens,np.uint16); total_len += int(np.sum(lens,dtype=np.int64))98  # whole-document binary support, one pair of files per corpus shard99  X.indices.astype(np.uint16).tofile(IDX/f'support_{sid:04d}.u16')100  X.indptr.astype(np.uint32).tofile(IDX/f'support_indptr_{sid:04d}.u32')101  with open(IDX/f'shard_{sid:04d}.json','w') as f: json.dump({'offset':offset,'n':n,'nnz':int(X.nnz),'seconds':time.time()-t},f)102  offset+=n; branches.flush(); memberships.flush(); res_terms.flush(); signbits.flush(); doc_lengths.flush()103  print(f'[{sid+1:02d}/36] n={n:,} nnz={X.nnz:,} offset={offset:,} sec={time.time()-t:.1f}',flush=True)104  del texts,lens,X,br,mm,rt,sb; gc.collect()105 assert offset==N,(offset,N)106 avg=total_len/N107 # Build branch postings by a global stable sort over 35.4M uint16 branch IDs.108 print('building branch postings...',flush=True); t=time.time(); flat=np.memmap(IDX/'branches.u16',dtype=np.uint16,mode='r',shape=(N*F,)); order=np.argsort(flat,kind='stable'); sorted_br=flat[order]; nvalid=int(np.searchsorted(sorted_br,SENT,side='left'))109 bo=np.memmap(IDX/'branch_order.u32',dtype=np.uint32,mode='w+',shape=(nvalid,)); bo[:]=order[:nvalid].astype(np.uint32); bo.flush(); counts=np.bincount(sorted_br[:nvalid].astype(np.int64),minlength=M); offsets=np.zeros(M+1,np.uint64); np.cumsum(counts,dtype=np.uint64,out=offsets[1:]); np.save(IDX/'branch_offsets.npy',offsets); print('postings valid memberships',nvalid,'sec',time.time()-t,flush=True)110 with open(IDX/'meta.json','w') as f: json.dump({'N':N,'M':M,'F':F,'S':S,'avg_doc_length':avg,'build_seconds':time.time()-t_all,'geometry':'geometry_1m_fullcorpus_vocab'},f,indent=2)111 print('FULL INDEX DONE avgdl',avg,'total sec',time.time()-t_all,flush=True)112