aruntandra commited on
Commit
252c8a3
Β·
verified Β·
1 Parent(s): 8a79f11

Upload 2 files

Browse files
Files changed (2) hide show
  1. app.py +1778 -0
  2. requirements.txt +15 -0
app.py ADDED
@@ -0,0 +1,1778 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # ── Cell 2: Imports ────────────────────────────────────────────────────────────
2
+ import os, re, json, time, random, shutil, unicodedata, numpy as np, pandas as pd
3
+ from getpass import getpass
4
+ from pymilvus import MilvusClient
5
+ from groq import Groq
6
+ from openai import OpenAI
7
+ from sentence_transformers import SentenceTransformer, CrossEncoder
8
+ from rank_bm25 import BM25Okapi
9
+ from sklearn.metrics import roc_auc_score
10
+ import torch
11
+ from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
12
+ from huggingface_hub import hf_hub_download, list_repo_files, HfFileSystem, login
13
+ from datasets import load_dataset
14
+ import gradio as gr
15
+
16
+
17
+ # ── Cell 3: API Keys ───────────────────────────────────────────────────────────
18
+ # Choose your provider: "groq" or "openrouter"
19
+ LLM_PROVIDER = "openrouter" # ← change to "groq" if preferred
20
+
21
+ import os
22
+ from getpass import getpass
23
+ from huggingface_hub import login
24
+
25
+ def get_secret_or_prompt(secret_name, prompt_text=None):
26
+ """
27
+ Try to read secret from Google Colab Secrets.
28
+ If not available, ask user securely using getpass().
29
+ """
30
+
31
+ value = None
32
+
33
+ # Try Colab Secrets first
34
+ try:
35
+ #from google.colab import userdata
36
+ value = os.environ.get(secret_name)
37
+ except Exception:
38
+ value = None
39
+
40
+ # Fallback to environment variable
41
+ if not value:
42
+ value = os.environ.get(secret_name)
43
+
44
+ # Fallback to manual secure input
45
+ if not value:
46
+ prompt_text = prompt_text or f"Enter {secret_name}: "
47
+ value = getpass(prompt_text)
48
+
49
+ return value
50
+
51
+
52
+ # ── HuggingFace Token ─────────────────────────────────────────────────────────
53
+
54
+ HF_TOKEN = get_secret_or_prompt(
55
+ "HF_TOKEN",
56
+ "Enter HuggingFace Token: "
57
+ )
58
+
59
+ login(token=HF_TOKEN)
60
+ os.environ["HF_TOKEN"] = HF_TOKEN
61
+
62
+ print("βœ… HuggingFace token loaded and login completed")
63
+
64
+
65
+ # ── LLM Provider API Key ──────────────────────────────────────────────────────
66
+
67
+ if LLM_PROVIDER == "groq":
68
+
69
+ GROQ_API_KEY = get_secret_or_prompt(
70
+ "GROQ_API_KEY",
71
+ "Enter GROQ API Key: "
72
+ )
73
+
74
+ OPENROUTER_API_KEY = None
75
+
76
+ os.environ["GROQ_API_KEY"] = GROQ_API_KEY
77
+
78
+ print("βœ… GROQ API key loaded")
79
+
80
+ elif LLM_PROVIDER == "openrouter":
81
+
82
+ OPENROUTER_API_KEY = get_secret_or_prompt(
83
+ "OPENROUTER_API_KEY",
84
+ "Enter OpenRouter API Key: "
85
+ )
86
+
87
+ GROQ_API_KEY = None
88
+
89
+ os.environ["OPENROUTER_API_KEY"] = OPENROUTER_API_KEY
90
+
91
+ print("βœ… OpenRouter API key loaded")
92
+
93
+ else:
94
+ raise ValueError(f"Unknown LLM_PROVIDER: {LLM_PROVIDER}")
95
+
96
+
97
+ # ── Cell 4: Global configuration ──────────────────────────────────────────────
98
+ BUCKET_ID = "Phani555/IIITH-Cohort26-RAG-Batch37-storage"
99
+ BUCKET_PREFIX = f"hf://buckets/{BUCKET_ID}/milvus_dbs"
100
+ MILVUS_DIR = "/content/milvus_store/milvus_dbs"
101
+ HF_REPO_ID = "Phani555/IIITH-Cohort26-RAG-Batch37-storage"
102
+ HF_REPO_TYPE = "dataset"
103
+ HF_FOLDER = "ablations"
104
+
105
+ # ── Download mode ─────────────────────────────────────────────────────────────
106
+ # "chunk_v5_domain_aware" : new advanced domain-aware chunk_v5 indexes (recommended)
107
+ # "llm_embedder_default" : legacy llm_embedder default indexes
108
+ DOWNLOAD_MODE = "chunk_v5_domain_aware"
109
+
110
+ if DOWNLOAD_MODE == "chunk_v5_domain_aware":
111
+ INDEX_VERSION = "chunk_v5_domain_aware"
112
+ DOMAIN_EMBEDDING_RECOMMENDATION = {
113
+ "Customer_Support": "qwen3_embedding_0_6b",
114
+ "Bio_Medical": "bge_m3",
115
+ "General_Knowledge":"qwen3_embedding_0_6b",
116
+ "Legal_Contracts": "bge_m3",
117
+ "Finance": "bge_m3",
118
+ }
119
+ EMBEDDING_TYPE = "bge_m3"
120
+ else: # llm_embedder_default
121
+ INDEX_VERSION = "default"
122
+ DOMAIN_EMBEDDING_RECOMMENDATION = None
123
+ EMBEDDING_TYPE = "llm_embedder"
124
+
125
+ # ── Embedding models ───────────────────────────────────────────────────────────
126
+ EMBED_MODELS = {
127
+ "bge_small": "BAAI/bge-small-en-v1.5",
128
+ "llm_embedder": "BAAI/llm-embedder",
129
+ "bge_m3": "BAAI/bge-m3",
130
+ "qwen3_embedding_0_6b": "Qwen/Qwen3-Embedding-0.6B",
131
+ }
132
+ EMBEDDING_CHOICES = list(EMBED_MODELS.keys())
133
+
134
+ # ── Model lists per LLM provider ──────────────────────────────────────────────
135
+ GROQ_LLM_CHOICES = [
136
+ "llama-3.1-8b-instant",
137
+ "gemma2-9b-it",
138
+ "llama-3.3-70b-versatile",
139
+ "mixtral-8x7b-32768",
140
+ "qwen/qwen3-32b",
141
+ "qwen-qwq-32b",
142
+ "deepseek-r1-distill-llama-70b",
143
+ ]
144
+ OPENROUTER_LLM_CHOICES = [
145
+ "meta-llama/llama-3.1-8b-instruct",
146
+ "meta-llama/llama-3.3-70b-instruct",
147
+ "openai/gpt-oss-20b",
148
+ "openai/gpt-oss-120b",
149
+ "qwen/qwen3-32b",
150
+ "deepseek/deepseek-r1",
151
+ "moonshotai/kimi-k2-instruct",
152
+ "openai/gpt-oss-safeguard-20b",
153
+ ]
154
+ LLM_CHOICES = OPENROUTER_LLM_CHOICES if LLM_PROVIDER == "openrouter" else GROQ_LLM_CHOICES
155
+
156
+ # ── Runtime globals ────────────────────────────────────────────────────────────
157
+ MODEL_NAME = LLM_CHOICES[0]
158
+ MODEL_NAME_BIG = LLM_CHOICES[4] if len(LLM_CHOICES) > 4 else LLM_CHOICES[-1]
159
+
160
+ # ── Feature flags ──────────────────────────────────────────────────────────────
161
+ ENABLE_HYBRID = True
162
+ ENABLE_HYDE = False
163
+ ENABLE_RERANKING = False
164
+ RERANKER_TYPE = "monot5" # monot5 | tilde
165
+ ENABLE_RRF = True # Reciprocal Rank Fusion inside hybrid search
166
+ RRF_K = 60 # standard RRF constant
167
+ PROMPT_STRATEGY = "short" # short | long | long_cot
168
+ ENABLE_REPACKING = False
169
+ REPACK_STRATEGY = "sides" # forward | reverse | sides
170
+ ENABLE_SUMMARIZATION = False
171
+ SUMMARIZATION_TYPE = "recomp" # recomp | longllmlingua
172
+ ENABLE_QUERY_REWRITING = False
173
+ ENABLE_QUERY_DECOMPOSITION = False
174
+ ENABLE_QUERY_CLASSIFICATION= False
175
+ MAX_SUBQUERIES = 3
176
+ QUERY_REWRITE_MODEL = None
177
+ QUERY_DECOMPOSE_MODEL = None
178
+ RETRIEVE_DEBUG = False
179
+
180
+ # ── Tunable knobs ─────────────────────────────────────────────────────────────
181
+ RETRIEVE_TOP_K = 10
182
+ RERANK_TOP_K = 3
183
+ HYBRID_ALPHA = 0.5
184
+ MONOT5_MODEL = "castorini/monot5-base-msmarco-10k"
185
+ TILDE_MODEL = "BAAI/bge-reranker-base"
186
+ RECOMP_TOP_K_SENTS = 6
187
+ RECOMP_MIN_SCORE = 0.00
188
+ RECOMP_GROUNDING_BOOST = 0.15
189
+ RECOMP_MIN_KEEP_RATIO = 0.30
190
+ RECOMP_KEEP_CRITICAL = True
191
+ LLMLINGUA_RATE = 0.5
192
+
193
+ # ── Runtime state ─────────────────────────────────────────────────────────────
194
+ milvus_clients = {}
195
+ bm25_indexes = {}
196
+ loaded_embedding_models = {} # keyed by embedding_type string
197
+ embed_model = None # single fallback embed model
198
+ llm_client = None
199
+ monot5_reranker = None
200
+ tilde_reranker = None
201
+ llmlingua_compressor = None
202
+ ragbench_by_domain = {}
203
+ LEGAL_SAMPLE_TO_CONTRACT_ID = {}
204
+
205
+ DOMAIN_NAMES = [
206
+ "Bio_Medical",
207
+ "General_Knowledge",
208
+ "Customer_Support",
209
+ "Finance",
210
+ "Legal_Contracts",
211
+ ]
212
+
213
+ GROUNDING_PATTERNS = [
214
+ r"\b(?:must|should|shall|cannot|can't|never|always|only|except|unless|required|recommended)\b",
215
+ r"\b(?:warning|caution|note|important|attention)\b",
216
+ r"\b(?:do not|don't|does not|did not|not allowed|not recommended|never)\b",
217
+ r"\b\d+(?:\.\d+)?\s*(?:%|percent|seconds?|minutes?|hours?|days?|weeks?|months?|years?)\b",
218
+ r"\b\d+(?:\.\d+)?\s*(?:GB|MB|KB|TB|kg|g|mg|mm|cm|m|km|degrees?|Β°C|Β°F)\b",
219
+ r"[$€£Β₯]\s*\d+(?:,\d{3})*(?:\.\d+)?",
220
+ r"\b\d+(?:,\d{3})*(?:\.\d+)?\s*(?:dollars?|rupees?|crores?|lakhs?|million|billion)\b",
221
+ r"\b(?:19|20)\d{2}\b",
222
+ r"\b\d{2,}\b",
223
+ r"\b[A-Z]{2,}[-_]?\d+[A-Z0-9-]*\b",
224
+ r"\b[A-Z0-9]{3,}[-_][A-Z0-9]{2,}\b",
225
+ r"\b[A-Z]{3,}\b",
226
+ ]
227
+ GROUNDING_REGEX = re.compile("|".join(GROUNDING_PATTERNS), re.IGNORECASE)
228
+
229
+ print(f"Config loaded. Mode: {DOWNLOAD_MODE} | Provider: {LLM_PROVIDER} | Models: {len(LLM_CHOICES)}")
230
+
231
+
232
+ # ── Cell 5: Pipeline functions (Advanced – chunk_v5 + RRF + Legal contract filtering) ──
233
+ # NOTE: _hf_fs is initialized in Cell 7. hf_path_exists() uses globals() so it
234
+ # safely resolves _hf_fs at call-time, not at definition-time.
235
+
236
+ # ── Utilities ──────────────────────────────────────────────────────────────────
237
+
238
+ def _safe_message_content(response):
239
+ try:
240
+ msg = response.choices[0].message
241
+ content = getattr(msg, "content", None)
242
+ return str(content).strip() if content else ""
243
+ except Exception:
244
+ return ""
245
+
246
+ def _sanitize(text):
247
+ if not text: return text
248
+ text = unicodedata.normalize("NFC", str(text))
249
+ return text.encode("ascii", errors="replace").decode("ascii")
250
+
251
+ def get_domain(dataset):
252
+ if dataset in ("covidqa","pubmedqa"): return "Bio_Medical"
253
+ elif dataset in ("expertqa","hagrid","hotpotqa","msmarco"): return "General_Knowledge"
254
+ elif dataset in ("delucionqa","emanual","techqa"): return "Customer_Support"
255
+ elif dataset in ("finqa","tatqa"): return "Finance"
256
+ else: return "Legal_Contracts"
257
+
258
+ def split_into_sentences(text):
259
+ return [s.strip() for s in re.split(r'(?<=[.!?])\s+', str(text).strip()) if s.strip()]
260
+
261
+ def _tokenize(text):
262
+ return re.findall(r'\w+', str(text).lower())
263
+
264
+ def _normalize(scores):
265
+ arr = np.array(scores, dtype=float)
266
+ if len(arr) == 0 or arr.max() == arr.min(): return np.zeros_like(arr)
267
+ return (arr - arr.min()) / (arr.max() - arr.min())
268
+
269
+ def _count_grounding_signals(sentence):
270
+ return len(GROUNDING_REGEX.findall(str(sentence)))
271
+
272
+ def _is_critical_sentence(sentence):
273
+ pat = re.compile(
274
+ r"\b(?:warning|caution|important|must|must not|cannot|can't|do not|don't|never|only|except|unless|required)\b",
275
+ re.IGNORECASE)
276
+ return bool(pat.search(str(sentence)))
277
+
278
+
279
+ # ── LLM client factory ─────────────────────────────────────────────────────────
280
+
281
+ def get_llm_client():
282
+ if LLM_PROVIDER == "groq":
283
+ return Groq(api_key=GROQ_API_KEY)
284
+ elif LLM_PROVIDER == "openrouter":
285
+ return OpenAI(api_key=OPENROUTER_API_KEY, base_url="https://openrouter.ai/api/v1")
286
+ raise ValueError(f"Unknown LLM_PROVIDER: {LLM_PROVIDER}")
287
+
288
+
289
+ # ── DB path helpers ────────────────────────────────────────────────────────────
290
+
291
+ def get_index_folder(embedding_type=None, index_version=None):
292
+ embedding_type = embedding_type or EMBEDDING_TYPE
293
+ index_version = index_version or INDEX_VERSION
294
+ return embedding_type if index_version == "default" else f"{embedding_type}_{index_version}"
295
+
296
+ def get_db_path(domain_name, embedding_type=None, index_version=None):
297
+ folder = get_index_folder(embedding_type, index_version)
298
+ db_dir = os.path.join(MILVUS_DIR, folder)
299
+ os.makedirs(db_dir, exist_ok=True)
300
+ return os.path.join(db_dir, f"{domain_name}.db")
301
+
302
+ def get_embedding_type_for_domain(domain_name):
303
+ rec = globals().get("DOMAIN_EMBEDDING_RECOMMENDATION")
304
+ if rec: return rec.get(domain_name, EMBEDDING_TYPE)
305
+ return EMBEDDING_TYPE
306
+
307
+ def get_embed_model_for_domain(domain_name):
308
+ emb_type = get_embedding_type_for_domain(domain_name)
309
+ models = globals().get("loaded_embedding_models", {})
310
+ if emb_type not in models:
311
+ raise ValueError(f"Embedding type '{emb_type}' not in loaded_embedding_models. Run Cell 7 first.")
312
+ return models[emb_type]
313
+
314
+
315
+ # ── HF filesystem helper ───────────────────────────────────────────────────────
316
+ # Uses globals() so _hf_fs is resolved at call-time (Cell 7), not import-time (Cell 5).
317
+
318
+ def hf_path_exists(path):
319
+ fs = globals().get("_hf_fs")
320
+ if fs is None:
321
+ raise RuntimeError("_hf_fs not initialised β€” run Cell 7 before Cell 8.")
322
+ try:
323
+ fs.ls(path); return True
324
+ except Exception:
325
+ return False
326
+
327
+
328
+ # ── Legal contract helpers ─────────────────────────────────────────────────────
329
+
330
+ def build_contract_filter_expr(contract_id):
331
+ contract_id = str(contract_id).replace('"', '\\"')
332
+ return f'contract_id == "{contract_id}"'
333
+
334
+ def load_legal_sample_to_contract_mapping(local_path=None):
335
+ if local_path is None:
336
+ legal_folder = get_index_folder(get_embedding_type_for_domain("Legal_Contracts"), INDEX_VERSION)
337
+ local_path = os.path.join(MILVUS_DIR, legal_folder, "legal_sample_to_contract_id.json")
338
+ if not os.path.exists(local_path):
339
+ print(f" WARNING: Legal mapping not found: {local_path}")
340
+ return {}
341
+ with open(local_path, "r") as f:
342
+ mapping = json.load(f)
343
+ print(f" Legal sample->contract mapping loaded: {len(mapping):,} entries")
344
+ return mapping
345
+
346
+ def get_contract_id_for_legal_sample(sample_id):
347
+ mapping = globals().get("LEGAL_SAMPLE_TO_CONTRACT_ID", {})
348
+ sid = str(sample_id)
349
+ if sid not in mapping:
350
+ raise ValueError(f"sample_id '{sid}' not found in Legal mapping.")
351
+ return mapping[sid]
352
+
353
+
354
+ # ── Query Classification ───────────────────────────────────────────────────────
355
+
356
+ def classify_query(query, domain_name=None):
357
+ if not ENABLE_QUERY_CLASSIFICATION: return "RAG"
358
+ rag_domains = {"Bio_Medical","General_Knowledge","Customer_Support","Finance","Legal_Contracts"}
359
+ if domain_name in rag_domains: return "RAG"
360
+ llm_keywords = ["who is","what is","when was","where is","define","explain",
361
+ "tell me about","what are","why is","how does","what does"]
362
+ if any(kw in str(query).lower() for kw in llm_keywords): return "LLM"
363
+ return "RAG"
364
+
365
+
366
+ # ── Query Rewriting ────────────────────────────────────────────────────────────
367
+
368
+ def rewrite_query(query, domain_name, llm_client, model_name=None):
369
+ if not ENABLE_QUERY_REWRITING: return query
370
+ model = model_name or QUERY_REWRITE_MODEL or MODEL_NAME
371
+ prompt = f"""Rewrite the question to improve document retrieval. Apply only when needed.
372
+ Domain: {domain_name}
373
+ Rules: Preserve meaning. Fix grammar. Expand abbreviations. Preserve all names/numbers/terms.
374
+ Do not answer. Return ONLY the rewritten query.
375
+ Original question: {query}""".strip()
376
+ try:
377
+ resp = llm_client.chat.completions.create(
378
+ model=model,
379
+ messages=[{"role":"system","content":"You rewrite questions to improve semantic document retrieval. Return only the rewritten question."},
380
+ {"role":"user","content":_sanitize(prompt)}],
381
+ temperature=0.0, max_tokens=150,
382
+ )
383
+ rewritten = _safe_message_content(resp).strip()
384
+ return rewritten if rewritten and len(rewritten) < 600 else query
385
+ except Exception as e:
386
+ print(f"Query rewriting failed: {e}"); return query
387
+
388
+
389
+ # ── Query Decomposition helpers ────────────────────────────────────────────────
390
+
391
+ def _clean_subquery_text(text):
392
+ if text is None: return ""
393
+ text = str(text).strip().replace("```json","").replace("```","").strip()
394
+ text = text.rstrip(",").strip('"').strip("'").strip()
395
+ text = re.sub(r"^\s*[-*]\s*","",text); text = re.sub(r"^\s*\d+[\).\:\-]\s*","",text)
396
+ return text.strip()
397
+
398
+ def _looks_like_explanation_line(text):
399
+ if not text: return True
400
+ tl = text.lower().strip()
401
+ bad = ["here are","here is","decomposed","search queries","the decomposed",
402
+ "queries:","subqueries:","output:","json:","answer:"]
403
+ if any(tl.startswith(p) for p in bad): return True
404
+ if tl in {"queries","subqueries","search queries","decomposed search queries"}: return True
405
+ return False
406
+
407
+ def _parse_json_object_line(line):
408
+ line = _clean_subquery_text(line)
409
+ if not line: return None
410
+ try:
411
+ obj = json.loads(line)
412
+ if isinstance(obj, dict):
413
+ for k in ["query","question","subquery","search_query"]:
414
+ if k in obj and str(obj[k]).strip(): return str(obj[k]).strip()
415
+ if isinstance(obj, str): return obj.strip()
416
+ except Exception: pass
417
+ m = re.search(r'"(?:query|question|subquery|search_query)"\s*:\s*"([^"]+)"', line)
418
+ if m: return m.group(1).strip()
419
+ return None
420
+
421
+ def _split_multi_question_locally(query, max_subqueries=None):
422
+ max_subqueries = max_subqueries or MAX_SUBQUERIES
423
+ parts = [p.strip() for p in re.split(r"\?\s*", str(query).strip()) if p.strip()]
424
+ if len(parts) <= 1: return None
425
+ return [(p+"?" if not p.endswith("?") else p) for p in parts[:max_subqueries]]
426
+
427
+ def _parse_subqueries(raw_text, original_query, max_subqueries=None):
428
+ max_subqueries = max_subqueries or MAX_SUBQUERIES
429
+ if not raw_text: return [original_query]
430
+ text = str(raw_text).strip().replace("```json","").replace("```","").strip()
431
+ try:
432
+ parsed = json.loads(text)
433
+ if isinstance(parsed, list):
434
+ subs = []
435
+ for item in parsed:
436
+ if isinstance(item, dict):
437
+ for k in ["query","question","subquery","search_query"]:
438
+ if k in item and str(item[k]).strip(): subs.append(str(item[k]).strip()); break
439
+ elif isinstance(item, str): subs.append(item.strip())
440
+ subs = [_clean_subquery_text(q) for q in subs if _clean_subquery_text(q)]
441
+ return subs[:max_subqueries] or [original_query]
442
+ elif isinstance(parsed, dict):
443
+ raw_list = parsed.get("subqueries") or parsed.get("queries") or parsed.get("questions") or []
444
+ if isinstance(raw_list, list):
445
+ subs = [_clean_subquery_text(q) for q in raw_list if _clean_subquery_text(q)]
446
+ return subs[:max_subqueries] or [original_query]
447
+ except Exception: pass
448
+ subqueries = []
449
+ for raw_line in text.splitlines():
450
+ line = _clean_subquery_text(raw_line)
451
+ if not line or _looks_like_explanation_line(line): continue
452
+ obj_q = _parse_json_object_line(line)
453
+ if obj_q:
454
+ obj_q = _clean_subquery_text(obj_q)
455
+ if obj_q and not _looks_like_explanation_line(obj_q): subqueries.append(obj_q)
456
+ continue
457
+ if line.startswith("{") or line.endswith("}") or line in {"[","]","{","}"}: continue
458
+ subqueries.append(line)
459
+ deduped = []
460
+ for q in subqueries:
461
+ q = _clean_subquery_text(q)
462
+ if q and q not in deduped: deduped.append(q)
463
+ return deduped[:max_subqueries] or [original_query]
464
+
465
+ def decompose_query(query, llm_client, domain=None, model=None, max_subqueries=None):
466
+ if not ENABLE_QUERY_DECOMPOSITION: return [query]
467
+ max_subqueries = max_subqueries or MAX_SUBQUERIES
468
+ local_split = _split_multi_question_locally(query, max_subqueries)
469
+ if local_split: return local_split
470
+ model = model or QUERY_DECOMPOSE_MODEL or MODEL_NAME
471
+ if not model: return [query]
472
+ prompt = f"""Decompose the question into at most {max_subqueries} retrieval-focused search queries.
473
+ Return ONLY a valid JSON list of strings. No explanations. No markdown.
474
+ Example: ["What caused the 2008 crisis?", "Which banks failed in 2008?"]
475
+ Rules: If already simple return list with original. Preserve all technical terms. Do not answer.
476
+ Domain: {domain}
477
+ Question: {query}""".strip()
478
+ try:
479
+ resp = llm_client.chat.completions.create(
480
+ model=model,
481
+ messages=[{"role":"system","content":"You decompose complex questions into retrieval subqueries and return only a JSON list of strings."},
482
+ {"role":"user","content":_sanitize(prompt)}],
483
+ temperature=0.0, max_tokens=300,
484
+ )
485
+ return _parse_subqueries(_safe_message_content(resp), original_query=query, max_subqueries=max_subqueries)
486
+ except Exception as e:
487
+ print(f"Query decomposition failed: {e}"); return [query]
488
+
489
+
490
+ # ── Reranking ──────────────────────────────────────────────────────────────────
491
+
492
+ class MonoT5Reranker:
493
+ def __init__(self, model_name=None):
494
+ model_name = model_name or MONOT5_MODEL
495
+ self.tokenizer = AutoTokenizer.from_pretrained(model_name)
496
+ self.model = AutoModelForSeq2SeqLM.from_pretrained(model_name)
497
+ self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
498
+ self.model.to(self.device); self.model.eval()
499
+ self.true_id = self.tokenizer.convert_tokens_to_ids("▁true")
500
+ self.false_id = self.tokenizer.convert_tokens_to_ids("▁false")
501
+ if not self.true_id or self.true_id < 0: self.true_id = self.tokenizer.encode("true", add_special_tokens=False)[0]
502
+ if not self.false_id or self.false_id < 0: self.false_id = self.tokenizer.encode("false", add_special_tokens=False)[0]
503
+ print(f"MonoT5 loaded: {model_name} on {self.device}")
504
+
505
+ def score(self, query, document):
506
+ text = f"Query: {query} Document: {document} Relevant:"
507
+ enc = self.tokenizer(text, return_tensors="pt", max_length=512, truncation=True).to(self.device)
508
+ with torch.no_grad():
509
+ out = self.model.generate(**enc, max_new_tokens=1, return_dict_in_generate=True, output_scores=True)
510
+ logits = out.scores[0][0]
511
+ probs = torch.softmax(torch.stack([logits[self.false_id], logits[self.true_id]]), dim=0)
512
+ return float(probs[1].item())
513
+
514
+ def compute_scores(self, query, texts):
515
+ return np.array([self.score(query, t) for t in texts], dtype=float)
516
+
517
+ def get_monot5_reranker():
518
+ global monot5_reranker
519
+ if monot5_reranker is None: monot5_reranker = MonoT5Reranker(MONOT5_MODEL)
520
+ return monot5_reranker
521
+
522
+ def get_tilde_reranker():
523
+ global tilde_reranker
524
+ if tilde_reranker is None:
525
+ dev = "cuda" if torch.cuda.is_available() else "cpu"
526
+ tilde_reranker = CrossEncoder(TILDE_MODEL, device=dev)
527
+ print(f"TILDE reranker loaded on {dev}")
528
+ return tilde_reranker
529
+
530
+ def rerank_documents(query, documents, top_k=3):
531
+ """Rerank documents; preserves all metadata fields including Legal contract_id."""
532
+ if not documents: return []
533
+ texts = [d["text"] if isinstance(d, dict) else d for d in documents]
534
+ rtype = RERANKER_TYPE.lower().strip()
535
+ scores = get_monot5_reranker().compute_scores(query, texts) if rtype == "monot5" \
536
+ else np.asarray(get_tilde_reranker().predict([(query, t) for t in texts], show_progress_bar=False), dtype=float).reshape(-1)
537
+ ranked_idx = np.argsort(scores)[::-1][:top_k]
538
+ reranked = []
539
+ for i in ranked_idx:
540
+ # dict() shallow-copies ALL fields (dense_score, bm25_score, contract_id, etc.)
541
+ item = dict(documents[i]) if isinstance(documents[i], dict) else {"text": documents[i]}
542
+ item["base_score"] = item.get("score") # preserve original retrieval score
543
+ item["score"] = float(scores[i])
544
+ item["rerank_score"] = float(scores[i])
545
+ item["reranker_type"] = rtype
546
+ reranked.append(item)
547
+ return reranked
548
+
549
+
550
+ # ── BM25 (stores contract_ids for Legal to enable per-contract filtering) ─────
551
+
552
+ def build_bm25_index(domain_name, clients):
553
+ client = clients[domain_name]; col = domain_name.lower()
554
+ try:
555
+ if "Loaded" not in str(client.get_load_state(col)): client.load_collection(col)
556
+ except Exception: pass
557
+ try: n = int(client.get_collection_stats(col).get("row_count", 0))
558
+ except Exception: return
559
+ if n == 0: return
560
+ output_fields = ["text"]
561
+ if domain_name == "Legal_Contracts": output_fields.append("contract_id")
562
+ rows = client.query(collection_name=col, filter="", limit=n, output_fields=output_fields)
563
+ if not rows: return
564
+ texts = []; contract_ids = []
565
+ for r in rows:
566
+ t = r.get("text","")
567
+ if not t: continue
568
+ texts.append(t)
569
+ if domain_name == "Legal_Contracts": contract_ids.append(r.get("contract_id"))
570
+ if not texts: return
571
+ bm25_indexes[domain_name] = {
572
+ "bm25": BM25Okapi([_tokenize(t) for t in texts]),
573
+ "texts": texts,
574
+ "contract_ids": contract_ids if domain_name == "Legal_Contracts" else None,
575
+ }
576
+ print(f" BM25 built: {len(texts)} docs [{domain_name}]")
577
+ if domain_name == "Legal_Contracts":
578
+ valid = sum(1 for c in contract_ids if c is not None)
579
+ print(f" Legal contract_ids: {valid:,}/{len(texts):,}")
580
+
581
+ def build_all_bm25_indexes(clients):
582
+ bm25_indexes.clear()
583
+ for d in clients: build_bm25_index(d, clients)
584
+
585
+
586
+ # ── HyDE ───────────────────────────────────────────────────────────────────────
587
+
588
+ def generate_hyde(query, llm_client, model_name=None):
589
+ model = model_name or MODEL_NAME
590
+ prompt = f"Write a brief factual passage answering this question (under 4 sentences).\nQuestion: {query}\nPassage:"
591
+ try:
592
+ resp = llm_client.chat.completions.create(
593
+ model=model,
594
+ messages=[{"role":"system","content":"You write hypothetical answer passages for retrieval."},
595
+ {"role":"user","content":_sanitize(prompt)}],
596
+ temperature=0.2, max_tokens=300,
597
+ )
598
+ return _safe_message_content(resp)
599
+ except Exception as e:
600
+ print(f"HyDE failed: {e}"); return ""
601
+
602
+
603
+ # ── Reciprocal Rank Fusion ─────────────────────────────────────────────────────
604
+
605
+ def reciprocal_rank_fusion(dense_results, bm25_results, top_k=20, rrf_k=60,
606
+ dense_meta=None, bm25_meta=None):
607
+ """
608
+ Fuse dense and BM25 ranked lists using RRF.
609
+ score(doc) = 1/(k + rank_dense) + 1/(k + rank_bm25)
610
+ dense_meta / bm25_meta: optional dicts of extra fields per text (e.g. Legal metadata).
611
+ """
612
+ dense_meta = dense_meta or {}; bm25_meta = bm25_meta or {}
613
+ rrf_scores = {}
614
+ for rank, (text, score) in enumerate(
615
+ sorted(dense_results.items(), key=lambda x: x[1], reverse=True), start=1):
616
+ rrf_scores.setdefault(text, {"text":text,"dense_score":float(score),"bm25_score":0.0,"score":0.0})
617
+ rrf_scores[text]["dense_score"] = float(score)
618
+ rrf_scores[text]["score"] += 1.0 / (rrf_k + rank)
619
+ if text in dense_meta: rrf_scores[text].update(dense_meta[text])
620
+ for rank, (text, score) in enumerate(
621
+ sorted(bm25_results.items(), key=lambda x: x[1], reverse=True), start=1):
622
+ rrf_scores.setdefault(text, {"text":text,"dense_score":0.0,"bm25_score":float(score),"score":0.0})
623
+ rrf_scores[text]["bm25_score"] = float(score)
624
+ rrf_scores[text]["score"] += 1.0 / (rrf_k + rank)
625
+ # dense_meta takes priority over bm25_meta for Legal metadata consistency
626
+ if text not in dense_meta and text in bm25_meta:
627
+ rrf_scores[text].update(bm25_meta[text])
628
+ fused = sorted(rrf_scores.values(), key=lambda x: x["score"], reverse=True)
629
+ return fused[:top_k]
630
+
631
+
632
+ # ── Hybrid search (dense + BM25, Legal contract filtering, RRF or alpha fusion) ─
633
+
634
+ def hybrid_search(query, domain_name, embed_model, top_k=20, alpha=0.5, contract_id=None):
635
+ client = milvus_clients[domain_name]; col = domain_name.lower()
636
+ try:
637
+ if "Loaded" not in str(client.get_load_state(col)): client.load_collection(col)
638
+ except Exception: pass
639
+
640
+ # ── Dense search ───────────────────────────────────────────���─────────────
641
+ q_emb = embed_model.encode([query], normalize_embeddings=True, convert_to_numpy=True).astype("float32")
642
+ output_fields = ["text"]
643
+ if domain_name == "Legal_Contracts": output_fields += ["contract_id","source_doc_id","source_hash"]
644
+ search_kwargs = dict(collection_name=col, data=q_emb.tolist(), limit=top_k,
645
+ output_fields=output_fields, search_params={"metric_type":"IP","params":{}})
646
+ if domain_name == "Legal_Contracts" and contract_id is not None:
647
+ search_kwargs["filter"] = build_contract_filter_expr(contract_id)
648
+ dense_hits = client.search(**search_kwargs)
649
+ dense_results = {}; dense_meta = {}
650
+ hit_list = dense_hits[0] if (dense_hits and isinstance(dense_hits[0], (list,tuple))) else dense_hits
651
+ for hit in hit_list:
652
+ entity = (hit.get("entity",{}) or hit) if isinstance(hit,dict) else (getattr(hit,"entity",{}) or {})
653
+ distance = hit.get("distance",0.0) if isinstance(hit,dict) else getattr(hit,"distance",0.0)
654
+ text = entity.get("text","")
655
+ if not text: continue
656
+ dense_results[text] = float(distance)
657
+ if domain_name == "Legal_Contracts":
658
+ dense_meta[text] = {"contract_id":entity.get("contract_id"),
659
+ "source_doc_id":entity.get("source_doc_id"),
660
+ "source_hash":entity.get("source_hash")}
661
+
662
+ # ── BM25 search ───────────────────────────────────────────────────────────
663
+ bm25_obj = bm25_indexes.get(domain_name)
664
+ if not bm25_obj: # fallback to dense-only
665
+ rows = [{"text":t,"score":s,"dense_score":s,"bm25_score":0.0,"hybrid_fallback":True}
666
+ for t,s in sorted(dense_results.items(),key=lambda x:-x[1])[:top_k]]
667
+ if domain_name == "Legal_Contracts":
668
+ for r in rows: r.update(dense_meta.get(r["text"],{}))
669
+ return rows
670
+
671
+ bm25_scores = bm25_obj["bm25"].get_scores(_tokenize(query))
672
+ bm25_texts = bm25_obj["texts"]
673
+ bm25_cids = bm25_obj.get("contract_ids")
674
+
675
+ # For Legal: filter BM25 candidates to the same contract before ranking
676
+ if domain_name == "Legal_Contracts" and contract_id is not None and bm25_cids:
677
+ candidate_idx = [i for i,cid in enumerate(bm25_cids) if str(cid)==str(contract_id)]
678
+ else:
679
+ candidate_idx = list(range(len(bm25_texts)))
680
+ top_bm25_idx = sorted(candidate_idx, key=lambda i: bm25_scores[i], reverse=True)[:top_k]
681
+
682
+ bm25_results = {}; bm25_meta = {}
683
+ for i in top_bm25_idx:
684
+ text = bm25_texts[i]
685
+ bm25_results[text] = float(bm25_scores[i])
686
+ if domain_name == "Legal_Contracts":
687
+ bm25_meta[text] = {"contract_id": str(contract_id) if contract_id else (bm25_cids[i] if bm25_cids else None)}
688
+
689
+ # ── Fusion ────────────────────────────────────────────────────────────────
690
+ if ENABLE_RRF:
691
+ return reciprocal_rank_fusion(dense_results, bm25_results, top_k=top_k, rrf_k=RRF_K,
692
+ dense_meta=dense_meta, bm25_meta=bm25_meta)
693
+
694
+ # Alpha-weighted min-max fusion (fallback when RRF disabled)
695
+ all_texts = sorted(set(dense_results)|set(bm25_results))
696
+ d_vals = [dense_results.get(t,0.0) for t in all_texts]
697
+ b_vals = [bm25_results.get(t,0.0) for t in all_texts]
698
+ d_norm, b_norm = _normalize(d_vals), _normalize(b_vals)
699
+ combined = []
700
+ for i, text in enumerate(all_texts):
701
+ row = {"text":text,"score":float(alpha*d_norm[i]+(1-alpha)*b_norm[i]),
702
+ "dense_score":float(d_vals[i]),"bm25_score":float(b_vals[i]),"alpha":alpha}
703
+ if domain_name == "Legal_Contracts":
704
+ row.update(dense_meta.get(text, bm25_meta.get(text,{})))
705
+ if "contract_id" not in row and contract_id is not None:
706
+ row["contract_id"] = str(contract_id)
707
+ combined.append(row)
708
+ combined.sort(key=lambda x: -x["score"])
709
+ return combined[:top_k]
710
+
711
+
712
+ # ── Repacking ──────────────────────────────────────────────────────────────────
713
+
714
+ def repack_documents(docs, strategy="sides"):
715
+ """
716
+ Reorder retrieved documents for LLM attention bias.
717
+ forward: most-relevant first (no change)
718
+ reverse: most-relevant last (benefits models that attend to end of context)
719
+ sides: U-shape β€” highest-relevance at both ends, lowest in the middle
720
+ """
721
+ if not docs: return []
722
+ if strategy == "forward": return docs
723
+ if strategy == "reverse": return docs[::-1]
724
+ if strategy == "sides":
725
+ n, result, left, right = len(docs), [None]*len(docs), 0, len(docs)-1
726
+ for i, doc in enumerate(docs):
727
+ if i % 2 == 0: result[left] = doc; left += 1
728
+ else: result[right] = doc; right -= 1
729
+ return result
730
+ raise ValueError(f"Unknown REPACK_STRATEGY: {strategy}")
731
+
732
+
733
+ # ── Summarization ──────────────────────────────────────────────────────────────
734
+
735
+ def recomp_summarize(query, docs, em, top_k=6, min_score=0.0,
736
+ grounding_boost=0.15, min_keep_ratio=0.30, keep_critical=True):
737
+ texts = [d.get("text","") if isinstance(d,dict) else d for d in docs]
738
+ sentences = [s.strip() for doc in texts for s in split_into_sentences(doc) if s.strip()]
739
+ if not sentences: return ""
740
+ q_emb = em.encode([query], normalize_embeddings=True)
741
+ s_emb = em.encode(sentences, normalize_embeddings=True)
742
+ scores = (q_emb @ s_emb.T).flatten() + np.array([_count_grounding_signals(s)*grounding_boost for s in sentences])
743
+ crits = {i for i,s in enumerate(sentences) if keep_critical and _is_critical_sentence(s)}
744
+ valid = np.where(scores >= min_score)[0]
745
+ if len(valid) == 0: valid = np.array([int(np.argmax(scores))])
746
+ keep = min(max(top_k, int(np.ceil(len(sentences)*min_keep_ratio))), len(sentences))
747
+ chosen = sorted(set(list(valid[np.argsort(scores[valid])[::-1][:keep]])) | crits)
748
+ return " ".join(sentences[i] for i in chosen)
749
+
750
+ def _get_llmlingua():
751
+ global llmlingua_compressor
752
+ if llmlingua_compressor is None:
753
+ from llmlingua import PromptCompressor
754
+ dev = "cuda" if torch.cuda.is_available() else "cpu"
755
+ llmlingua_compressor = PromptCompressor(
756
+ model_name="microsoft/llmlingua-2-bert-base-multilingual-cased-meetingbank",
757
+ use_llmlingua2=True, device_map=dev)
758
+ print(f"LLMLingua loaded on {dev}")
759
+ return llmlingua_compressor
760
+
761
+ def llmlingua_compress(query, docs, rate=0.5):
762
+ texts = [d.get("text","") if isinstance(d,dict) else d for d in docs]
763
+ comp = _get_llmlingua()
764
+ parts = []
765
+ for t in texts:
766
+ if not t or not t.strip(): continue
767
+ try: parts.append(comp.compress_prompt(t, question=query, rate=rate)["compressed_prompt"])
768
+ except Exception as e: print(f"LLMLingua chunk failed: {e}"); parts.append(t)
769
+ return "\n\n".join(parts)
770
+
771
+ def summarize_docs(query, docs, em=None, llm_client=None):
772
+ """
773
+ FIX: explicit None guard on em before calling encode().
774
+ Falls back to global embed_model, then raises a clear error.
775
+ """
776
+ if not ENABLE_SUMMARIZATION:
777
+ return [d.get("text","") if isinstance(d,dict) else d for d in docs]
778
+
779
+ # Resolve embed model β€” must be non-None before encode()
780
+ if em is None:
781
+ em = globals().get("embed_model")
782
+ if em is None:
783
+ raise RuntimeError("summarize_docs: no embed_model available. Run Cell 7 first.")
784
+
785
+ if SUMMARIZATION_TYPE == "recomp":
786
+ s = recomp_summarize(query, docs, em, RECOMP_TOP_K_SENTS, RECOMP_MIN_SCORE,
787
+ RECOMP_GROUNDING_BOOST, RECOMP_MIN_KEEP_RATIO, RECOMP_KEEP_CRITICAL)
788
+ return [s] if s else []
789
+ elif SUMMARIZATION_TYPE == "longllmlingua":
790
+ c = llmlingua_compress(query, docs, LLMLINGUA_RATE)
791
+ return [c] if c else []
792
+ raise ValueError(f"Unknown SUMMARIZATION_TYPE: {SUMMARIZATION_TYPE}")
793
+
794
+
795
+ # ── Main retrieve ──────────────────────────────────────────────────────────────
796
+ # Pipeline order: HyDE β†’ Retrieve β†’ Rerank β†’ Repack β†’ Summarize
797
+
798
+ def retrieve(query, domain_name, embed_model=None, llm_client=None, top_k=None,
799
+ rewritten_query=None, sample_id=None, contract_id=None):
800
+ """
801
+ FIX: top_k now defaults to RETRIEVE_TOP_K (10), not RERANK_TOP_K (3).
802
+ The fetch_k logic already enlarges the initial pool; top_k is the final
803
+ count after reranking.
804
+ """
805
+ if domain_name not in milvus_clients:
806
+ raise ValueError(f"Domain '{domain_name}' not loaded.")
807
+
808
+ # FIX: default to RETRIEVE_TOP_K for initial fetch, not RERANK_TOP_K
809
+ if top_k is None:
810
+ top_k = globals().get("RETRIEVE_TOP_K", 10)
811
+ top_k = int(top_k)
812
+
813
+ llm = llm_client or globals().get("llm_client")
814
+
815
+ # Resolve domain-specific embed model
816
+ em = embed_model
817
+ if em is None:
818
+ try: em = get_embed_model_for_domain(domain_name)
819
+ except Exception: em = globals().get("embed_model")
820
+ if em is None:
821
+ raise ValueError(f"No embed_model available for domain '{domain_name}'. Run Cell 7 first.")
822
+
823
+ # Legal contract filtering (section 1.5)
824
+ legal_contract_id = None
825
+ if domain_name == "Legal_Contracts":
826
+ if contract_id is not None:
827
+ legal_contract_id = str(contract_id)
828
+ elif sample_id is not None:
829
+ try: legal_contract_id = get_contract_id_for_legal_sample(sample_id)
830
+ except Exception as e:
831
+ print(f" Could not resolve Legal contract_id for sample_id={sample_id}: {e}")
832
+ if legal_contract_id is None:
833
+ print(" WARNING: Legal_Contracts retrieval without contract filter β€” cross-contract contamination possible")
834
+
835
+ # Fetch more candidates if downstream processing will reduce count
836
+ fetch_k = max(int(RETRIEVE_TOP_K if (ENABLE_HYBRID or ENABLE_RERANKING or ENABLE_SUMMARIZATION) else top_k), top_k)
837
+
838
+ eff_q = rewritten_query or query; search_q = eff_q
839
+
840
+ # HyDE query expansion
841
+ if ENABLE_HYDE and llm:
842
+ try:
843
+ hyde = generate_hyde(eff_q, llm)
844
+ if hyde: search_q = f"{eff_q} {hyde}"
845
+ except Exception as e: print(f"HyDE failed: {e}")
846
+
847
+ # Retrieve
848
+ if ENABLE_HYBRID:
849
+ retrieved = hybrid_search(search_q, domain_name, em, top_k=fetch_k,
850
+ alpha=HYBRID_ALPHA, contract_id=legal_contract_id) or []
851
+ for d in retrieved:
852
+ if isinstance(d, dict): d.setdefault("retrieval_type","hybrid")
853
+ else:
854
+ client = milvus_clients[domain_name]; col = domain_name.lower()
855
+ try:
856
+ if "Loaded" not in str(client.get_load_state(col)): client.load_collection(col)
857
+ except Exception: pass
858
+ q_emb = em.encode([search_q], normalize_embeddings=True, convert_to_numpy=True).astype("float32")
859
+ output_fields = ["text"]
860
+ if domain_name == "Legal_Contracts": output_fields += ["contract_id","source_doc_id","source_hash"]
861
+ skw = dict(collection_name=col, data=q_emb.tolist(), limit=fetch_k,
862
+ output_fields=output_fields, search_params={"metric_type":"IP","params":{}})
863
+ if domain_name == "Legal_Contracts" and legal_contract_id:
864
+ skw["filter"] = build_contract_filter_expr(legal_contract_id)
865
+ hits = client.search(**skw)
866
+ hit_list = hits[0] if (hits and isinstance(hits[0],(list,tuple))) else hits
867
+ seen, retrieved = set(), []
868
+ for hit in hit_list:
869
+ entity = (hit.get("entity",{}) or hit) if isinstance(hit,dict) else (getattr(hit,"entity",{}) or {})
870
+ distance = hit.get("distance",0.0) if isinstance(hit,dict) else getattr(hit,"distance",0.0)
871
+ text = entity.get("text","")
872
+ if text and text not in seen:
873
+ item = {"text":text,"score":float(distance),"retrieval_type":"dense"}
874
+ if domain_name == "Legal_Contracts":
875
+ item.update({"contract_id":entity.get("contract_id"),
876
+ "source_doc_id":entity.get("source_doc_id"),
877
+ "source_hash":entity.get("source_hash")})
878
+ retrieved.append(item); seen.add(text)
879
+
880
+ if not retrieved: return []
881
+
882
+ # Rerank β†’ trim to top_k
883
+ if ENABLE_RERANKING: retrieved = rerank_documents(eff_q, retrieved, top_k)
884
+ else: retrieved = retrieved[:top_k]
885
+
886
+ # Repack (reorder for LLM attention)
887
+ if ENABLE_REPACKING: retrieved = repack_documents(retrieved, REPACK_STRATEGY)
888
+
889
+ # Summarize / compress context
890
+ if ENABLE_SUMMARIZATION:
891
+ orig = retrieved
892
+ summarized = summarize_docs(eff_q, retrieved, em=em, llm_client=llm)
893
+ if not summarized: return orig
894
+ scores_list = [d.get("rerank_score",d.get("score",0.0)) for d in retrieved if isinstance(d,dict)]
895
+ avg = float(np.mean(scores_list)) if scores_list else 1.0
896
+ mx = float(max(scores_list)) if scores_list else avg
897
+ rtype = retrieved[0].get("reranker_type") if retrieved and isinstance(retrieved[0],dict) else None
898
+ rettype = retrieved[0].get("retrieval_type") if retrieved and isinstance(retrieved[0],dict) else None
899
+ # Preserve Legal metadata from the first (highest-relevance) source chunk
900
+ smeta = {}
901
+ if domain_name == "Legal_Contracts" and retrieved and isinstance(retrieved[0],dict):
902
+ smeta = {k: retrieved[0].get(k) for k in ("contract_id","source_doc_id","source_hash")}
903
+ retrieved = [{"text":s,"score":avg,"rerank_score":mx,"reranker_type":rtype,
904
+ "retrieval_type":rettype,"summarized":True,"summary_type":SUMMARIZATION_TYPE,**smeta}
905
+ for s in summarized if s and str(s).strip()]
906
+ if not retrieved: return orig
907
+
908
+ return retrieved
909
+
910
+
911
+ # ── Prompt / generation ────────────────────────────────────────────────────────
912
+
913
+ def _build_prompt(context, question, strategy="short"):
914
+ context = _sanitize(context); question = _sanitize(question)
915
+ if strategy == "short":
916
+ return f"Answer the question using the provided context.\n\nContext:\n{context}\n\nQuestion:\n{question}".strip()
917
+ elif strategy == "long":
918
+ return ("You are a chatbot providing answers to user queries. Use the context documents to answer the question.\n"
919
+ 'If the documents do not provide enough information, say "The documents are missing some of the information required to answer the question."\n'
920
+ f"Do not use external knowledge. Do not make up an answer.\n\nContext Documents:\n{context}\n\nQuestion: {question}").strip()
921
+ elif strategy == "long_cot":
922
+ return ("You are a chatbot providing answers to user queries. Use the context documents to answer the question.\n"
923
+ 'If the documents do not provide enough information, say "The documents are missing some of the information required to answer the question."\n'
924
+ f"Do not use external knowledge. Think step by step and quote documents when necessary.\n\nContext Documents:\n{context}\n\nQuestion: {question}").strip()
925
+ raise ValueError(f"Unknown PROMPT_STRATEGY: {strategy}")
926
+
927
+ def ask_rag(context, question, llm_client, strategy=None):
928
+ strategy = strategy or PROMPT_STRATEGY
929
+ resp = llm_client.chat.completions.create(
930
+ model=MODEL_NAME,
931
+ messages=[{"role":"system","content":"You are a helpful RAG assistant"},
932
+ {"role":"user","content":_build_prompt(context, question, strategy)}],
933
+ temperature=0.3,
934
+ )
935
+ return _safe_message_content(resp)
936
+
937
+
938
+ print("Pipeline functions defined.")
939
+
940
+ # ── Cell 6: Initialize LLM client ─────────────────────────────────────────────
941
+ llm_client = get_llm_client()
942
+ print(f"LLM client ready. Provider: {LLM_PROVIDER}")
943
+
944
+
945
+ # ── Cell 7: HF filesystem + embedding model loading ───────────────────────────
946
+ # NOTE: _hf_fs must be initialized here before Cell 8 calls hf_path_exists()
947
+ from sentence_transformers import SentenceTransformer
948
+ import torch
949
+ from huggingface_hub import HfFileSystem
950
+
951
+ device = "cuda" if torch.cuda.is_available() else "cpu"
952
+ _hf_fs = HfFileSystem()
953
+ print(f"Device: {device} | HfFileSystem ready")
954
+
955
+ # Determine which embedding types to load
956
+ if DOMAIN_EMBEDDING_RECOMMENDATION:
957
+ models_to_load = sorted(set(DOMAIN_EMBEDDING_RECOMMENDATION.values()))
958
+ print(f"Domain-aware mode β€” loading: {models_to_load}")
959
+ else:
960
+ # Legacy single-model mode: always load EMBEDDING_TYPE
961
+ models_to_load = [EMBEDDING_TYPE]
962
+ print(f"Single-model mode β€” loading: {models_to_load}")
963
+
964
+ loaded_embedding_models = {}
965
+ for emb_type in models_to_load:
966
+ if emb_type not in EMBED_MODELS:
967
+ raise ValueError(f"Unknown embedding type: {emb_type}. Available: {list(EMBED_MODELS.keys())}")
968
+ model_name = EMBED_MODELS[emb_type]
969
+ print(f" Loading {emb_type}: {model_name}")
970
+ m = SentenceTransformer(model_name, device=device)
971
+ dim = m.get_sentence_embedding_dimension() if hasattr(m, "get_sentence_embedding_dimension") else getattr(m, "get_embedding_dimension", lambda: "?")()
972
+ loaded_embedding_models[emb_type] = m
973
+ print(f" OK β€” dim={dim}")
974
+
975
+ if not loaded_embedding_models:
976
+ raise RuntimeError("No embedding models were loaded. Check DOWNLOAD_MODE and EMBED_MODELS.")
977
+
978
+ # Fallback single embed_model used by RECOMP summarization
979
+ embed_model = loaded_embedding_models.get(EMBEDDING_TYPE) or next(iter(loaded_embedding_models.values()))
980
+ print(f"\nAll embedding models ready. Fallback embed_model: {EMBEDDING_TYPE}")
981
+
982
+ # ── Cell 8: Download Milvus DBs + Legal mapping from HuggingFace ──────────────
983
+ # Uses download_indexes() from Cell 5 which mirrors the Advanced notebook logic:
984
+ # preferred path: BUCKET_PREFIX/{embedding_type}_{index_version}/{domain}.db
985
+ # fallback path: BUCKET_PREFIX/{embedding_type}/{domain}.db
986
+ os.makedirs(MILVUS_DIR, exist_ok=True)
987
+
988
+ def download_indexes():
989
+ """Download all domain DBs using domain-aware embedding types."""
990
+ report = []
991
+ for domain_name in DOMAIN_NAMES:
992
+ embedding_type = get_embedding_type_for_domain(domain_name)
993
+ local_target = get_db_path(domain_name, embedding_type=embedding_type, index_version=INDEX_VERSION)
994
+ os.makedirs(os.path.dirname(local_target), exist_ok=True)
995
+
996
+ preferred_remote = f"{BUCKET_PREFIX}/{get_index_folder(embedding_type, INDEX_VERSION)}/{domain_name}.db"
997
+ fallback_remote = f"{BUCKET_PREFIX}/{embedding_type}/{domain_name}.db"
998
+
999
+ # Remove stale file before re-download
1000
+ if os.path.exists(local_target):
1001
+ if os.path.isdir(local_target): shutil.rmtree(local_target)
1002
+ else: os.remove(local_target)
1003
+
1004
+ selected_remote, source_type = None, None
1005
+ if hf_path_exists(preferred_remote):
1006
+ selected_remote = preferred_remote
1007
+ source_type = get_index_folder(embedding_type, INDEX_VERSION)
1008
+ elif hf_path_exists(fallback_remote):
1009
+ selected_remote = fallback_remote
1010
+ source_type = embedding_type
1011
+
1012
+ if selected_remote is None:
1013
+ print(f" MISSING: {domain_name} ({embedding_type})")
1014
+ report.append({"domain":domain_name,"status":"missing","source_type":None,"local_target":local_target})
1015
+ continue
1016
+
1017
+ print(f" Downloading: {domain_name} [{source_type}]")
1018
+ try:
1019
+ _hf_fs.get(selected_remote, local_target, recursive=True)
1020
+ ok = os.path.exists(local_target) and os.path.getsize(local_target) > 0
1021
+ status = "downloaded" if ok else "empty"
1022
+ print(f" {'OK' if ok else 'EMPTY'}: {local_target}")
1023
+ report.append({"domain":domain_name,"status":status,"source_type":source_type,"local_target":local_target})
1024
+ except Exception as e:
1025
+ print(f" FAILED: {e}")
1026
+ report.append({"domain":domain_name,"status":"failed","source_type":source_type,"local_target":local_target,"error":str(e)})
1027
+ return report
1028
+
1029
+ print("Downloading domain DBs...")
1030
+ dl_report = download_indexes()
1031
+
1032
+ # ── Download Legal sampleβ†’contract mapping ────────────────────────────────────
1033
+ legal_emb = get_embedding_type_for_domain("Legal_Contracts")
1034
+ legal_folder = get_index_folder(legal_emb, INDEX_VERSION)
1035
+ remote_mapping = f"{BUCKET_PREFIX}/{legal_folder}/legal_sample_to_contract_id.json"
1036
+ local_mapping = os.path.join(MILVUS_DIR, legal_folder, "legal_sample_to_contract_id.json")
1037
+ os.makedirs(os.path.dirname(local_mapping), exist_ok=True)
1038
+ print(f"\nDownloading Legal mapping: {remote_mapping}")
1039
+ try:
1040
+ _hf_fs.get(remote_mapping, local_mapping)
1041
+ if os.path.exists(local_mapping) and os.path.getsize(local_mapping) > 0:
1042
+ print(f" OK: {local_mapping}")
1043
+ else:
1044
+ print(" WARNING: Legal mapping download failed or empty")
1045
+ except Exception as e:
1046
+ print(f" WARNING: Could not download Legal mapping: {e}")
1047
+
1048
+ # ── Sanity check ──────────────────────────────────────────────────────────────
1049
+ print("\nSanity check:")
1050
+ for domain in DOMAIN_NAMES:
1051
+ p = get_db_path(domain, get_embedding_type_for_domain(domain), INDEX_VERSION)
1052
+ print(f" {'OK' if os.path.exists(p) else 'MISSING'}: {p}")
1053
+
1054
+
1055
+ # ── Cell 9: Open Milvus clients + BM25 indexes + Legal mapping ────────────────
1056
+ def load_milvus_clients():
1057
+ global milvus_clients, LEGAL_SAMPLE_TO_CONTRACT_ID
1058
+ milvus_clients = {}
1059
+ for domain_name in DOMAIN_NAMES:
1060
+ embedding_type = get_embedding_type_for_domain(domain_name)
1061
+ db_path = get_db_path(domain_name, embedding_type=embedding_type, index_version=INDEX_VERSION)
1062
+ col = domain_name.lower()
1063
+ if not os.path.exists(db_path):
1064
+ print(f" DB not found, skipping: {db_path}")
1065
+ continue
1066
+ try:
1067
+ client = MilvusClient(db_path)
1068
+ if not client.has_collection(col):
1069
+ print(f" Collection missing in {db_path}, skipping")
1070
+ continue
1071
+ client.load_collection(col)
1072
+ stats = client.get_collection_stats(col)
1073
+ rows = int(stats.get("row_count", 0))
1074
+ milvus_clients[domain_name] = client
1075
+ print(f" {domain_name}: {rows:,} rows [{embedding_type}]")
1076
+ except Exception as e:
1077
+ print(f" Failed to open {domain_name}: {e}")
1078
+ print(f"\nLoaded {len(milvus_clients)} domain clients: {list(milvus_clients.keys())}")
1079
+
1080
+ # Legal sample→contract mapping
1081
+ LEGAL_SAMPLE_TO_CONTRACT_ID = load_legal_sample_to_contract_mapping()
1082
+
1083
+ load_milvus_clients()
1084
+
1085
+ # Build BM25 indexes (stores contract_ids for Legal)
1086
+ build_all_bm25_indexes(milvus_clients)
1087
+ print("BM25 indexes ready.")
1088
+
1089
+ # ── Cell 10: Load RAGBench (test split only) + sample catalogue ───────────────
1090
+ DATASET_BY_DOMAIN = {
1091
+ "Bio_Medical": ["covidqa", "pubmedqa"],
1092
+ "General_Knowledge": ["expertqa", "hagrid", "hotpotqa", "msmarco"],
1093
+ "Customer_Support": ["delucionqa", "emanual", "techqa"],
1094
+ "Finance": ["finqa", "tatqa"],
1095
+ "Legal_Contracts": ["cuad"],
1096
+ }
1097
+
1098
+ # sample_store[domain][dataset] = list of row dicts from the test split
1099
+ sample_store = {}
1100
+
1101
+ def load_ragbench(domains=None):
1102
+ global ragbench_by_domain, sample_store
1103
+ domains = domains or list(DATASET_BY_DOMAIN.keys())
1104
+ for domain in domains:
1105
+ ragbench_by_domain[domain] = {}
1106
+ sample_store[domain] = {}
1107
+ for ds_name in DATASET_BY_DOMAIN.get(domain, []):
1108
+ try:
1109
+ ds = load_dataset("rungalileo/ragbench", ds_name)
1110
+ ragbench_by_domain[domain][ds_name] = ds
1111
+ if "test" not in ds:
1112
+ print(f" WARNING: no 'test' split for {domain}/{ds_name}, skipping")
1113
+ continue
1114
+ rows = []
1115
+ for idx, row in enumerate(ds["test"]):
1116
+ # For Legal_Contracts resolve contract_id from mapping
1117
+ contract_id = None
1118
+ if domain == "Legal_Contracts":
1119
+ try: contract_id = get_contract_id_for_legal_sample(idx)
1120
+ except Exception: pass
1121
+ rows.append({
1122
+ "idx": idx,
1123
+ "question": row.get("question", ""),
1124
+ "response": row.get("response", ""),
1125
+ "documents": row.get("documents", []),
1126
+ "contract_id": contract_id,
1127
+ "gold_relevance": row.get("relevance_score"),
1128
+ "gold_utilization": row.get("utilization_score"),
1129
+ "gold_completeness": row.get("completeness_score"),
1130
+ "gold_adherence": row.get("adherence_score"),
1131
+ })
1132
+ sample_store[domain][ds_name] = rows
1133
+ print(f" Loaded: {domain}/{ds_name} test rows={len(rows)}")
1134
+ except Exception as e:
1135
+ print(f" Failed: {domain}/{ds_name}: {e}")
1136
+ print(f"\nRAGBench loaded (test only). Domains: {list(sample_store.keys())}")
1137
+
1138
+ load_ragbench()
1139
+
1140
+
1141
+ # ── Helpers for cascading dropdowns ───────────────────────────────────────────
1142
+
1143
+ def get_datasets_for_domain(domain):
1144
+ return list(sample_store.get(domain, {}).keys())
1145
+
1146
+ def get_sample_ids_for_dataset(domain, dataset):
1147
+ """
1148
+ Returns label strings for the Sample ID dropdown.
1149
+ For Legal_Contracts uses 'Contract ID' wording and shows contract hash.
1150
+ """
1151
+ rows = sample_store.get(domain, {}).get(dataset, [])
1152
+ is_legal = (domain == "Legal_Contracts")
1153
+ labels = []
1154
+ for r in rows:
1155
+ q = r["question"]
1156
+ cid = r.get("contract_id")
1157
+ if is_legal and cid:
1158
+ prefix = f"Contract {str(cid)[:8]}… | idx={r['idx']} – "
1159
+ else:
1160
+ prefix = f"{r['idx']} – "
1161
+ labels.append(f"{prefix}{q[:70]}{'…' if len(q)>70 else ''}")
1162
+ return labels
1163
+
1164
+ def get_row_by_label(domain, dataset, label):
1165
+ """Retrieve a stored row dict from a label string."""
1166
+ if not label: return None
1167
+ rows = sample_store.get(domain, {}).get(dataset, [])
1168
+ is_legal = (domain == "Legal_Contracts")
1169
+ # Legal labels: "Contract <hash8>… | idx=N – ..."
1170
+ # Regular labels: "N – ..."
1171
+ if is_legal:
1172
+ m = re.search(r"idx=(\d+)", label)
1173
+ try: idx = int(m.group(1)) if m else int(label.split("–")[0].strip())
1174
+ except ValueError: return None
1175
+ else:
1176
+ try: idx = int(label.split("–")[0].strip())
1177
+ except ValueError: return None
1178
+ return next((r for r in rows if r["idx"] == idx), None)
1179
+
1180
+ print("Sample catalogue ready.")
1181
+
1182
+ # ── Cell 11: Evaluation helpers + all Gradio handlers ─────────────────────────
1183
+
1184
+ # ── Judge / evaluation ─────────────────────────────────────────────────────────
1185
+
1186
+ def build_keyed_response(answer):
1187
+ return {f"r_{i}": s for i, s in enumerate(split_into_sentences(answer))}
1188
+
1189
+ def build_sentence_keyed_docs(retrieved_docs):
1190
+ keyed = {}
1191
+ for di, doc in enumerate(retrieved_docs):
1192
+ text = doc.get("text","") if isinstance(doc, dict) else doc
1193
+ for si, s in enumerate(split_into_sentences(text)):
1194
+ keyed[f"{di}_{si}"] = s
1195
+ return keyed
1196
+
1197
+ def build_evaluation_prompt(documents_text, question, answer_text):
1198
+ return f"""Evaluate the RAG response using the provided documents.
1199
+
1200
+ Documents (sentence-keyed):
1201
+ {documents_text}
1202
+
1203
+ Question:
1204
+ {question}
1205
+
1206
+ Response (sentence-keyed):
1207
+ {answer_text}
1208
+
1209
+ Return ONLY valid JSON:
1210
+ {{
1211
+ "overall_supported": true,
1212
+ "all_relevant_sentence_keys": ["0_0"],
1213
+ "all_utilized_sentence_keys": ["0_0"],
1214
+ "sentence_support_information": [
1215
+ {{"response_sentence_key": "r_0", "supporting_sentence_keys": ["0_0"], "fully_supported": true}}
1216
+ ]
1217
+ }}
1218
+ Rules: document keys look like 0_0; response keys like r_0. Return only JSON.""".strip()
1219
+
1220
+ def ask_judge(prompt, llm_client, judge_model, max_retries=5):
1221
+ last_error = None
1222
+ for attempt in range(max_retries):
1223
+ try:
1224
+ resp = llm_client.chat.completions.create(
1225
+ model=judge_model,
1226
+ messages=[
1227
+ {"role":"system","content":"You are a strict RAG evaluation judge. Return ONLY valid JSON. No markdown. No <think> tags."},
1228
+ {"role":"user","content":_sanitize(prompt)},
1229
+ ],
1230
+ temperature=0.0, max_tokens=3000,
1231
+ )
1232
+ return _safe_message_content(resp)
1233
+ except Exception as e:
1234
+ last_error = e; msg = str(e)
1235
+ wait = 2**attempt
1236
+ if "429" in msg or "rate_limit" in msg:
1237
+ m = re.search(r"try again in ([\\d.]+)s", msg)
1238
+ if m: wait = float(m.group(1))
1239
+ elif not any(x in msg for x in ["503","502","504","over capacity","gateway"]):
1240
+ raise
1241
+ time.sleep(wait + random.uniform(0.1, 0.5))
1242
+ raise RuntimeError(f"Judge failed after {max_retries} retries: {last_error}")
1243
+
1244
+ def parse_judge_json(raw):
1245
+ if not raw: raise ValueError("Judge output empty")
1246
+ cleaned = re.sub(r"<think>.*?</think>","",str(raw),flags=re.DOTALL).strip()
1247
+ cleaned = cleaned.replace("```json","").replace("```","").strip()
1248
+ s, e = cleaned.find("{"), cleaned.rfind("}")
1249
+ if s == -1 or e == -1: raise ValueError(f"No JSON: {cleaned[:300]}")
1250
+ cleaned = cleaned[s:e+1]
1251
+ cleaned = re.sub(r"}\s*{","}, {",cleaned)
1252
+ cleaned = re.sub(r",\s*([}\]])",r"\1",cleaned)
1253
+ return json.loads(cleaned)
1254
+
1255
+ def evaluate_ragbench_json(judge_json, keyed_docs):
1256
+ vk = set(keyed_docs.keys())
1257
+ rel = set(judge_json.get("all_relevant_sentence_keys", [])) & vk
1258
+ utl = set(judge_json.get("all_utilized_sentence_keys", [])) & vk
1259
+ ovl = rel & utl; n = len(vk)
1260
+ return {
1261
+ "adherence_score": int(bool(judge_json.get("overall_supported", False))),
1262
+ "hallucination_flag": 1 - int(bool(judge_json.get("overall_supported", False))),
1263
+ "relevance_score": float(np.clip(len(rel)/n if n else 0, 0, 1)),
1264
+ "utilization_score": float(np.clip(len(utl)/n if n else 0, 0, 1)),
1265
+ "completeness_score": float(np.clip(len(ovl)/len(rel) if rel else 0, 0, 1)),
1266
+ }
1267
+
1268
+
1269
+ # ── Source badge helper ────────────────────────────────────────────────────────
1270
+
1271
+ def _source_badge(source, model, extra=None):
1272
+ parts = [f"[Source: {source} | model: {model}"]
1273
+ if extra: parts += [f" | {k}: {v}" for k, v in extra.items()]
1274
+ parts.append("]")
1275
+ return "".join(parts)
1276
+
1277
+
1278
+ # ── DB status helper ───────────────────────────────────────────────────────────
1279
+
1280
+ def db_status_md():
1281
+ if not milvus_clients:
1282
+ return ("> **No vector DBs loaded.** Re-run Cell 8 (download) then Cell 9 (open), then re-run Cell 12.")
1283
+ rows = []
1284
+ for d in sorted(milvus_clients.keys()):
1285
+ emb = get_embedding_type_for_domain(d)
1286
+ rows.append(f"`{d}` ({emb})")
1287
+ return f"> **Loaded domains ({len(milvus_clients)}):** {', '.join(rows)}"
1288
+
1289
+
1290
+ # ── Config applier ─────────────────────────────────────────────────────────────
1291
+
1292
+ def apply_config(llm_choice, embed_choice,
1293
+ enable_hybrid, enable_hyde, enable_reranking, reranker_type,
1294
+ enable_rrf, rrf_k,
1295
+ enable_repacking, repack_strategy,
1296
+ enable_summarization, summarization_type,
1297
+ prompt_strategy, hybrid_alpha, top_k,
1298
+ enable_query_classification, enable_query_rewriting, enable_query_decomp):
1299
+ global MODEL_NAME, EMBEDDING_TYPE, embed_model
1300
+ global ENABLE_HYBRID, ENABLE_HYDE, ENABLE_RERANKING, RERANKER_TYPE
1301
+ global ENABLE_RRF, RRF_K
1302
+ global ENABLE_REPACKING, REPACK_STRATEGY, ENABLE_SUMMARIZATION, SUMMARIZATION_TYPE
1303
+ global PROMPT_STRATEGY, HYBRID_ALPHA
1304
+ global ENABLE_QUERY_CLASSIFICATION, ENABLE_QUERY_REWRITING, ENABLE_QUERY_DECOMPOSITION
1305
+
1306
+ MODEL_NAME = llm_choice
1307
+ ENABLE_HYBRID = enable_hybrid
1308
+ ENABLE_HYDE = enable_hyde
1309
+ ENABLE_RERANKING = enable_reranking
1310
+ RERANKER_TYPE = reranker_type
1311
+ ENABLE_RRF = enable_rrf
1312
+ RRF_K = int(rrf_k)
1313
+ ENABLE_REPACKING = enable_repacking
1314
+ REPACK_STRATEGY = repack_strategy
1315
+ ENABLE_SUMMARIZATION = enable_summarization
1316
+ SUMMARIZATION_TYPE = summarization_type
1317
+ PROMPT_STRATEGY = prompt_strategy
1318
+ HYBRID_ALPHA = float(hybrid_alpha)
1319
+ ENABLE_QUERY_CLASSIFICATION = enable_query_classification
1320
+ ENABLE_QUERY_REWRITING = enable_query_rewriting
1321
+ ENABLE_QUERY_DECOMPOSITION = enable_query_decomp
1322
+
1323
+ # Update single fallback embed_model if user changes embedding choice
1324
+ if embed_choice != EMBEDDING_TYPE:
1325
+ EMBEDDING_TYPE = embed_choice
1326
+ if embed_choice in loaded_embedding_models:
1327
+ embed_model = loaded_embedding_models[embed_choice]
1328
+ else:
1329
+ print(f"Embedding type '{embed_choice}' not preloaded; loading now...")
1330
+ embed_model = SentenceTransformer(EMBED_MODELS[embed_choice], device=device)
1331
+ loaded_embedding_models[embed_choice] = embed_model
1332
+
1333
+
1334
+ # ── Cascading dropdown callbacks ───────────────────────────────────────────────
1335
+
1336
+ _NONE_DOMAIN = "None (direct LLM, no retrieval)"
1337
+
1338
+ def on_domain_change(domain):
1339
+ if domain == _NONE_DOMAIN:
1340
+ return gr.update(choices=[], value=None), gr.update(choices=[], value=None), gr.update()
1341
+ datasets = get_datasets_for_domain(domain)
1342
+ ds = datasets[0] if datasets else None
1343
+ sample_ids = get_sample_ids_for_dataset(domain, ds) if ds else []
1344
+ label_text = "Contract ID (contract hash | idx – question preview)" if domain == "Legal_Contracts" else "Sample ID (idx – question preview)"
1345
+ return (
1346
+ gr.update(choices=datasets, value=ds),
1347
+ gr.update(choices=sample_ids, value=None, label=label_text),
1348
+ gr.update(value=""),
1349
+ )
1350
+
1351
+ def on_dataset_change(domain, dataset):
1352
+ if domain == _NONE_DOMAIN or not dataset:
1353
+ return gr.update(choices=[], value=None), gr.update(value="")
1354
+ sample_ids = get_sample_ids_for_dataset(domain, dataset)
1355
+ label_text = "Contract ID (contract hash | idx – question preview)" if domain == "Legal_Contracts" else "Sample ID (idx – question preview)"
1356
+ return gr.update(choices=sample_ids, value=None, label=label_text), gr.update(value="")
1357
+
1358
+ def on_sample_select(domain, dataset, label):
1359
+ if domain == _NONE_DOMAIN or not label: return gr.update()
1360
+ row = get_row_by_label(domain, dataset, label)
1361
+ if row is None: return gr.update()
1362
+ return gr.update(value=row["question"])
1363
+
1364
+
1365
+ # ── Chunk display helpers ──────────────────────────────────────────────────────
1366
+
1367
+ def _format_chunks(docs, title="Retrieved"):
1368
+ if not docs: return f"_No documents for {title}._"
1369
+ parts = []
1370
+ for i, doc in enumerate(docs):
1371
+ if isinstance(doc, dict):
1372
+ text = doc.get("text", str(doc))
1373
+ score = doc.get("rerank_score", doc.get("score", 0.0))
1374
+ tags = []
1375
+ if doc.get("summarized"): tags.append(f"summarized/{doc.get('summary_type','')}")
1376
+ if doc.get("reranker_type"): tags.append(f"reranked/{doc.get('reranker_type','')}")
1377
+ if doc.get("retrieval_type"): tags.append(doc.get("retrieval_type",""))
1378
+ if ENABLE_HYBRID and ENABLE_RRF:
1379
+ tags.append(f"RRF score={score:.4f}")
1380
+ elif ENABLE_HYBRID:
1381
+ tags.append(f"d={doc.get('dense_score',0):.3f} b={doc.get('bm25_score',0):.3f}")
1382
+ if doc.get("contract_id"): tags.append(f"contract={str(doc.get('contract_id',''))[:8]}…")
1383
+ tag_str = f" `{' | '.join(tags)}`" if tags else ""
1384
+ else:
1385
+ text, score, tag_str = str(doc), 0.0, ""
1386
+ parts.append(f"**{title} Chunk {i+1}** β€” score: `{score:.4f}`{tag_str}\n\n{text}")
1387
+ return "\n\n---\n\n".join(parts)
1388
+
1389
+ def _format_gt_docs(doc_list):
1390
+ if not doc_list: return "_No ground-truth documents stored for this sample._"
1391
+ parts = []
1392
+ for i, text in enumerate(doc_list):
1393
+ parts.append(f"**GT Doc {i+1}**\n\n{str(text)}")
1394
+ return "\n\n---\n\n".join(parts)
1395
+
1396
+
1397
+ # ── Main run handler ───────────────────────────────────────────────────────────
1398
+
1399
+ def run_query(
1400
+ query, domain,
1401
+ dataset_sel, sample_label,
1402
+ llm_choice, judge_llm_choice, embed_choice,
1403
+ enable_hybrid, enable_hyde, enable_reranking, reranker_type,
1404
+ enable_rrf, rrf_k,
1405
+ enable_repacking, repack_strategy,
1406
+ enable_summarization, summarization_type,
1407
+ prompt_strategy, hybrid_alpha, top_k,
1408
+ enable_query_classification, enable_query_rewriting, enable_query_decomp,
1409
+ run_judge,
1410
+ ):
1411
+ query = _sanitize(query)
1412
+ if not query.strip():
1413
+ return ("Please enter a query.",) + ("",)*4
1414
+ if llm_client is None:
1415
+ return ("LLM client not initialised. Re-run Cell 6 then Cell 12.",) + ("",)*4
1416
+
1417
+ apply_config(
1418
+ llm_choice, embed_choice,
1419
+ enable_hybrid, enable_hyde, enable_reranking, reranker_type,
1420
+ enable_rrf, rrf_k,
1421
+ enable_repacking, repack_strategy,
1422
+ enable_summarization, summarization_type,
1423
+ prompt_strategy, float(hybrid_alpha), int(top_k),
1424
+ enable_query_classification, enable_query_rewriting, enable_query_decomp,
1425
+ )
1426
+
1427
+ # ── Domain = None β†’ direct LLM ──────────────────────────────���────────────
1428
+ if domain == _NONE_DOMAIN or not domain:
1429
+ try:
1430
+ direct_ans = _safe_message_content(llm_client.chat.completions.create(
1431
+ model=MODEL_NAME,
1432
+ messages=[{"role":"system","content":"You are a helpful assistant."},
1433
+ {"role":"user","content":query}],
1434
+ temperature=0.3, max_tokens=800,
1435
+ ))
1436
+ except Exception as e: direct_ans = f"Direct LLM error: {e}"
1437
+ badge = _source_badge("Direct LLM (no retrieval)", MODEL_NAME)
1438
+ note = "_[Domain = None β€” answered directly by LLM without vector DB retrieval]_"
1439
+ return f"{badge}\n\n{direct_ans}", note, note, note, note
1440
+
1441
+ # ── Guard ──────────────────────────────────────────────────────────────────
1442
+ if not milvus_clients:
1443
+ return ("No vector DBs loaded. Re-run Cell 8 then Cell 9, then re-run Cell 12.",) + ("",)*4
1444
+ if domain not in milvus_clients:
1445
+ return (f"Domain '{domain}' not loaded. Loaded: {list(milvus_clients.keys())}",) + ("",)*4
1446
+
1447
+ # ── Query Classification ──────────────────────────────────────────────────
1448
+ route = classify_query(query, domain_name=domain)
1449
+ if route == "LLM":
1450
+ try:
1451
+ direct_ans = _safe_message_content(llm_client.chat.completions.create(
1452
+ model=MODEL_NAME,
1453
+ messages=[{"role":"system","content":"You are a concise factual assistant."},
1454
+ {"role":"user","content":query}],
1455
+ temperature=0.2, max_tokens=500,
1456
+ ))
1457
+ except Exception as e: direct_ans = f"Direct LLM error: {e}"
1458
+ badge = _source_badge("Direct LLM", MODEL_NAME)
1459
+ note = "_[Query Classifier routed to direct LLM β€” no retrieval]_"
1460
+ return f"{badge}\n\n{direct_ans}", note, note, note, note
1461
+
1462
+ # ── Query Rewriting + Decomposition ──────────────────────────────────────
1463
+ rewritten = rewrite_query(query, domain, llm_client) if ENABLE_QUERY_REWRITING else query
1464
+ subqueries = decompose_query(rewritten, llm_client, domain=domain) if ENABLE_QUERY_DECOMPOSITION else [rewritten]
1465
+
1466
+ # ── Resolve row + contract_id for Legal ───────────────────────────────────
1467
+ row = get_row_by_label(domain, dataset_sel, sample_label) if sample_label else None
1468
+ if row is None:
1469
+ for ds_name, rows in sample_store.get(domain, {}).items():
1470
+ match = next((r for r in rows if r["question"].strip().lower() == query.strip().lower()), None)
1471
+ if match: row = match; break
1472
+
1473
+ legal_contract_id = None
1474
+ legal_sample_id = None
1475
+ if domain == "Legal_Contracts" and row is not None:
1476
+ legal_contract_id = row.get("contract_id")
1477
+ legal_sample_id = row.get("idx")
1478
+
1479
+ # ── Retrieve + Generate ───────────────────────────────────────────────────
1480
+ all_retrieved, all_answers = [], []
1481
+ for sq in subqueries:
1482
+ try:
1483
+ docs = retrieve(sq, domain,
1484
+ llm_client=llm_client, top_k=int(top_k),
1485
+ sample_id=legal_sample_id, contract_id=legal_contract_id)
1486
+ except Exception as e:
1487
+ return (f"Retrieval error: {e}",) + ("",)*4
1488
+ if not docs: continue
1489
+ all_retrieved.extend(docs)
1490
+ ctx = _sanitize("\n\n".join(d.get("text","") if isinstance(d,dict) else d for d in docs))
1491
+ sq = _sanitize(sq)
1492
+ try:
1493
+ all_answers.append(ask_rag(ctx, sq, llm_client, strategy=PROMPT_STRATEGY))
1494
+ except Exception as e:
1495
+ return (f"Generation error: {e}",) + ("",)*4
1496
+
1497
+ if not all_retrieved:
1498
+ return ("No documents retrieved.",) + ("",)*4
1499
+
1500
+ raw_answer = "\n\n".join(all_answers)
1501
+
1502
+ # ── Source badge ──────────────────────────────────────────────────────────
1503
+ active = {"prompt": PROMPT_STRATEGY, "chunks": len(all_retrieved)}
1504
+ if ENABLE_HYBRID:
1505
+ active["hybrid"] = f"RRF(k={RRF_K})" if ENABLE_RRF else f"alpha={HYBRID_ALPHA}"
1506
+ if ENABLE_HYDE: active["hyde"] = "on"
1507
+ if ENABLE_RERANKING: active["rerank"] = RERANKER_TYPE
1508
+ if ENABLE_SUMMARIZATION: active["summ"] = SUMMARIZATION_TYPE
1509
+ if ENABLE_REPACKING: active["repack"] = REPACK_STRATEGY
1510
+ if len(subqueries) > 1: active["subq"] = len(subqueries)
1511
+ if legal_contract_id: active["contract"] = str(legal_contract_id)[:8] + "…"
1512
+ rag_response = f"{_source_badge('RAG', MODEL_NAME, extra=active)}\n\n{raw_answer}"
1513
+
1514
+ ground_truth = row["response"] if row else "_(no matching sample found)_"
1515
+ gt_docs_md = _format_gt_docs(row["documents"] if row else [])
1516
+ rag_docs_md = _format_chunks(all_retrieved, title="RAG")
1517
+
1518
+ # ── Judge evaluation ──────────────────────────────────────────────────────
1519
+ metrics_md = "_Judge evaluation not requested._"
1520
+ if run_judge:
1521
+ try:
1522
+ keyed_docs = build_sentence_keyed_docs(all_retrieved)
1523
+ keyed_answer = build_keyed_response(raw_answer)
1524
+ docs_text = _sanitize("\n".join(f"{k}: {v}" for k,v in keyed_docs.items()))
1525
+ ans_text = _sanitize("\n".join(f"{k}: {v}" for k,v in keyed_answer.items()))
1526
+ raw = ask_judge(build_evaluation_prompt(docs_text, query, ans_text), llm_client, judge_llm_choice)
1527
+ pred = evaluate_ragbench_json(parse_judge_json(raw), keyed_docs)
1528
+ gold = {k: row.get(f"gold_{k}") for k in ("relevance","utilization","completeness","adherence")} if row else {}
1529
+ def _f(v): return f"{v:.3f}" if isinstance(v, float) else (str(v) if v is not None else "β€”")
1530
+ metrics_md = "\n".join([
1531
+ "| Metric | Predicted | Gold |",
1532
+ "|--------|-----------|------|",
1533
+ f"| Relevance | {_f(pred['relevance_score'])} | {_f(gold.get('relevance'))} |",
1534
+ f"| Utilization | {_f(pred['utilization_score'])} | {_f(gold.get('utilization'))} |",
1535
+ f"| Completeness | {_f(pred['completeness_score'])} | {_f(gold.get('completeness'))} |",
1536
+ f"| Adherence | {_f(pred['adherence_score'])} | {_f(gold.get('adherence'))} |",
1537
+ f"| Hallucination| {_f(pred['hallucination_flag'])} | β€” |",
1538
+ ])
1539
+ except Exception as e:
1540
+ metrics_md = f"Judge error: {e}"
1541
+
1542
+ return ground_truth, rag_response, gt_docs_md, rag_docs_md, metrics_md
1543
+
1544
+
1545
+ print("Handlers ready.")
1546
+
1547
+ # ── Cell 12: Gradio UI ────────────────────────────────────────────────────────
1548
+
1549
+ AVAILABLE_DOMAINS = list(milvus_clients.keys())
1550
+ DEFAULT_DOMAIN = AVAILABLE_DOMAINS[0] if AVAILABLE_DOMAINS else None
1551
+
1552
+ _init_datasets = get_datasets_for_domain(DEFAULT_DOMAIN) if DEFAULT_DOMAIN else []
1553
+ _init_ds = _init_datasets[0] if _init_datasets else None
1554
+ _init_samples = get_sample_ids_for_dataset(DEFAULT_DOMAIN, _init_ds) if _init_ds else []
1555
+ _legal_first = DEFAULT_DOMAIN == "Legal_Contracts"
1556
+
1557
+ CSS = """
1558
+ footer { display: none !important; }
1559
+ """
1560
+
1561
+ with gr.Blocks(title="RAG Capstone β€” Advanced Demo") as demo:
1562
+
1563
+ # ── Header ────────────────────────────────────────────────────────────────
1564
+ gr.Markdown("# πŸ” RAG Capstone β€” Advanced Interactive Demo")
1565
+ gr.Markdown(
1566
+ f"**Provider:** {LLM_PROVIDER.upper()} &nbsp;|&nbsp; "
1567
+ f"**Index:** `{INDEX_VERSION}` &nbsp;|&nbsp; "
1568
+ "Type any question, or expand **Sample Selector** to load a test-split example."
1569
+ )
1570
+ gr.Markdown(db_status_md())
1571
+
1572
+ # ══════════════════════════════════════════════════════════════════════════
1573
+ # SECTION 1 β€” Query + Domain (always visible)
1574
+ # ══════════════════════════════════════════════════════════════════════════
1575
+ with gr.Row():
1576
+ query_input = gr.Textbox(
1577
+ lines=3,
1578
+ placeholder="Type any question here… or expand Sample Selector below to auto-fill.",
1579
+ label="Query",
1580
+ scale=4,
1581
+ )
1582
+ domain_dd = gr.Dropdown(
1583
+ choices=["None (direct LLM, no retrieval)"] + AVAILABLE_DOMAINS,
1584
+ value="None (direct LLM, no retrieval)" if not AVAILABLE_DOMAINS else DEFAULT_DOMAIN,
1585
+ label="Domain",
1586
+ info="None = direct LLM; pick a domain to run full RAG retrieval",
1587
+ scale=1,
1588
+ )
1589
+
1590
+ # ══════════════════════════════════════════════════════════════════════════
1591
+ # SECTION 2 β€” Sample Selector (collapsed, optional)
1592
+ # ═══════════════════════════════════════════════��══════════════════════════
1593
+ with gr.Accordion("πŸ“‹ Sample Selector (optional β€” expand to load a test-split example)", open=False):
1594
+ gr.Markdown(
1595
+ "_Select a sample to auto-fill Query above. "
1596
+ "For **Legal_Contracts** the dropdown shows Contract ID (hash prefix) instead of plain Sample ID β€” "
1597
+ "retrieval is automatically scoped to that contract._"
1598
+ )
1599
+ with gr.Row():
1600
+ dataset_dd = gr.Dropdown(
1601
+ choices=_init_datasets, value=_init_ds, label="Dataset", scale=1)
1602
+ sample_dd = gr.Dropdown(
1603
+ choices=_init_samples, value=None,
1604
+ label="Contract ID (contract hash | idx – question preview)" if _legal_first else "Sample ID (idx – question preview)",
1605
+ scale=4)
1606
+
1607
+ # ══════════════════════════════════════════════════════════════════════════
1608
+ # SECTION 3 β€” Control Panel
1609
+ # ══════════════════════════════════════════════════════════════════════════
1610
+ with gr.Accordion("βš™οΈ Control Panel", open=False):
1611
+ with gr.Tabs():
1612
+
1613
+ # ── Models ───────────────────────────────────────────────────────
1614
+ with gr.Tab("πŸ€– Models"):
1615
+ gr.Markdown(
1616
+ f"**Domain embedding assignment** (chunk_v5_domain_aware): \n"
1617
+ + " \n".join(
1618
+ [f"- `{d}` β†’ `{get_embedding_type_for_domain(d)}` ({EMBED_MODELS[get_embedding_type_for_domain(d)]})"
1619
+ for d in DOMAIN_NAMES]
1620
+ )
1621
+ )
1622
+ with gr.Row():
1623
+ llm_choice = gr.Dropdown(
1624
+ choices=LLM_CHOICES, value=LLM_CHOICES[0],
1625
+ label="Generator LLM", info="Produces the RAG answer")
1626
+ judge_llm_choice = gr.Dropdown(
1627
+ choices=LLM_CHOICES,
1628
+ value=LLM_CHOICES[4] if len(LLM_CHOICES) > 4 else LLM_CHOICES[-1],
1629
+ label="Judge LLM", info="Used for evaluation scoring")
1630
+ embed_choice = gr.Dropdown(
1631
+ choices=EMBEDDING_CHOICES, value=EMBEDDING_TYPE,
1632
+ label="Fallback Embedding Model",
1633
+ info="Used only when domain-specific model is unavailable")
1634
+
1635
+ # ── Query Processing ──────────────────────────────────────────────
1636
+ with gr.Tab("πŸ”„ Query Processing"):
1637
+ gr.Markdown("Applied **before** retrieval: Classify β†’ Rewrite β†’ Decompose")
1638
+ with gr.Row():
1639
+ enable_query_classification = gr.Checkbox(
1640
+ label="Query Classification", value=False,
1641
+ info="Route simple factual queries to LLM directly; benchmark domains always use RAG")
1642
+ with gr.Row():
1643
+ enable_query_rewriting = gr.Checkbox(
1644
+ label="Query Rewriting", value=False,
1645
+ info="LLM rewrites the query for better retrieval")
1646
+ enable_query_decomp = gr.Checkbox(
1647
+ label="Query Decomposition", value=False,
1648
+ info="Break multi-part queries into subqueries")
1649
+
1650
+ # ── Retrieval ─────────────────────────────────────────────────────
1651
+ with gr.Tab("πŸ”Ž Retrieval"):
1652
+ with gr.Row():
1653
+ top_k = gr.Slider(minimum=1, maximum=10, step=1, value=3,
1654
+ label="Top-K chunks returned")
1655
+ hybrid_alpha = gr.Slider(minimum=0.0, maximum=1.0, step=0.05, value=0.5,
1656
+ label="Hybrid Alpha (1=dense, 0=BM25) β€” used only when RRF is OFF")
1657
+ with gr.Row():
1658
+ enable_hybrid = gr.Checkbox(label="Hybrid Search (Dense + BM25)", value=True)
1659
+ enable_hyde = gr.Checkbox(label="HyDE (query expansion)", value=False)
1660
+
1661
+ # ── RRF ───────────────────────────────────────────────────────────
1662
+ with gr.Tab("πŸ”€ RRF"):
1663
+ gr.Markdown(
1664
+ "**Reciprocal Rank Fusion** replaces the weighted alpha fusion inside Hybrid Search. \n"
1665
+ "Score formula: `1/(k + rank_dense) + 1/(k + rank_bm25)` \n"
1666
+ "Standard literature value for k is **60** β€” lower k boosts top-ranked docs more aggressively."
1667
+ )
1668
+ with gr.Row():
1669
+ enable_rrf = gr.Checkbox(
1670
+ label="Enable RRF (replaces alpha fusion inside Hybrid Search)",
1671
+ value=True,
1672
+ info="RRF is only active when Hybrid Search is also enabled")
1673
+ rrf_k = gr.Slider(
1674
+ minimum=1, maximum=200, step=1, value=60,
1675
+ label="RRF k (rank smoothing constant)")
1676
+
1677
+ # ── Reranking ─────────────────────────────────────────────────────
1678
+ with gr.Tab("↕️ Reranking"):
1679
+ with gr.Row():
1680
+ enable_reranking = gr.Checkbox(label="Enable Reranking", value=False)
1681
+ reranker_type = gr.Radio(choices=["monot5","tilde"], value="monot5",
1682
+ label="Reranker", info="MonoT5: seq2seq | TILDE: cross-encoder")
1683
+
1684
+ # ── Repacking ─────────────────────────────────────────────────────
1685
+ with gr.Tab("πŸ“¦ Repacking"):
1686
+ with gr.Row():
1687
+ enable_repacking = gr.Checkbox(label="Enable Repacking", value=False)
1688
+ repack_strategy = gr.Radio(choices=["forward","reverse","sides"], value="sides",
1689
+ label="Strategy", info="forward | reverse | U-shape sides")
1690
+
1691
+ # ── Summarization ─────────────────────────────────────────────────
1692
+ with gr.Tab("πŸ“ Summarization"):
1693
+ with gr.Row():
1694
+ enable_summarization = gr.Checkbox(label="Enable Summarization", value=False)
1695
+ summarization_type = gr.Radio(choices=["recomp","longllmlingua"], value="recomp",
1696
+ label="Method", info="RECOMP: extractive | LLMLingua: token compression")
1697
+
1698
+ # ── Prompt ────────────────────────────────────────────────────────
1699
+ with gr.Tab("πŸ’¬ Prompt"):
1700
+ prompt_strategy = gr.Radio(
1701
+ choices=["short","long","long_cot"], value="short",
1702
+ label="Prompt Strategy",
1703
+ info="short: minimal | long: strict no-hallucination | long_cot: step-by-step")
1704
+
1705
+ # ── Judge ─────────────────────────────────────────────────────────
1706
+ with gr.Tab("βš–οΈ Judge"):
1707
+ run_judge = gr.Checkbox(
1708
+ label="Run Judge evaluation after generation", value=False,
1709
+ info="~1 extra LLM call. Gold scores shown only for preloaded samples.")
1710
+ gr.Markdown("_Judge LLM is configured in the **Models** tab._")
1711
+
1712
+ # ── Run button ────────────────────────────────────────────────────────────
1713
+ run_btn = gr.Button("β–Ά Run Query", variant="primary", size="lg")
1714
+
1715
+ # ══════════════════════════════════════════════════════════════════════════
1716
+ # SECTION 4 β€” Responses (Ground Truth LEFT, RAG RIGHT)
1717
+ # ══════════════════════════════════════════════════════════════════════════
1718
+ gr.Markdown("## πŸ’¬ Responses")
1719
+ with gr.Row(equal_height=True):
1720
+ gt_out = gr.Textbox(label="Ground Truth Response", lines=10, interactive=False, scale=1)
1721
+ rag_out = gr.Textbox(label="RAG Response", lines=10, interactive=False, scale=1)
1722
+
1723
+ # ══════════════════════════════════════════════════════════════════════════
1724
+ # SECTION 5 β€” Retrieved Documents (GT LEFT, RAG RIGHT)
1725
+ # ══════════════════════════════════════════════════════════════════════════
1726
+ gr.Markdown("## πŸ“„ Retrieved Documents")
1727
+ with gr.Row(equal_height=True):
1728
+ with gr.Column(scale=1):
1729
+ gr.Markdown("### Ground Truth Documents")
1730
+ gt_docs_out = gr.Markdown(value="_Select a preloaded sample to see GT documents._")
1731
+ with gr.Column(scale=1):
1732
+ gr.Markdown("### RAG Retrieved Documents")
1733
+ rag_docs_out = gr.Markdown(value="_Run a query to see RAG retrieved chunks._")
1734
+
1735
+ # ══════════════════════════════════════════════════════════════════════════
1736
+ # SECTION 6 β€” Metrics
1737
+ # ══════════════════════════════════════════════════════════════════════════
1738
+ with gr.Accordion("πŸ“Š Metrics (Gold vs Predicted)", open=False):
1739
+ metrics_out = gr.Markdown(value="_Enable the Judge in the Control Panel and run a query._")
1740
+
1741
+ # ── Cascading sample selector wiring ─────────────────────────────────────
1742
+ domain_dd.change(
1743
+ fn=on_domain_change, inputs=[domain_dd],
1744
+ outputs=[dataset_dd, sample_dd, query_input],
1745
+ )
1746
+ dataset_dd.change(
1747
+ fn=on_dataset_change, inputs=[domain_dd, dataset_dd],
1748
+ outputs=[sample_dd, query_input],
1749
+ )
1750
+ sample_dd.change(
1751
+ fn=on_sample_select, inputs=[domain_dd, dataset_dd, sample_dd],
1752
+ outputs=[query_input],
1753
+ )
1754
+
1755
+ # ── Run wiring ────────────────────────────────────────────────────────────
1756
+ _config_inputs = [
1757
+ llm_choice, judge_llm_choice, embed_choice,
1758
+ enable_hybrid, enable_hyde, enable_reranking, reranker_type,
1759
+ enable_rrf, rrf_k,
1760
+ enable_repacking, repack_strategy,
1761
+ enable_summarization, summarization_type,
1762
+ prompt_strategy, hybrid_alpha, top_k,
1763
+ enable_query_classification, enable_query_rewriting, enable_query_decomp,
1764
+ run_judge,
1765
+ ]
1766
+ _all_inputs = [query_input, domain_dd, dataset_dd, sample_dd] + _config_inputs
1767
+ _all_outputs = [gt_out, rag_out, gt_docs_out, rag_docs_out, metrics_out]
1768
+
1769
+ run_btn.click(fn=run_query, inputs=_all_inputs, outputs=_all_outputs)
1770
+ query_input.submit(fn=run_query, inputs=_all_inputs, outputs=_all_outputs)
1771
+
1772
+ demo.launch(
1773
+ share=True,
1774
+ debug=True,
1775
+ theme=gr.themes.Soft(),
1776
+ css=CSS,
1777
+ )
1778
+
requirements.txt ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ numpy
2
+ pandas
3
+ pymilvus
4
+ groq
5
+ openai
6
+ sentence_transformers
7
+ rank_bm25
8
+ scikit-learn
9
+ torch
10
+ transformers
11
+ huggingface_hub
12
+ datasets
13
+ gradio
14
+ llmlingua
15
+ milvus-lite