quantumbit commited on
Commit
6504f2b
·
verified ·
1 Parent(s): 4013785

Upload folder using huggingface_hub

Browse files
rag_system/cache.py CHANGED
@@ -101,6 +101,10 @@ def _exact_key(query: str, collection_key: str, params_key: str) -> str:
101
  return EXACT_PREFIX + hashlib.sha256(payload.encode()).hexdigest()[:32]
102
 
103
 
 
 
 
 
104
  def _vector_bytes(vec: list[float]) -> bytes:
105
  return np.array(vec, dtype=np.float32).tobytes()
106
 
@@ -146,9 +150,10 @@ def get_semantic(query_vec: list[float], collection_key: str, params_key: str) -
146
 
147
  vec = _vector_bytes(query_vec)
148
  k = 4
 
149
  query = (
150
  f"@{COLLECTION_FIELD}:{{{collection_key}}} "
151
- f"@{PARAMS_FIELD}:{{{params_key}}}"
152
  f"=>[KNN {k} @{VECTOR_FIELD} $vec AS score]"
153
  )
154
 
@@ -223,7 +228,7 @@ def set_semantic(
223
  mapping={
224
  "query": query,
225
  COLLECTION_FIELD: collection_key,
226
- PARAMS_FIELD: params_key,
227
  VECTOR_FIELD: _vector_bytes(query_vec),
228
  PAYLOAD_FIELD: payload,
229
  },
 
101
  return EXACT_PREFIX + hashlib.sha256(payload.encode()).hexdigest()[:32]
102
 
103
 
104
+ def _tag_hash(value: str) -> str:
105
+ return hashlib.sha1(value.encode()).hexdigest()[:16]
106
+
107
+
108
  def _vector_bytes(vec: list[float]) -> bytes:
109
  return np.array(vec, dtype=np.float32).tobytes()
110
 
 
150
 
151
  vec = _vector_bytes(query_vec)
152
  k = 4
153
+ params_tag = _tag_hash(params_key)
154
  query = (
155
  f"@{COLLECTION_FIELD}:{{{collection_key}}} "
156
+ f"@{PARAMS_FIELD}:{{{params_tag}}} "
157
  f"=>[KNN {k} @{VECTOR_FIELD} $vec AS score]"
158
  )
159
 
 
228
  mapping={
229
  "query": query,
230
  COLLECTION_FIELD: collection_key,
231
+ PARAMS_FIELD: _tag_hash(params_key),
232
  VECTOR_FIELD: _vector_bytes(query_vec),
233
  PAYLOAD_FIELD: payload,
234
  },
rag_system/models.py CHANGED
@@ -47,6 +47,10 @@ class QueryRequest(BaseModel):
47
  retrieval_mode: RetrievalMode = RetrievalMode.HYBRID
48
  embedding_mode: Optional[EmbeddingMode] = None
49
  top_k: Optional[int] = None
 
 
 
 
50
  doc_collections: Optional[list[str]] = None # per-doc sub-collections; None = legacy single-collection mode
51
  history: list[ChatMessage] = Field(default_factory=list)
52
  stream: bool = False
@@ -56,6 +60,24 @@ class QueryRequest(BaseModel):
56
  def sanitize_query(cls,v):
57
  return v.strip()
58
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
59
  class SourceDocument(BaseModel):
60
  doc_id: str
61
  content: str
 
47
  retrieval_mode: RetrievalMode = RetrievalMode.HYBRID
48
  embedding_mode: Optional[EmbeddingMode] = None
49
  top_k: Optional[int] = None
50
+ top_k_retrieval: Optional[int] = None
51
+ mmr_lambda: Optional[float] = None
52
+ bm25_weight: Optional[float] = None
53
+ vector_weight: Optional[float] = None
54
  doc_collections: Optional[list[str]] = None # per-doc sub-collections; None = legacy single-collection mode
55
  history: list[ChatMessage] = Field(default_factory=list)
56
  stream: bool = False
 
60
  def sanitize_query(cls,v):
61
  return v.strip()
62
 
63
+ @field_validator("top_k", "top_k_retrieval")
64
+ @classmethod
65
+ def validate_top_k(cls, v):
66
+ if v is None:
67
+ return v
68
+ if v < 1:
69
+ raise ValueError("top_k values must be >= 1")
70
+ return v
71
+
72
+ @field_validator("mmr_lambda", "bm25_weight", "vector_weight")
73
+ @classmethod
74
+ def validate_weights(cls, v):
75
+ if v is None:
76
+ return v
77
+ if v < 0 or v > 1:
78
+ raise ValueError("weights must be between 0 and 1")
79
+ return v
80
+
81
  class SourceDocument(BaseModel):
82
  doc_id: str
83
  content: str
rag_system/query_engine.py CHANGED
@@ -71,9 +71,26 @@ def _cache_collection_key(collections: list[str]) -> str:
71
  return hashlib.sha1(raw.encode()).hexdigest()[:16]
72
 
73
 
74
- def _cache_params_key(mode: str, top_k: Optional[int]) -> str:
75
- k = top_k if top_k is not None else settings.top_k_rerank
76
- return f"{mode}:{k}"
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
77
 
78
  # Query rewriting
79
  async def rewrite_query(query: str) -> str:
@@ -195,7 +212,14 @@ async def query(
195
  mode_val = request.retrieval_mode.value if hasattr(request.retrieval_mode, "value") else str(request.retrieval_mode)
196
  cache_allowed = settings.cache_enabled and _is_try_docs_scope(collections)
197
  cache_collection_key = _cache_collection_key(collections)
198
- cache_params_key = _cache_params_key(mode_val, request.top_k)
 
 
 
 
 
 
 
199
  cache_query_vec = None
200
 
201
  # 1. Input guardrail
@@ -250,6 +274,10 @@ async def query(
250
  collections=scoped,
251
  mode=request.retrieval_mode.value if hasattr(request.retrieval_mode, "value") else str(request.retrieval_mode),
252
  k_per_collection=k_per,
 
 
 
 
253
  use_reranker=True,
254
  expand_context=True,
255
  )
@@ -260,6 +288,10 @@ async def query(
260
  collection=collections[0],
261
  mode=request.retrieval_mode,
262
  top_k=request.top_k,
 
 
 
 
263
  use_reranker=True,
264
  expand_context=True,
265
  )
@@ -368,7 +400,14 @@ async def pipeline_stream_query(request: QueryRequest) -> AsyncIterator[str]:
368
  embedding_mode = resolve_embedding_mode_for_collections(collections, request.embedding_mode)
369
  cache_allowed = settings.cache_enabled and _is_try_docs_scope(collections)
370
  cache_collection_key = _cache_collection_key(collections)
371
- cache_params_key = _cache_params_key(mode_val, request.top_k)
 
 
 
 
 
 
 
372
 
373
  yield emit("pipeline_start", "in_progress", {
374
  "query": request.query,
@@ -451,7 +490,8 @@ async def pipeline_stream_query(request: QueryRequest) -> AsyncIterator[str]:
451
  # --- Retrieval ---
452
  yield emit("retrieval_start", "in_progress", {
453
  "mode": mode_val,
454
- "top_k": request.top_k or settings.top_k_retrieval,
 
455
  "collections": len(scoped),
456
  })
457
 
@@ -462,6 +502,10 @@ async def pipeline_stream_query(request: QueryRequest) -> AsyncIterator[str]:
462
  collections=scoped,
463
  mode=mode_val,
464
  k_per_collection=k_per,
 
 
 
 
465
  use_reranker=True,
466
  expand_context=True,
467
  )
@@ -471,6 +515,10 @@ async def pipeline_stream_query(request: QueryRequest) -> AsyncIterator[str]:
471
  collection=scoped[0],
472
  mode=request.retrieval_mode,
473
  top_k=request.top_k,
 
 
 
 
474
  use_reranker=True,
475
  expand_context=True,
476
  )
 
71
  return hashlib.sha1(raw.encode()).hexdigest()[:16]
72
 
73
 
74
+ def _fmt_param(value: Optional[float]) -> str:
75
+ if value is None:
76
+ return "-"
77
+ return f"{value:.3f}"
78
+
79
+
80
+ def _cache_params_key_v2(
81
+ mode: str,
82
+ top_k: Optional[int],
83
+ top_k_retrieval: Optional[int],
84
+ mmr_lambda: Optional[float],
85
+ bm25_weight: Optional[float],
86
+ vector_weight: Optional[float],
87
+ ) -> str:
88
+ k_final = top_k if top_k is not None else settings.top_k_rerank
89
+ k_retrieve = top_k_retrieval if top_k_retrieval is not None else settings.top_k_retrieval
90
+ return (
91
+ f"{mode}:{k_final}:{k_retrieve}:"
92
+ f"{_fmt_param(mmr_lambda)}:{_fmt_param(bm25_weight)}:{_fmt_param(vector_weight)}"
93
+ )
94
 
95
  # Query rewriting
96
  async def rewrite_query(query: str) -> str:
 
212
  mode_val = request.retrieval_mode.value if hasattr(request.retrieval_mode, "value") else str(request.retrieval_mode)
213
  cache_allowed = settings.cache_enabled and _is_try_docs_scope(collections)
214
  cache_collection_key = _cache_collection_key(collections)
215
+ cache_params_key = _cache_params_key_v2(
216
+ mode_val,
217
+ request.top_k,
218
+ request.top_k_retrieval,
219
+ request.mmr_lambda,
220
+ request.bm25_weight,
221
+ request.vector_weight,
222
+ )
223
  cache_query_vec = None
224
 
225
  # 1. Input guardrail
 
274
  collections=scoped,
275
  mode=request.retrieval_mode.value if hasattr(request.retrieval_mode, "value") else str(request.retrieval_mode),
276
  k_per_collection=k_per,
277
+ top_k_retrieval=request.top_k_retrieval,
278
+ mmr_lambda=request.mmr_lambda,
279
+ bm25_weight=request.bm25_weight,
280
+ vector_weight=request.vector_weight,
281
  use_reranker=True,
282
  expand_context=True,
283
  )
 
288
  collection=collections[0],
289
  mode=request.retrieval_mode,
290
  top_k=request.top_k,
291
+ top_k_retrieval=request.top_k_retrieval,
292
+ mmr_lambda=request.mmr_lambda,
293
+ bm25_weight=request.bm25_weight,
294
+ vector_weight=request.vector_weight,
295
  use_reranker=True,
296
  expand_context=True,
297
  )
 
400
  embedding_mode = resolve_embedding_mode_for_collections(collections, request.embedding_mode)
401
  cache_allowed = settings.cache_enabled and _is_try_docs_scope(collections)
402
  cache_collection_key = _cache_collection_key(collections)
403
+ cache_params_key = _cache_params_key_v2(
404
+ mode_val,
405
+ request.top_k,
406
+ request.top_k_retrieval,
407
+ request.mmr_lambda,
408
+ request.bm25_weight,
409
+ request.vector_weight,
410
+ )
411
 
412
  yield emit("pipeline_start", "in_progress", {
413
  "query": request.query,
 
490
  # --- Retrieval ---
491
  yield emit("retrieval_start", "in_progress", {
492
  "mode": mode_val,
493
+ "top_k": request.top_k or settings.top_k_rerank,
494
+ "top_k_retrieval": request.top_k_retrieval or settings.top_k_retrieval,
495
  "collections": len(scoped),
496
  })
497
 
 
502
  collections=scoped,
503
  mode=mode_val,
504
  k_per_collection=k_per,
505
+ top_k_retrieval=request.top_k_retrieval,
506
+ mmr_lambda=request.mmr_lambda,
507
+ bm25_weight=request.bm25_weight,
508
+ vector_weight=request.vector_weight,
509
  use_reranker=True,
510
  expand_context=True,
511
  )
 
515
  collection=scoped[0],
516
  mode=request.retrieval_mode,
517
  top_k=request.top_k,
518
+ top_k_retrieval=request.top_k_retrieval,
519
+ mmr_lambda=request.mmr_lambda,
520
+ bm25_weight=request.bm25_weight,
521
+ vector_weight=request.vector_weight,
522
  use_reranker=True,
523
  expand_context=True,
524
  )
rag_system/retriever.py CHANGED
@@ -49,13 +49,30 @@ def bm25_retrieve(query: str, collection: str,k: int) -> list[tuple[Document,flo
49
  def _rrf_score(rank: int, k: int = 60) -> float:
50
  return 1.0 / (k + rank + 1) #here 1 is added to handle rank 1 which here comes as 0
51
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
52
  #Hybrid Retrieval
53
  def hybrid_retrieve(
54
  query: str,
55
  collection: str,
56
- k: int
 
 
57
  ) -> list[tuple[Document,float]]:
58
  """Reciprocal Rank fusion of BM25 and Vector results"""
 
59
  pool_size = k*3 #casting a wide net before fusing
60
  vec_results = similarity_search_with_scores(query,collection,k=pool_size)
61
  bm25_results = bm25_retrieve(query,collection,k=pool_size)
@@ -65,12 +82,12 @@ def hybrid_retrieve(
65
 
66
  for rank, (doc, _) in enumerate(vec_results):
67
  did = doc.metadata.get("doc_id",id(doc))
68
- rrf_scores[did] = rrf_scores.get(did,0) + settings.vector_weight * _rrf_score(rank) #check the existing score first and then add the fresh score
69
  doc_map[did] = doc
70
 
71
  for rank, (doc, _) in enumerate(bm25_results):
72
  did = doc.metadata.get("doc_id",id(doc))
73
- rrf_scores[did] = rrf_scores.get(did,0) + settings.bm25_weight * _rrf_score(rank)
74
  doc_map[did] = doc
75
 
76
  sorted_ids = sorted(rrf_scores, key=lambda x: rrf_scores[x], reverse=True)[:k]
@@ -81,14 +98,15 @@ async def mmr_retrieve(
81
  query: str,
82
  collection: str,
83
  k: int,
84
- lambda_mult: float = None,
 
85
  ) -> list[tuple[Document, float]]:
86
  """
87
  Maximal Marginal Relevance: balance relevance vs diversity.
88
  """
89
  # 1. Setup parameters
90
  lam = lambda_mult or settings.mmr_lambda
91
- fetch_k = settings.top_k_retrieval
92
  store = get_store(collection)
93
 
94
  # 2. Get original scores to map them back later
@@ -195,10 +213,14 @@ async def retrieve(
195
  collection: str = "default",
196
  mode: str = "hybrid",
197
  top_k: Optional[int] = None,
 
 
 
 
198
  use_reranker: bool = True,
199
  expand_context: bool = True,
200
  ) -> list[tuple[Document,float]]:
201
- k_retrieve = settings.top_k_retrieval
202
  k_final = top_k or settings.top_k_rerank
203
 
204
  if mode == "vector":
@@ -206,10 +228,16 @@ async def retrieve(
206
  elif mode == "bm25":
207
  results = bm25_retrieve(query,collection,k=k_retrieve)
208
  elif mode == "mmr":
209
- results = await mmr_retrieve(query,collection,k=k_final)
210
  return results # MMR alrady handles diversity , skip reranker
211
  else: #go with hybrid
212
- results = hybrid_retrieve(query,collection,k=k_retrieve)
 
 
 
 
 
 
213
 
214
  #Rerank
215
  if use_reranker:
@@ -259,6 +287,10 @@ async def multi_collection_retrieve(
259
  collections: list[str],
260
  mode: str = "hybrid",
261
  k_per_collection: int = 5,
 
 
 
 
262
  use_reranker: bool = True,
263
  expand_context: bool = True,
264
  ) -> list[tuple[Document, float]]:
@@ -275,6 +307,10 @@ async def multi_collection_retrieve(
275
  collection=coll,
276
  mode=mode,
277
  top_k=k_per_collection,
 
 
 
 
278
  use_reranker=False, # defer reranking until after merge
279
  expand_context=expand_context,
280
  )
 
49
  def _rrf_score(rank: int, k: int = 60) -> float:
50
  return 1.0 / (k + rank + 1) #here 1 is added to handle rank 1 which here comes as 0
51
 
52
+
53
+ def _resolve_hybrid_weights(
54
+ bm25_weight: Optional[float],
55
+ vector_weight: Optional[float],
56
+ ) -> tuple[float, float]:
57
+ bw = settings.bm25_weight if bm25_weight is None else bm25_weight
58
+ vw = settings.vector_weight if vector_weight is None else vector_weight
59
+ total = bw + vw
60
+ if total <= 0:
61
+ bw = settings.bm25_weight
62
+ vw = settings.vector_weight
63
+ total = bw + vw
64
+ return bw / total, vw / total
65
+
66
  #Hybrid Retrieval
67
  def hybrid_retrieve(
68
  query: str,
69
  collection: str,
70
+ k: int,
71
+ bm25_weight: Optional[float] = None,
72
+ vector_weight: Optional[float] = None,
73
  ) -> list[tuple[Document,float]]:
74
  """Reciprocal Rank fusion of BM25 and Vector results"""
75
+ bw, vw = _resolve_hybrid_weights(bm25_weight, vector_weight)
76
  pool_size = k*3 #casting a wide net before fusing
77
  vec_results = similarity_search_with_scores(query,collection,k=pool_size)
78
  bm25_results = bm25_retrieve(query,collection,k=pool_size)
 
82
 
83
  for rank, (doc, _) in enumerate(vec_results):
84
  did = doc.metadata.get("doc_id",id(doc))
85
+ rrf_scores[did] = rrf_scores.get(did,0) + vw * _rrf_score(rank) #check the existing score first and then add the fresh score
86
  doc_map[did] = doc
87
 
88
  for rank, (doc, _) in enumerate(bm25_results):
89
  did = doc.metadata.get("doc_id",id(doc))
90
+ rrf_scores[did] = rrf_scores.get(did,0) + bw * _rrf_score(rank)
91
  doc_map[did] = doc
92
 
93
  sorted_ids = sorted(rrf_scores, key=lambda x: rrf_scores[x], reverse=True)[:k]
 
98
  query: str,
99
  collection: str,
100
  k: int,
101
+ lambda_mult: Optional[float] = None,
102
+ fetch_k: Optional[int] = None,
103
  ) -> list[tuple[Document, float]]:
104
  """
105
  Maximal Marginal Relevance: balance relevance vs diversity.
106
  """
107
  # 1. Setup parameters
108
  lam = lambda_mult or settings.mmr_lambda
109
+ fetch_k = fetch_k or settings.top_k_retrieval
110
  store = get_store(collection)
111
 
112
  # 2. Get original scores to map them back later
 
213
  collection: str = "default",
214
  mode: str = "hybrid",
215
  top_k: Optional[int] = None,
216
+ top_k_retrieval: Optional[int] = None,
217
+ mmr_lambda: Optional[float] = None,
218
+ bm25_weight: Optional[float] = None,
219
+ vector_weight: Optional[float] = None,
220
  use_reranker: bool = True,
221
  expand_context: bool = True,
222
  ) -> list[tuple[Document,float]]:
223
+ k_retrieve = top_k_retrieval or settings.top_k_retrieval
224
  k_final = top_k or settings.top_k_rerank
225
 
226
  if mode == "vector":
 
228
  elif mode == "bm25":
229
  results = bm25_retrieve(query,collection,k=k_retrieve)
230
  elif mode == "mmr":
231
+ results = await mmr_retrieve(query,collection,k=k_final,lambda_mult=mmr_lambda,fetch_k=k_retrieve)
232
  return results # MMR alrady handles diversity , skip reranker
233
  else: #go with hybrid
234
+ results = hybrid_retrieve(
235
+ query,
236
+ collection,
237
+ k=k_retrieve,
238
+ bm25_weight=bm25_weight,
239
+ vector_weight=vector_weight,
240
+ )
241
 
242
  #Rerank
243
  if use_reranker:
 
287
  collections: list[str],
288
  mode: str = "hybrid",
289
  k_per_collection: int = 5,
290
+ top_k_retrieval: Optional[int] = None,
291
+ mmr_lambda: Optional[float] = None,
292
+ bm25_weight: Optional[float] = None,
293
+ vector_weight: Optional[float] = None,
294
  use_reranker: bool = True,
295
  expand_context: bool = True,
296
  ) -> list[tuple[Document, float]]:
 
307
  collection=coll,
308
  mode=mode,
309
  top_k=k_per_collection,
310
+ top_k_retrieval=top_k_retrieval,
311
+ mmr_lambda=mmr_lambda,
312
+ bm25_weight=bm25_weight,
313
+ vector_weight=vector_weight,
314
  use_reranker=False, # defer reranking until after merge
315
  expand_context=expand_context,
316
  )