Spaces:
Sleeping
Sleeping
File size: 4,382 Bytes
f36c047 dfbb2bc | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 | import logging
from typing import List
from langchain_pinecone import PineconeVectorStore
from langchain_core.documents import Document
from langchain_core.messages import HumanMessage
from flashrank import Ranker, RerankRequest
from config import *
# Logging setup
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
# ⚡ OPTIMIZATION 1: Use FlashRank (Ultra-fast CPU Re-ranking)
ranker = Ranker(model_name="ms-marco-MiniLM-L-12-v2", cache_dir="/tmp/flashrank_cache")
def reciprocal_rank_fusion(results: List[List[Document]], k=60):
fused_scores = {}
doc_map = {}
for docs in results:
for rank, doc in enumerate(docs):
doc_id = doc.metadata.get("doc_id")
if doc_id not in doc_map: doc_map[doc_id] = doc
if doc_id not in fused_scores: fused_scores[doc_id] = 0
fused_scores[doc_id] += 1 / (rank + k)
reranked_ids = sorted(fused_scores, key=fused_scores.get, reverse=True)
return [doc_map[doc_id] for doc_id in reranked_ids]
def rerank_documents(query: str, docs: List[Document], top_n=5):
"""
Optimized Re-ranking using FlashRank (runs in milliseconds).
"""
if not docs: return []
try:
# Prepare format for FlashRank
passages = [
{"id": str(i), "text": doc.page_content, "meta": doc.metadata}
for i, doc in enumerate(docs)
]
# Rerank
rerank_request = RerankRequest(query=query, passages=passages)
results = ranker.rerank(rerank_request)
# Convert back to Document objects
final_docs = []
for res in results[:top_n]:
final_docs.append(Document(page_content=res["text"], metadata=res["meta"]))
return final_docs
except Exception as e:
logger.error(f"FlashRank failed: {e}")
return docs[:top_n] # Fallback
def run_advanced_rag(query: str, session_id: str, bm25_retriever, doc_store, chat_history):
# ⚡ OPTIMIZATION 2: Skip LLM Query Decomposition (Saves 3-5s)
# We treat the user query as the only query.
queries = [query]
# ⚡ OPTIMIZATION 3: Standard Similarity Search (Faster than MMR)
vectorstore = PineconeVectorStore.from_existing_index(
index_name=INDEX_NAME, embedding=get_embeddings(), namespace=session_id
)
dense_retriever = vectorstore.as_retriever(search_type="similarity", search_kwargs={"k": 3})
all_docs = []
# Hybrid Search
for q in queries:
dense_docs = dense_retriever.invoke(f"query: {q}")
sparse_docs = bm25_retriever.invoke(q)
# Fuse results
all_docs.extend(reciprocal_rank_fusion([dense_docs, sparse_docs]))
# Deduplicate by ID
unique_docs = {d.metadata["doc_id"]: d for d in all_docs}
# Re-rank (Fast)
final_docs = rerank_documents(query, list(unique_docs.values()), top_n=5)
# Context Construction
context_text = ""
retrieved_images = []
seen_imgs = set()
for i, doc in enumerate(final_docs):
# Fetch heavy content from in-memory store
heavy = doc_store.get_chunk(doc.metadata["doc_id"])
context_text += f"\n--- Source {i+1} ---\n{heavy.get('raw_text', '')}\n"
for t in heavy.get('tables', []):
context_text += f"[Table]: {t}\n"
# Collect images
for img in heavy.get('images', []):
if img not in seen_imgs:
seen_imgs.add(img)
retrieved_images.append(img)
# Final Generation
llm = get_llm()
# Limit history to last 2 turns to save tokens/time
hist_text = "\n".join([f"{m.type}: {m.content}" for m in chat_history[-2:]])
prompt = f"""
Answer the user question based on the provided context.
CHAT HISTORY: {hist_text}
CONTEXT: {context_text[:5000]}
QUESTION: {query}
"""
msg_content = [{"type": "text", "text": prompt}]
# ⚡ OPTIMIZATION 4: Limit Images to top 2
for b64 in retrieved_images[:2]:
if "," in b64: b64 = b64.split(",")[1]
msg_content.append({
"type": "image_url",
"image_url": {"url": f"data:image/jpeg;base64,{b64}"}
})
response = llm.invoke([HumanMessage(content=msg_content)])
return response.content |