exp-chat-rag / src /retrieval.py
Arnic's picture
fix: dummy @spaces.GPU placeholder — no GPU allocation, no quota consumption (same pattern as neurorag-2026)
acbca58
Raw
History Blame Contribute Delete
4.82 kB
"""Embedding and FAISS retrieval — REVISED v3."""
import os
import pickle
from typing import List, Dict
import numpy as np
import faiss
from sentence_transformers import SentenceTransformer
from config import (
EMBEDDING_MODEL, EMBEDDING_DIM, INDEX_DIR, TOP_K, TOP_K_FETCH, get_boosted_sections
)
class PortfolioRetriever:
"""Dual-tier retriever: FAISS semantic + keyword fallback."""
def __init__(self, build: bool = False):
self.encoder = SentenceTransformer(EMBEDDING_MODEL, device="cpu")
self.index = None
self.chunks = []
if build or not self._index_exists():
self._build_index()
else:
self._load_index()
def _index_exists(self) -> bool:
return os.path.exists(os.path.join(INDEX_DIR, "faiss.index"))
def _build_index(self):
"""Build FAISS index from KB chunks."""
from ingestion import parse_kb
print("Building FAISS index...")
os.makedirs(INDEX_DIR, exist_ok=True)
self.chunks = parse_kb()
texts = [c["text"] for c in self.chunks]
print(f"Encoding {len(texts)} chunks with {EMBEDDING_MODEL}...")
embeddings = self.encoder.encode(
texts,
show_progress_bar=True,
normalize_embeddings=True
)
self.index = faiss.IndexFlatIP(EMBEDDING_DIM)
self.index.add(np.array(embeddings).astype('float32'))
faiss.write_index(self.index, os.path.join(INDEX_DIR, "faiss.index"))
with open(os.path.join(INDEX_DIR, "chunks.pkl"), 'wb') as f:
pickle.dump(self.chunks, f)
print(f"Index saved: {len(self.chunks)} chunks, {EMBEDDING_DIM}d")
def _load_index(self):
"""Load pre-built index."""
print("Loading pre-built FAISS index...")
self.index = faiss.read_index(os.path.join(INDEX_DIR, "faiss.index"))
with open(os.path.join(INDEX_DIR, "chunks.pkl"), 'rb') as f:
self.chunks = pickle.load(f)
print(f"Loaded {len(self.chunks)} chunks")
def retrieve(self, query: str, persona: str = "general", k: int = TOP_K) -> List[Dict]:
"""Retrieve relevant chunks with keyword fallback boosting."""
query_vec = self.encoder.encode([query], normalize_embeddings=True)
query_vec = np.array(query_vec).astype('float32')
scores, indices = self.index.search(query_vec, min(TOP_K_FETCH, len(self.chunks)))
results = []
seen_ids = set()
boosted_sections = get_boosted_sections(query)
for score, idx in zip(scores[0], indices[0]):
if idx < 0 or idx >= len(self.chunks):
continue
chunk = self.chunks[idx].copy()
chunk_id = chunk["id"]
if chunk_id in seen_ids:
continue
seen_ids.add(chunk_id)
chunk["score"] = float(score)
if chunk["section_type"] in boosted_sections:
chunk["score"] *= 1.5
if persona == "recruiter" and chunk["section_type"] in [
"identity", "education", "skills", "skills_summary",
"experience", "experience_summary_v2", "certifications",
"kaggle_summary"
]:
chunk["score"] *= 1.15
elif persona == "technical" and chunk["section_type"] in [
"architecture_deep_dive_axiomis", "projects", "skills", "skills_summary",
"publications", "publications_summary", "philosophy",
"clinical_ai_projects", "marketing_ai_projects", "supply_chain_projects",
"energy_projects", "computational_biology_projects", "finance_regtech_projects",
"education_projects", "neuro_symbolic_projects", "projects_master_list"
]:
chunk["score"] *= 1.15
results.append(chunk)
results.sort(key=lambda x: x["score"], reverse=True)
final_results = []
section_counts = {}
for r in results:
st = r["section_type"]
section_counts[st] = section_counts.get(st, 0) + 1
if section_counts[st] <= 3:
final_results.append(r)
if len(final_results) >= k:
break
return final_results
def get_stats(self) -> Dict:
from collections import Counter
dist = Counter(c["section_type"] for c in self.chunks)
return {
"total_chunks": len(self.chunks),
"embedding_dim": EMBEDDING_DIM,
"encoder": EMBEDDING_MODEL,
"section_distribution": dict(dist)
}