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